k3-node 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- k3_node/__init__.py +122 -0
- k3_node/applications/__init__.py +17 -0
- k3_node/applications/bio/__init__.py +21 -0
- k3_node/applications/chemistry/__init__.py +155 -0
- k3_node/applications/materials/__init__.py +127 -0
- k3_node/applications/materials/basis.py +449 -0
- k3_node/applications/materials/chgnet.py +360 -0
- k3_node/applications/materials/core.py +351 -0
- k3_node/applications/materials/grace.py +246 -0
- k3_node/applications/materials/io.py +230 -0
- k3_node/applications/materials/m3gnet.py +462 -0
- k3_node/applications/materials/megnet.py +395 -0
- k3_node/applications/materials/qet.py +220 -0
- k3_node/applications/materials/readout.py +235 -0
- k3_node/applications/materials/so3net.py +234 -0
- k3_node/applications/materials/tensornet.py +381 -0
- k3_node/applications/materials/test_materials.py +167 -0
- k3_node/applications/materials/wrappers.py +95 -0
- k3_node/data/__init__.py +47 -0
- k3_node/data/batch.py +102 -0
- k3_node/data/collate.py +282 -0
- k3_node/data/data.py +532 -0
- k3_node/data/database.py +154 -0
- k3_node/data/dataset.py +182 -0
- k3_node/data/download.py +49 -0
- k3_node/data/extract.py +45 -0
- k3_node/data/feature_store.py +70 -0
- k3_node/data/graph_store.py +92 -0
- k3_node/data/hetero_data.py +374 -0
- k3_node/data/hypergraph_data.py +59 -0
- k3_node/data/in_memory_dataset.py +177 -0
- k3_node/data/makedirs.py +7 -0
- k3_node/data/on_disk_dataset.py +77 -0
- k3_node/data/separate.py +115 -0
- k3_node/data/storage.py +593 -0
- k3_node/data/temporal.py +154 -0
- k3_node/data/test_batch.py +67 -0
- k3_node/data/test_data.py +68 -0
- k3_node/data/test_dataset_and_stores.py +111 -0
- k3_node/data/test_hetero_data.py +33 -0
- k3_node/data/test_temporal_and_hyper.py +32 -0
- k3_node/data/view.py +43 -0
- k3_node/datasets/__init__.py +88 -0
- k3_node/datasets/actor.py +101 -0
- k3_node/datasets/airports.py +84 -0
- k3_node/datasets/amazon.py +66 -0
- k3_node/datasets/ba2motif_dataset.py +73 -0
- k3_node/datasets/ba_shapes.py +81 -0
- k3_node/datasets/bitcoin_otc.py +77 -0
- k3_node/datasets/citation_full.py +81 -0
- k3_node/datasets/coauthor.py +66 -0
- k3_node/datasets/dblp.py +106 -0
- k3_node/datasets/digits.py +63 -0
- k3_node/datasets/email_eu_core.py +60 -0
- k3_node/datasets/entities.py +158 -0
- k3_node/datasets/explainer_dataset.py +101 -0
- k3_node/datasets/facebook.py +51 -0
- k3_node/datasets/fake.py +256 -0
- k3_node/datasets/freebase.py +90 -0
- k3_node/datasets/geometric_shapes.py +69 -0
- k3_node/datasets/github.py +51 -0
- k3_node/datasets/graph_generator/__init__.py +6 -0
- k3_node/datasets/graph_generator/ba_graph.py +20 -0
- k3_node/datasets/graph_generator/base.py +29 -0
- k3_node/datasets/graph_generator/er_graph.py +21 -0
- k3_node/datasets/icews.py +58 -0
- k3_node/datasets/imdb.py +96 -0
- k3_node/datasets/jodie.py +56 -0
- k3_node/datasets/karate.py +56 -0
- k3_node/datasets/lastfm_asia.py +51 -0
- k3_node/datasets/mesh_correspondence.py +50 -0
- k3_node/datasets/molecule_net.py +148 -0
- k3_node/datasets/motif_generator/__init__.py +7 -0
- k3_node/datasets/motif_generator/base.py +29 -0
- k3_node/datasets/motif_generator/custom.py +17 -0
- k3_node/datasets/motif_generator/cycle.py +25 -0
- k3_node/datasets/motif_generator/house.py +27 -0
- k3_node/datasets/movielens.py +55 -0
- k3_node/datasets/planetoid.py +137 -0
- k3_node/datasets/polblogs.py +63 -0
- k3_node/datasets/ppi.py +189 -0
- k3_node/datasets/qm7.py +65 -0
- k3_node/datasets/qm9.py +132 -0
- k3_node/datasets/reddit.py +121 -0
- k3_node/datasets/sbm_dataset.py +165 -0
- k3_node/datasets/seal.py +74 -0
- k3_node/datasets/shape_scenes.py +92 -0
- k3_node/datasets/test_datasets.py +322 -0
- k3_node/datasets/tu_dataset.py +131 -0
- k3_node/datasets/twitch.py +66 -0
- k3_node/datasets/webkb.py +102 -0
- k3_node/datasets/wikics.py +85 -0
- k3_node/datasets/word_net.py +184 -0
- k3_node/etl/__init__.py +37 -0
- k3_node/etl/encoders.py +248 -0
- k3_node/etl/graph_builders.py +270 -0
- k3_node/etl/relational_to_graph.py +201 -0
- k3_node/etl/table_to_graph.py +244 -0
- k3_node/etl/test_etl.py +318 -0
- k3_node/export/__init__.py +15 -0
- k3_node/export/cross_backend.py +172 -0
- k3_node/export/onnx_exporter.py +190 -0
- k3_node/export/runtime.py +254 -0
- k3_node/export/tensorrt_exporter.py +201 -0
- k3_node/export/test_export.py +337 -0
- k3_node/export/tflite_exporter.py +112 -0
- k3_node/hub/__init__.py +29 -0
- k3_node/hub/dataset_hub.py +242 -0
- k3_node/hub/hub_mixin.py +599 -0
- k3_node/hub/model_card.py +133 -0
- k3_node/hub/test_hub.py +419 -0
- k3_node/io/__init__.py +22 -0
- k3_node/io/fs.py +117 -0
- k3_node/io/npz.py +45 -0
- k3_node/io/off.py +29 -0
- k3_node/io/planetoid.py +98 -0
- k3_node/io/tu.py +137 -0
- k3_node/io/txt_array.py +58 -0
- k3_node/layers/__init__.py +14 -0
- k3_node/layers/aggr/__init__.py +70 -0
- k3_node/layers/aggr/attention.py +77 -0
- k3_node/layers/aggr/base.py +403 -0
- k3_node/layers/aggr/basic.py +412 -0
- k3_node/layers/aggr/deep_sets.py +65 -0
- k3_node/layers/aggr/deepsets.py +29 -0
- k3_node/layers/aggr/equilibrium.py +107 -0
- k3_node/layers/aggr/fused.py +43 -0
- k3_node/layers/aggr/gmt.py +89 -0
- k3_node/layers/aggr/gru.py +58 -0
- k3_node/layers/aggr/lcm.py +143 -0
- k3_node/layers/aggr/lstm.py +58 -0
- k3_node/layers/aggr/mlp.py +75 -0
- k3_node/layers/aggr/multi.py +154 -0
- k3_node/layers/aggr/patch_transformer.py +137 -0
- k3_node/layers/aggr/quantile.py +125 -0
- k3_node/layers/aggr/resolver.py +68 -0
- k3_node/layers/aggr/scaler.py +133 -0
- k3_node/layers/aggr/set2set.py +87 -0
- k3_node/layers/aggr/set_transformer.py +107 -0
- k3_node/layers/aggr/sort.py +68 -0
- k3_node/layers/aggr/test_aggr.py +337 -0
- k3_node/layers/aggr/utils.py +210 -0
- k3_node/layers/aggr/variance_preserving.py +54 -0
- k3_node/layers/attention/__init__.py +5 -0
- k3_node/layers/attention/pair_attention.py +448 -0
- k3_node/layers/attention/performer.py +187 -0
- k3_node/layers/attention/polynormer.py +160 -0
- k3_node/layers/attention/qformer.py +143 -0
- k3_node/layers/attention/sgformer.py +106 -0
- k3_node/layers/attention/test_attention.py +68 -0
- k3_node/layers/attention/test_pair_attention.py +91 -0
- k3_node/layers/conv/__init__.py +149 -0
- k3_node/layers/conv/agnn_conv.py +120 -0
- k3_node/layers/conv/antisymmetric_conv.py +94 -0
- k3_node/layers/conv/appnp.py +105 -0
- k3_node/layers/conv/appnp_conv.py +157 -0
- k3_node/layers/conv/arma_conv.py +231 -0
- k3_node/layers/conv/cg_conv.py +92 -0
- k3_node/layers/conv/cheb_conv.py +137 -0
- k3_node/layers/conv/cluster_gcn_conv.py +102 -0
- k3_node/layers/conv/conv.py +100 -0
- k3_node/layers/conv/crystal_conv.py +140 -0
- k3_node/layers/conv/cugraph.py +84 -0
- k3_node/layers/conv/diffusion_conv.py +144 -0
- k3_node/layers/conv/dir_gnn_conv.py +93 -0
- k3_node/layers/conv/dna_conv.py +192 -0
- k3_node/layers/conv/edge_conv.py +107 -0
- k3_node/layers/conv/eg_conv.py +155 -0
- k3_node/layers/conv/fa_conv.py +107 -0
- k3_node/layers/conv/feast_conv.py +126 -0
- k3_node/layers/conv/film_conv.py +143 -0
- k3_node/layers/conv/gat_conv.py +244 -0
- k3_node/layers/conv/gated_graph_conv.py +136 -0
- k3_node/layers/conv/gatv2_conv.py +205 -0
- k3_node/layers/conv/gcn.py +144 -0
- k3_node/layers/conv/gcn2_conv.py +126 -0
- k3_node/layers/conv/gcn_conv.py +135 -0
- k3_node/layers/conv/gen_conv.py +163 -0
- k3_node/layers/conv/general_conv.py +218 -0
- k3_node/layers/conv/gin_conv.py +218 -0
- k3_node/layers/conv/gmm_conv.py +172 -0
- k3_node/layers/conv/gps_conv.py +153 -0
- k3_node/layers/conv/graph_attention.py +262 -0
- k3_node/layers/conv/graph_conv.py +84 -0
- k3_node/layers/conv/gravnet_conv.py +93 -0
- k3_node/layers/conv/han_conv.py +175 -0
- k3_node/layers/conv/heat_conv.py +131 -0
- k3_node/layers/conv/hetero_conv.py +128 -0
- k3_node/layers/conv/hgt_conv.py +218 -0
- k3_node/layers/conv/hypergraph_conv.py +182 -0
- k3_node/layers/conv/le_conv.py +81 -0
- k3_node/layers/conv/lg_conv.py +58 -0
- k3_node/layers/conv/meshcnn_conv.py +84 -0
- k3_node/layers/conv/message_passing.py +451 -0
- k3_node/layers/conv/mf_conv.py +95 -0
- k3_node/layers/conv/mixhop_conv.py +108 -0
- k3_node/layers/conv/nn_conv.py +110 -0
- k3_node/layers/conv/pan_conv.py +100 -0
- k3_node/layers/conv/pdn_conv.py +109 -0
- k3_node/layers/conv/pna_conv.py +177 -0
- k3_node/layers/conv/point_conv.py +101 -0
- k3_node/layers/conv/point_gnn_conv.py +90 -0
- k3_node/layers/conv/point_transformer_conv.py +132 -0
- k3_node/layers/conv/ppf_conv.py +135 -0
- k3_node/layers/conv/ppnp.py +89 -0
- k3_node/layers/conv/res_gated_graph_conv.py +126 -0
- k3_node/layers/conv/rgat_conv.py +251 -0
- k3_node/layers/conv/rgcn_conv.py +321 -0
- k3_node/layers/conv/sage_conv.py +154 -0
- k3_node/layers/conv/sg_conv.py +96 -0
- k3_node/layers/conv/signed_conv.py +100 -0
- k3_node/layers/conv/simple_conv.py +75 -0
- k3_node/layers/conv/spline_conv.py +182 -0
- k3_node/layers/conv/ssg_conv.py +101 -0
- k3_node/layers/conv/supergat_conv.py +195 -0
- k3_node/layers/conv/tag_conv.py +98 -0
- k3_node/layers/conv/test_backend_consistency.py +164 -0
- k3_node/layers/conv/test_conv.py +176 -0
- k3_node/layers/conv/test_conv_pyg.py +566 -0
- k3_node/layers/conv/transformer_conv.py +168 -0
- k3_node/layers/conv/utils.py +403 -0
- k3_node/layers/conv/wl_conv.py +151 -0
- k3_node/layers/conv/x_conv.py +187 -0
- k3_node/layers/dense/__init__.py +40 -0
- k3_node/layers/dense/dense_gat_conv.py +149 -0
- k3_node/layers/dense/dense_gcn_conv.py +117 -0
- k3_node/layers/dense/dense_gin_conv.py +88 -0
- k3_node/layers/dense/dense_graph_conv.py +95 -0
- k3_node/layers/dense/dense_sage_conv.py +85 -0
- k3_node/layers/dense/diff_pool.py +76 -0
- k3_node/layers/dense/dmon_pool.py +223 -0
- k3_node/layers/dense/linear.py +327 -0
- k3_node/layers/dense/mincut_pool.py +92 -0
- k3_node/layers/dense/test_dense.py +377 -0
- k3_node/layers/functional/__init__.py +13 -0
- k3_node/layers/functional/bro.py +49 -0
- k3_node/layers/functional/edge_dropout.py +55 -0
- k3_node/layers/functional/gini.py +44 -0
- k3_node/layers/functional/test_functional.py +34 -0
- k3_node/layers/kge/__init__.py +17 -0
- k3_node/layers/kge/base.py +255 -0
- k3_node/layers/kge/complex.py +98 -0
- k3_node/layers/kge/distmult.py +79 -0
- k3_node/layers/kge/loader.py +50 -0
- k3_node/layers/kge/rotate.py +103 -0
- k3_node/layers/kge/test_kge.py +76 -0
- k3_node/layers/kge/transe.py +96 -0
- k3_node/layers/norm/__init__.py +23 -0
- k3_node/layers/norm/batch_norm.py +328 -0
- k3_node/layers/norm/diff_group_norm.py +141 -0
- k3_node/layers/norm/graph_norm.py +105 -0
- k3_node/layers/norm/graph_size_norm.py +57 -0
- k3_node/layers/norm/instance_norm.py +163 -0
- k3_node/layers/norm/layer_norm.py +245 -0
- k3_node/layers/norm/mean_subtraction_norm.py +57 -0
- k3_node/layers/norm/msg_norm.py +58 -0
- k3_node/layers/norm/pair_norm.py +94 -0
- k3_node/layers/norm/test_norm.py +275 -0
- k3_node/layers/pool/__init__.py +83 -0
- k3_node/layers/pool/approx_knn.py +101 -0
- k3_node/layers/pool/asap.py +173 -0
- k3_node/layers/pool/avg_pool.py +165 -0
- k3_node/layers/pool/cluster_pool.py +168 -0
- k3_node/layers/pool/connect/__init__.py +10 -0
- k3_node/layers/pool/connect/base.py +103 -0
- k3_node/layers/pool/connect/filter_edges.py +113 -0
- k3_node/layers/pool/consecutive.py +30 -0
- k3_node/layers/pool/decimation.py +48 -0
- k3_node/layers/pool/edge_pool.py +189 -0
- k3_node/layers/pool/glob.py +139 -0
- k3_node/layers/pool/graclus.py +66 -0
- k3_node/layers/pool/knn.py +253 -0
- k3_node/layers/pool/max_pool.py +159 -0
- k3_node/layers/pool/mem_pool.py +145 -0
- k3_node/layers/pool/pan_pool.py +144 -0
- k3_node/layers/pool/point_cloud.py +212 -0
- k3_node/layers/pool/pool.py +119 -0
- k3_node/layers/pool/sag_pool.py +174 -0
- k3_node/layers/pool/select/__init__.py +10 -0
- k3_node/layers/pool/select/base.py +112 -0
- k3_node/layers/pool/select/topk.py +206 -0
- k3_node/layers/pool/test_pool.py +456 -0
- k3_node/layers/pool/topk_pool.py +103 -0
- k3_node/layers/pool/voxel_grid.py +70 -0
- k3_node/layers/unpool/__init__.py +9 -0
- k3_node/layers/unpool/knn_interpolate.py +57 -0
- k3_node/layers/unpool/test_unpool.py +31 -0
- k3_node/loader/__init__.py +62 -0
- k3_node/loader/base.py +69 -0
- k3_node/loader/cache.py +68 -0
- k3_node/loader/cluster.py +127 -0
- k3_node/loader/data_list_loader.py +45 -0
- k3_node/loader/dataloader.py +117 -0
- k3_node/loader/dense_data_loader.py +62 -0
- k3_node/loader/dynamic_batch_sampler.py +93 -0
- k3_node/loader/graph_saint.py +188 -0
- k3_node/loader/hgt_loader.py +90 -0
- k3_node/loader/imbalanced_sampler.py +87 -0
- k3_node/loader/keras_dataset.py +334 -0
- k3_node/loader/link_loader.py +179 -0
- k3_node/loader/link_neighbor_loader.py +202 -0
- k3_node/loader/mixin.py +190 -0
- k3_node/loader/neighbor_loader.py +159 -0
- k3_node/loader/neighbor_sampler.py +167 -0
- k3_node/loader/node_loader.py +185 -0
- k3_node/loader/prefetch.py +115 -0
- k3_node/loader/random_node_loader.py +89 -0
- k3_node/loader/sampler_utils.py +499 -0
- k3_node/loader/shadow.py +115 -0
- k3_node/loader/temporal_dataloader.py +98 -0
- k3_node/loader/test_dataloader.py +113 -0
- k3_node/loader/test_keras_dataset.py +221 -0
- k3_node/loader/test_neighbor_loader.py +122 -0
- k3_node/loader/test_sampler_utils.py +82 -0
- k3_node/loader/test_samplers.py +96 -0
- k3_node/loader/test_subgraph_loaders.py +89 -0
- k3_node/loader/utils.py +232 -0
- k3_node/loader/zip_loader.py +88 -0
- k3_node/metrics.py +94 -0
- k3_node/models/__init__.py +424 -0
- k3_node/models/attentive_fp.py +232 -0
- k3_node/models/attract_repel.py +108 -0
- k3_node/models/autoencoder.py +318 -0
- k3_node/models/basic_gnn.py +443 -0
- k3_node/models/bio/__init__.py +4 -0
- k3_node/models/captum.py +52 -0
- k3_node/models/chemistry/__init__.py +4 -0
- k3_node/models/correct_and_smooth.py +146 -0
- k3_node/models/deep_graph_infomax.py +113 -0
- k3_node/models/deepgcn.py +121 -0
- k3_node/models/dimenet.py +737 -0
- k3_node/models/dimenet_utils.py +153 -0
- k3_node/models/gnnff.py +263 -0
- k3_node/models/gps_model.py +1122 -0
- k3_node/models/gpse.py +638 -0
- k3_node/models/graph_unet.py +199 -0
- k3_node/models/graphmae2.py +954 -0
- k3_node/models/graphormer.py +1258 -0
- k3_node/models/graphormer_3d.py +868 -0
- k3_node/models/grover.py +1066 -0
- k3_node/models/jumping_knowledge.py +200 -0
- k3_node/models/label_prop.py +110 -0
- k3_node/models/lightgcn.py +171 -0
- k3_node/models/linkx.py +181 -0
- k3_node/models/lpformer.py +404 -0
- k3_node/models/mask_label.py +114 -0
- k3_node/models/materials/__init__.py +33 -0
- k3_node/models/meta.py +133 -0
- k3_node/models/metapath2vec.py +234 -0
- k3_node/models/mlp.py +264 -0
- k3_node/models/mole_bert.py +379 -0
- k3_node/models/neural_fingerprint.py +95 -0
- k3_node/models/node2vec.py +213 -0
- k3_node/models/pmlp.py +157 -0
- k3_node/models/polynormer.py +229 -0
- k3_node/models/rect.py +93 -0
- k3_node/models/renet.py +221 -0
- k3_node/models/rev_gnn.py +128 -0
- k3_node/models/schnet.py +484 -0
- k3_node/models/sgformer.py +195 -0
- k3_node/models/signed_gcn.py +185 -0
- k3_node/models/test_attentive_fp.py +32 -0
- k3_node/models/test_attract_repel.py +33 -0
- k3_node/models/test_autoencoder.py +119 -0
- k3_node/models/test_basic_gnn.py +102 -0
- k3_node/models/test_correct_and_smooth.py +40 -0
- k3_node/models/test_deep_graph_infomax.py +68 -0
- k3_node/models/test_deepgcn.py +21 -0
- k3_node/models/test_dimenet.py +86 -0
- k3_node/models/test_domain_apis.py +138 -0
- k3_node/models/test_gnnff.py +24 -0
- k3_node/models/test_gps_model.py +271 -0
- k3_node/models/test_gpse.py +34 -0
- k3_node/models/test_graph_unet.py +26 -0
- k3_node/models/test_graphmae2.py +226 -0
- k3_node/models/test_graphormer.py +233 -0
- k3_node/models/test_graphormer3d.py +163 -0
- k3_node/models/test_grover.py +287 -0
- k3_node/models/test_jumping_knowledge.py +129 -0
- k3_node/models/test_label_prop.py +37 -0
- k3_node/models/test_lightgcn.py +38 -0
- k3_node/models/test_linkx.py +31 -0
- k3_node/models/test_lpformer.py +22 -0
- k3_node/models/test_mask_label.py +90 -0
- k3_node/models/test_meta.py +159 -0
- k3_node/models/test_metapath2vec.py +45 -0
- k3_node/models/test_mlp.py +62 -0
- k3_node/models/test_mole_bert.py +164 -0
- k3_node/models/test_neural_fingerprint.py +13 -0
- k3_node/models/test_node2vec.py +57 -0
- k3_node/models/test_pmlp.py +81 -0
- k3_node/models/test_polynormer.py +104 -0
- k3_node/models/test_rect.py +23 -0
- k3_node/models/test_renet.py +32 -0
- k3_node/models/test_rev_gnn.py +24 -0
- k3_node/models/test_schnet.py +43 -0
- k3_node/models/test_sgformer.py +48 -0
- k3_node/models/test_signed_gcn.py +28 -0
- k3_node/models/test_tgn.py +77 -0
- k3_node/models/test_unimol.py +179 -0
- k3_node/models/test_unimol2.py +114 -0
- k3_node/models/test_unimol_plus.py +131 -0
- k3_node/models/test_visnet.py +44 -0
- k3_node/models/tgn.py +382 -0
- k3_node/models/unimol.py +1156 -0
- k3_node/models/unimol2.py +616 -0
- k3_node/models/unimol_docking_v2.py +301 -0
- k3_node/models/unimol_plus.py +456 -0
- k3_node/models/utils.py +97 -0
- k3_node/models/visnet.py +759 -0
- k3_node/ops/__init__.py +4 -0
- k3_node/ops/conv.py +56 -0
- k3_node/ops/creation.py +43 -0
- k3_node/ops/graph.py +27 -0
- k3_node/ops/host.py +41 -0
- k3_node/ops/matmul.py +49 -0
- k3_node/ops/numpy.py +24 -0
- k3_node/ops/segment.py +54 -0
- k3_node/ops/sparse.py +51 -0
- k3_node/rag/__init__.py +49 -0
- k3_node/rag/encoders.py +312 -0
- k3_node/rag/pipeline.py +192 -0
- k3_node/rag/projector.py +184 -0
- k3_node/rag/subgraph.py +270 -0
- k3_node/rag/test_rag.py +347 -0
- k3_node/rag/verbalizer.py +162 -0
- k3_node/tasks/__init__.py +19 -0
- k3_node/tasks/backbone_resolver.py +125 -0
- k3_node/tasks/base.py +67 -0
- k3_node/tasks/graph_classification.py +270 -0
- k3_node/tasks/graph_regression.py +228 -0
- k3_node/tasks/link_prediction.py +306 -0
- k3_node/tasks/node_classification.py +194 -0
- k3_node/tasks/node_regression.py +138 -0
- k3_node/tasks/test_tasks.py +319 -0
- k3_node/test_docstring_examples.py +106 -0
- k3_node/test_training_forwarding.py +116 -0
- k3_node/training.py +115 -0
- k3_node/transforms/__init__.py +166 -0
- k3_node/transforms/base_transform.py +32 -0
- k3_node/transforms/compose.py +58 -0
- k3_node/transforms/general.py +676 -0
- k3_node/transforms/graph.py +1070 -0
- k3_node/transforms/spatial.py +797 -0
- k3_node/transforms/test_random_link_split.py +45 -0
- k3_node/transforms/test_spatial_transforms.py +65 -0
- k3_node/transforms/test_transforms.py +253 -0
- k3_node/transforms/utils.py +102 -0
- k3_node/utils/__init__.py +5 -0
- k3_node/utils/backend_import.py +12 -0
- k3_node/utils/graph.py +286 -0
- k3_node/utils/keras.py +94 -0
- k3_node/utils/random.py +103 -0
- k3_node/utils/smiles.py +235 -0
- k3_node-1.0.0.dist-info/METADATA +284 -0
- k3_node-1.0.0.dist-info/RECORD +459 -0
- k3_node-1.0.0.dist-info/WHEEL +5 -0
- k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
- k3_node-1.0.0.dist-info/top_level.txt +1 -0
k3_node/data/temporal.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Any, Dict, Iterable, List, NamedTuple, Optional, Sequence, Tuple, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.data.data import BaseData, size_repr
|
|
8
|
+
from k3_node.data.storage import BaseStorage, GlobalStorage, get_shape, is_tensor_like
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class TemporalData(BaseData):
|
|
12
|
+
"""A data object composed of a stream of events describing a temporal graph."""
|
|
13
|
+
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
src: Optional[Any] = None,
|
|
17
|
+
dst: Optional[Any] = None,
|
|
18
|
+
t: Optional[Any] = None,
|
|
19
|
+
msg: Optional[Any] = None,
|
|
20
|
+
y: Optional[Any] = None,
|
|
21
|
+
**kwargs,
|
|
22
|
+
):
|
|
23
|
+
self.__dict__["_store"] = GlobalStorage(_parent=self)
|
|
24
|
+
if src is not None:
|
|
25
|
+
self.src = src
|
|
26
|
+
if dst is not None:
|
|
27
|
+
self.dst = dst
|
|
28
|
+
if t is not None:
|
|
29
|
+
self.t = t
|
|
30
|
+
if msg is not None:
|
|
31
|
+
self.msg = msg
|
|
32
|
+
if y is not None:
|
|
33
|
+
self.y = y
|
|
34
|
+
for key, value in kwargs.items():
|
|
35
|
+
setattr(self, key, value)
|
|
36
|
+
|
|
37
|
+
@classmethod
|
|
38
|
+
def from_dict(cls, mapping: Dict[str, Any]) -> "TemporalData":
|
|
39
|
+
return cls(**mapping)
|
|
40
|
+
|
|
41
|
+
@property
|
|
42
|
+
def num_events(self) -> int:
|
|
43
|
+
for key in ("src", "dst", "t", "msg"):
|
|
44
|
+
if key in self._store and self._store[key] is not None:
|
|
45
|
+
return get_shape(self._store[key])[0]
|
|
46
|
+
return 0
|
|
47
|
+
|
|
48
|
+
def train_val_test_split(self, val_ratio: float = 0.15, test_ratio: float = 0.15):
|
|
49
|
+
r"""Splits the events chronologically into training, validation and test events (the
|
|
50
|
+
last ``val_ratio + test_ratio`` of the time span go to validation and test)."""
|
|
51
|
+
t = np.asarray(ops.convert_to_numpy(self.t))
|
|
52
|
+
val_time, test_time = np.quantile(t, [1.0 - val_ratio - test_ratio, 1.0 - test_ratio])
|
|
53
|
+
val_idx, test_idx = int((t <= val_time).sum()), int((t <= test_time).sum())
|
|
54
|
+
return self[:val_idx], self[val_idx:test_idx], self[test_idx:]
|
|
55
|
+
|
|
56
|
+
@property
|
|
57
|
+
def num_nodes(self) -> int:
|
|
58
|
+
nodes = []
|
|
59
|
+
if "src" in self._store and self._store.src is not None:
|
|
60
|
+
src_np = ops.convert_to_numpy(self._store.src)
|
|
61
|
+
if src_np.size > 0:
|
|
62
|
+
nodes.append(int(np.max(src_np)))
|
|
63
|
+
if "dst" in self._store and self._store.dst is not None:
|
|
64
|
+
dst_np = ops.convert_to_numpy(self._store.dst)
|
|
65
|
+
if dst_np.size > 0:
|
|
66
|
+
nodes.append(int(np.max(dst_np)))
|
|
67
|
+
return max(nodes) + 1 if len(nodes) > 0 else 0
|
|
68
|
+
|
|
69
|
+
def __len__(self) -> int:
|
|
70
|
+
return self.num_events
|
|
71
|
+
|
|
72
|
+
def __getitem__(self, idx: Any) -> Any:
|
|
73
|
+
if isinstance(idx, str):
|
|
74
|
+
return self._store[idx]
|
|
75
|
+
data = copy.copy(self)
|
|
76
|
+
num_events = self.num_events
|
|
77
|
+
for key, value in data._store.items():
|
|
78
|
+
if is_tensor_like(value) and get_shape(value)[0] == num_events:
|
|
79
|
+
data[key] = value[idx]
|
|
80
|
+
return data
|
|
81
|
+
|
|
82
|
+
def __setitem__(self, key: str, value: Any):
|
|
83
|
+
self._store[key] = value
|
|
84
|
+
|
|
85
|
+
def __delitem__(self, key: str):
|
|
86
|
+
if key in self._store:
|
|
87
|
+
del self._store[key]
|
|
88
|
+
|
|
89
|
+
def __getattr__(self, key: str) -> Any:
|
|
90
|
+
if "_store" not in self.__dict__:
|
|
91
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
|
|
92
|
+
try:
|
|
93
|
+
return getattr(self._store, key)
|
|
94
|
+
except AttributeError:
|
|
95
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'") from None
|
|
96
|
+
|
|
97
|
+
def __setattr__(self, key: str, value: Any):
|
|
98
|
+
if key == "_store":
|
|
99
|
+
self.__dict__["_store"] = value
|
|
100
|
+
elif "_store" in self.__dict__:
|
|
101
|
+
setattr(self._store, key, value)
|
|
102
|
+
else:
|
|
103
|
+
self.__dict__[key] = value
|
|
104
|
+
|
|
105
|
+
def __delattr__(self, key: str):
|
|
106
|
+
if key == "_store":
|
|
107
|
+
del self.__dict__["_store"]
|
|
108
|
+
elif "_store" in self.__dict__:
|
|
109
|
+
delattr(self._store, key)
|
|
110
|
+
else:
|
|
111
|
+
del self.__dict__[key]
|
|
112
|
+
|
|
113
|
+
def __copy__(self):
|
|
114
|
+
out = self.__class__.__new__(self.__class__)
|
|
115
|
+
for key, value in self.__dict__.items():
|
|
116
|
+
out.__dict__[key] = value
|
|
117
|
+
out.__dict__["_store"] = copy.copy(self._store)
|
|
118
|
+
out._store._parent = out
|
|
119
|
+
return out
|
|
120
|
+
|
|
121
|
+
def __deepcopy__(self, memo=None):
|
|
122
|
+
out = self.__class__.__new__(self.__class__)
|
|
123
|
+
for key, value in self.__dict__.items():
|
|
124
|
+
out.__dict__[key] = copy.deepcopy(value, memo)
|
|
125
|
+
out._store._parent = out
|
|
126
|
+
return out
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def stores(self) -> List[BaseStorage]:
|
|
130
|
+
return [self._store]
|
|
131
|
+
|
|
132
|
+
@property
|
|
133
|
+
def node_stores(self) -> List[Any]:
|
|
134
|
+
return [self._store]
|
|
135
|
+
|
|
136
|
+
@property
|
|
137
|
+
def edge_stores(self) -> List[Any]:
|
|
138
|
+
return [self._store]
|
|
139
|
+
|
|
140
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
141
|
+
return self._store.to_dict()
|
|
142
|
+
|
|
143
|
+
def to_namedtuple(self) -> NamedTuple:
|
|
144
|
+
fields = sorted(list(self.keys()))
|
|
145
|
+
import collections
|
|
146
|
+
|
|
147
|
+
TemporalTuple = collections.namedtuple("TemporalTuple", fields)
|
|
148
|
+
return TemporalTuple(**{f: self[f] for f in fields})
|
|
149
|
+
|
|
150
|
+
def __repr__(self) -> str:
|
|
151
|
+
cls = self.__class__.__name__
|
|
152
|
+
attrs = [size_repr(k, v) for k, v in self._store.items()]
|
|
153
|
+
return f"{cls}({', '.join(attrs)})"
|
|
154
|
+
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Batch, Data, HeteroData
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def test_batch_homogeneous():
|
|
9
|
+
d1 = Data(
|
|
10
|
+
x=ops.convert_to_tensor([[1.0, 2.0], [3.0, 4.0]]),
|
|
11
|
+
edge_index=ops.convert_to_tensor([[0, 1], [1, 0]], dtype="int64"),
|
|
12
|
+
y=ops.convert_to_tensor([0]),
|
|
13
|
+
)
|
|
14
|
+
d2 = Data(
|
|
15
|
+
x=ops.convert_to_tensor([[5.0, 6.0], [7.0, 8.0], [9.0, 10.0]]),
|
|
16
|
+
edge_index=ops.convert_to_tensor([[0, 1, 2], [1, 2, 0]], dtype="int64"),
|
|
17
|
+
y=ops.convert_to_tensor([1]),
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
batch = Batch.from_data_list([d1, d2])
|
|
21
|
+
assert batch.num_graphs == 2
|
|
22
|
+
assert batch.num_nodes == 5
|
|
23
|
+
assert batch.num_edges == 5
|
|
24
|
+
|
|
25
|
+
# Check batch vector
|
|
26
|
+
batch_vec = ops.convert_to_numpy(batch.batch)
|
|
27
|
+
assert np.array_equal(batch_vec, [0, 0, 1, 1, 1])
|
|
28
|
+
|
|
29
|
+
# Check ptr vector
|
|
30
|
+
ptr_vec = ops.convert_to_numpy(batch.ptr)
|
|
31
|
+
assert np.array_equal(ptr_vec, [0, 2, 5])
|
|
32
|
+
|
|
33
|
+
# Check offset edge_index
|
|
34
|
+
ei = ops.convert_to_numpy(batch.edge_index)
|
|
35
|
+
assert np.array_equal(ei[:, :2], [[0, 1], [1, 0]])
|
|
36
|
+
assert np.array_equal(ei[:, 2:], [[2, 3, 4], [3, 4, 2]])
|
|
37
|
+
|
|
38
|
+
# Separate back
|
|
39
|
+
rec1 = batch[0]
|
|
40
|
+
rec2 = batch[1]
|
|
41
|
+
assert rec1.num_nodes == 2
|
|
42
|
+
assert rec2.num_nodes == 3
|
|
43
|
+
assert np.allclose(ops.convert_to_numpy(rec1.x), ops.convert_to_numpy(d1.x))
|
|
44
|
+
assert np.allclose(ops.convert_to_numpy(rec2.x), ops.convert_to_numpy(d2.x))
|
|
45
|
+
|
|
46
|
+
data_list = batch.to_data_list()
|
|
47
|
+
assert len(data_list) == 2
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def test_batch_heterogeneous():
|
|
51
|
+
h1 = HeteroData()
|
|
52
|
+
h1["v"].x = ops.convert_to_tensor([[1.0], [2.0]])
|
|
53
|
+
h1["v", "e", "v"].edge_index = ops.convert_to_tensor([[0], [1]], dtype="int64")
|
|
54
|
+
|
|
55
|
+
h2 = HeteroData()
|
|
56
|
+
h2["v"].x = ops.convert_to_tensor([[3.0], [4.0], [5.0]])
|
|
57
|
+
h2["v", "e", "v"].edge_index = ops.convert_to_tensor([[0, 1], [1, 2]], dtype="int64")
|
|
58
|
+
|
|
59
|
+
batch = Batch.from_data_list([h1, h2])
|
|
60
|
+
assert batch.num_graphs == 2
|
|
61
|
+
assert batch["v"].num_nodes == 5
|
|
62
|
+
assert batch["v", "e", "v"].num_edges == 3
|
|
63
|
+
|
|
64
|
+
rec1 = batch[0]
|
|
65
|
+
assert rec1["v"].num_nodes == 2
|
|
66
|
+
assert rec1["v", "e", "v"].num_edges == 1
|
|
67
|
+
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def test_data_basic():
|
|
9
|
+
x = ops.convert_to_tensor([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
|
|
10
|
+
edge_index = ops.convert_to_tensor([[0, 1, 1, 2], [1, 0, 2, 1]], dtype="int64")
|
|
11
|
+
y = ops.convert_to_tensor([0, 1, 0], dtype="int64")
|
|
12
|
+
|
|
13
|
+
data = Data(x=x, edge_index=edge_index, y=y)
|
|
14
|
+
assert data.num_nodes == 3
|
|
15
|
+
assert data.num_edges == 4
|
|
16
|
+
assert data.num_features == 2
|
|
17
|
+
assert data.num_classes == 2
|
|
18
|
+
assert data.is_undirected()
|
|
19
|
+
assert not data.is_directed()
|
|
20
|
+
assert not data.has_self_loops()
|
|
21
|
+
assert not data.has_isolated_nodes()
|
|
22
|
+
|
|
23
|
+
# Attribute and dict access
|
|
24
|
+
assert "x" in data
|
|
25
|
+
assert "edge_index" in data
|
|
26
|
+
assert "edge_attr" not in data
|
|
27
|
+
assert len(data.keys()) == 3
|
|
28
|
+
assert ops.convert_to_numpy(data["x"]).shape == (3, 2)
|
|
29
|
+
assert ops.convert_to_numpy(data.x).shape == (3, 2)
|
|
30
|
+
|
|
31
|
+
# Clone
|
|
32
|
+
clone = data.clone()
|
|
33
|
+
assert clone.num_nodes == 3
|
|
34
|
+
assert clone.num_edges == 4
|
|
35
|
+
|
|
36
|
+
# Dict and namedtuple
|
|
37
|
+
d = data.to_dict()
|
|
38
|
+
assert "x" in d and "edge_index" in d and "y" in d
|
|
39
|
+
nt = data.to_namedtuple()
|
|
40
|
+
assert hasattr(nt, "x") and hasattr(nt, "edge_index")
|
|
41
|
+
|
|
42
|
+
# String repr
|
|
43
|
+
repr_str = str(data)
|
|
44
|
+
assert "Data(" in repr_str and "x=" in repr_str and "edge_index=" in repr_str
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def test_data_subgraph():
|
|
48
|
+
x = ops.convert_to_tensor([[10.0], [20.0], [30.0], [40.0]])
|
|
49
|
+
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
|
|
50
|
+
data = Data(x=x, edge_index=edge_index)
|
|
51
|
+
|
|
52
|
+
subset = ops.convert_to_tensor([0, 2], dtype="int64")
|
|
53
|
+
sub = data.subgraph(subset)
|
|
54
|
+
assert sub.num_nodes == 2
|
|
55
|
+
assert ops.convert_to_numpy(sub.x).tolist() == [[10.0], [30.0]]
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def test_data_to_heterogeneous():
|
|
59
|
+
x = ops.convert_to_tensor([[1.0], [2.0]])
|
|
60
|
+
edge_index = ops.convert_to_tensor([[0, 1], [1, 0]], dtype="int64")
|
|
61
|
+
data = Data(x=x, edge_index=edge_index)
|
|
62
|
+
|
|
63
|
+
hetero = data.to_heterogeneous(node_type="v", edge_type=("v", "e", "v"))
|
|
64
|
+
assert "v" in hetero.node_types
|
|
65
|
+
assert ("v", "e", "v") in hetero.edge_types
|
|
66
|
+
assert ops.convert_to_numpy(hetero["v"].x).shape == (2, 1)
|
|
67
|
+
assert ops.convert_to_numpy(hetero["v", "e", "v"].edge_index).shape == (2, 2)
|
|
68
|
+
|
|
@@ -0,0 +1,111 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import tempfile
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pytest
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.data import (
|
|
8
|
+
Data,
|
|
9
|
+
EdgeAttr,
|
|
10
|
+
EdgeLayout,
|
|
11
|
+
FeatureStore,
|
|
12
|
+
GraphStore,
|
|
13
|
+
InMemoryDataset,
|
|
14
|
+
OnDiskDataset,
|
|
15
|
+
SQLiteDatabase,
|
|
16
|
+
TensorAttr,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class MyInMemoryDataset(InMemoryDataset):
|
|
21
|
+
def __init__(self, data_list, root=None):
|
|
22
|
+
super().__init__(root)
|
|
23
|
+
self._data_list = data_list
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def test_in_memory_dataset():
|
|
27
|
+
d1 = Data(x=ops.convert_to_tensor([[1.0]]), edge_index=ops.convert_to_tensor([[0], [0]], dtype="int64"))
|
|
28
|
+
d2 = Data(x=ops.convert_to_tensor([[2.0]]), edge_index=ops.convert_to_tensor([[0], [0]], dtype="int64"))
|
|
29
|
+
|
|
30
|
+
ds = MyInMemoryDataset([d1, d2])
|
|
31
|
+
assert len(ds) == 2
|
|
32
|
+
assert ops.convert_to_numpy(ds[0].x).item() == 1.0
|
|
33
|
+
assert ops.convert_to_numpy(ds[1].x).item() == 2.0
|
|
34
|
+
|
|
35
|
+
# Slicing
|
|
36
|
+
sub_ds = ds[1:]
|
|
37
|
+
assert len(sub_ds) == 1
|
|
38
|
+
assert ops.convert_to_numpy(sub_ds[0].x).item() == 2.0
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def test_sqlite_database_and_on_disk_dataset():
|
|
42
|
+
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
43
|
+
db_path = os.path.join(tmp_dir, "test.db")
|
|
44
|
+
db = SQLiteDatabase(path=db_path)
|
|
45
|
+
d1 = Data(x=ops.convert_to_tensor([[1.0]]))
|
|
46
|
+
d2 = Data(x=ops.convert_to_tensor([[2.0]]))
|
|
47
|
+
db[0] = d1
|
|
48
|
+
db[1] = d2
|
|
49
|
+
assert len(db) == 2
|
|
50
|
+
rec1 = db[0]
|
|
51
|
+
assert ops.convert_to_numpy(rec1.x).item() == 1.0
|
|
52
|
+
db.close()
|
|
53
|
+
|
|
54
|
+
# Test OnDiskDataset
|
|
55
|
+
ds = OnDiskDataset(root=tmp_dir, backend="sqlite")
|
|
56
|
+
ds.append(d1)
|
|
57
|
+
ds.append(d2)
|
|
58
|
+
assert len(ds) == 2
|
|
59
|
+
rec = ds[0]
|
|
60
|
+
assert ops.convert_to_numpy(rec.x).item() == 1.0
|
|
61
|
+
ds.close()
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class SimpleFeatureStore(FeatureStore):
|
|
65
|
+
def _put_tensor(self, tensor, attr):
|
|
66
|
+
self._feat_dict[(attr.group_name, attr.attr_name)] = tensor
|
|
67
|
+
return True
|
|
68
|
+
|
|
69
|
+
def _get_tensor(self, attr):
|
|
70
|
+
return self._feat_dict.get((attr.group_name, attr.attr_name))
|
|
71
|
+
|
|
72
|
+
def _remove_tensor(self, attr):
|
|
73
|
+
return self._feat_dict.pop((attr.group_name, attr.attr_name), None) is not None
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class SimpleGraphStore(GraphStore):
|
|
77
|
+
def __init__(self):
|
|
78
|
+
self._store = {}
|
|
79
|
+
|
|
80
|
+
def _put_edge_index(self, edge_index, edge_attr):
|
|
81
|
+
self._store[(edge_attr.edge_type, edge_attr.layout)] = edge_index
|
|
82
|
+
return True
|
|
83
|
+
|
|
84
|
+
def _get_edge_index(self, edge_attr):
|
|
85
|
+
return self._store.get((edge_attr.edge_type, edge_attr.layout))
|
|
86
|
+
|
|
87
|
+
def _remove_edge_index(self, edge_attr):
|
|
88
|
+
return self._store.pop((edge_attr.edge_type, edge_attr.layout), None) is not None
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def test_feature_and_graph_store():
|
|
92
|
+
fs = SimpleFeatureStore()
|
|
93
|
+
x = ops.convert_to_tensor([[1.0, 2.0]])
|
|
94
|
+
fs.put_tensor(x, group_name="user", attr_name="feat")
|
|
95
|
+
ret = fs.get_tensor(group_name="user", attr_name="feat")
|
|
96
|
+
assert ops.convert_to_numpy(ret).shape == (1, 2)
|
|
97
|
+
|
|
98
|
+
gs = SimpleGraphStore()
|
|
99
|
+
ei = ops.convert_to_tensor([[0, 1], [1, 0]], dtype="int64")
|
|
100
|
+
gs.put_edge_index(ei, edge_type=("u", "follows", "u"), layout=EdgeLayout.COO)
|
|
101
|
+
ret_ei = gs.get_edge_index(edge_type=("u", "follows", "u"), layout=EdgeLayout.COO)
|
|
102
|
+
assert ops.convert_to_numpy(ret_ei).shape == (2, 2)
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def test_fractional_slicing():
|
|
107
|
+
from k3_node.datasets import FakeDataset
|
|
108
|
+
|
|
109
|
+
dataset = FakeDataset(num_graphs=10)
|
|
110
|
+
assert len(dataset[:0.9]) == 9 and len(dataset[0.9:]) == 1
|
|
111
|
+
assert len(dataset[0.2:0.5]) == 3
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import HeteroData
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def test_hetero_data_basic():
|
|
9
|
+
data = HeteroData()
|
|
10
|
+
|
|
11
|
+
data["paper"].x = ops.convert_to_tensor([[1.0, 2.0], [3.0, 4.0]])
|
|
12
|
+
data["author"].x = ops.convert_to_tensor([[5.0, 6.0], [7.0, 8.0], [9.0, 10.0]])
|
|
13
|
+
data["author", "writes", "paper"].edge_index = ops.convert_to_tensor(
|
|
14
|
+
[[0, 1, 2], [0, 1, 1]], dtype="int64"
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
assert set(data.node_types) == {"paper", "author"}
|
|
18
|
+
assert set(data.edge_types) == {("author", "writes", "paper")}
|
|
19
|
+
assert data["paper"].num_nodes == 2
|
|
20
|
+
assert data["author"].num_nodes == 3
|
|
21
|
+
assert data["author", "writes", "paper"].num_edges == 3
|
|
22
|
+
|
|
23
|
+
meta = data.metadata()
|
|
24
|
+
assert len(meta[0]) == 2
|
|
25
|
+
assert len(meta[1]) == 1
|
|
26
|
+
|
|
27
|
+
# Homogeneous conversion
|
|
28
|
+
homo = data.to_homogeneous()
|
|
29
|
+
assert homo.num_nodes == 5
|
|
30
|
+
assert homo.num_edges == 3
|
|
31
|
+
assert ops.convert_to_numpy(homo.x).shape == (5, 2)
|
|
32
|
+
assert ops.convert_to_numpy(homo.edge_index).shape == (2, 3)
|
|
33
|
+
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import HypergraphData, TemporalData
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def test_temporal_data():
|
|
9
|
+
src = ops.convert_to_tensor([0, 1, 2, 3], dtype="int64")
|
|
10
|
+
dst = ops.convert_to_tensor([1, 2, 3, 4], dtype="int64")
|
|
11
|
+
t = ops.convert_to_tensor([10, 20, 30, 40], dtype="int64")
|
|
12
|
+
msg = ops.convert_to_tensor([[1.0], [2.0], [3.0], [4.0]])
|
|
13
|
+
|
|
14
|
+
data = TemporalData(src=src, dst=dst, t=t, msg=msg)
|
|
15
|
+
assert data.num_events == 4
|
|
16
|
+
assert data.num_nodes == 5
|
|
17
|
+
assert len(data) == 4
|
|
18
|
+
|
|
19
|
+
sub = data[:2]
|
|
20
|
+
assert sub.num_events == 2
|
|
21
|
+
assert ops.convert_to_numpy(sub.src).tolist() == [0, 1]
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def test_hypergraph_data():
|
|
25
|
+
x = ops.convert_to_tensor([[1.0], [2.0], [3.0], [4.0], [5.0]])
|
|
26
|
+
edge_index = ops.convert_to_tensor(
|
|
27
|
+
[[0, 1, 2, 1, 2, 3, 4], [0, 0, 0, 1, 1, 1, 1]], dtype="int64"
|
|
28
|
+
)
|
|
29
|
+
data = HypergraphData(x=x, edge_index=edge_index)
|
|
30
|
+
assert data.num_nodes == 5
|
|
31
|
+
assert data.num_edges == 2
|
|
32
|
+
|
k3_node/data/view.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
from typing import Any, Iterator, List, Mapping, Tuple
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class MappingView:
|
|
5
|
+
def __init__(self, mapping: Mapping[str, Any], *args: str):
|
|
6
|
+
self._mapping = mapping
|
|
7
|
+
self._args = args
|
|
8
|
+
|
|
9
|
+
def _keys(self) -> List[str]:
|
|
10
|
+
if len(self._args) == 0:
|
|
11
|
+
return list(self._mapping.keys())
|
|
12
|
+
else:
|
|
13
|
+
return [arg for arg in self._args if arg in self._mapping]
|
|
14
|
+
|
|
15
|
+
def __len__(self) -> int:
|
|
16
|
+
return len(self._keys())
|
|
17
|
+
|
|
18
|
+
def __contains__(self, item: Any) -> bool:
|
|
19
|
+
return item in self._keys()
|
|
20
|
+
|
|
21
|
+
def __repr__(self) -> str:
|
|
22
|
+
mapping = {key: self._mapping[key] for key in self._keys()}
|
|
23
|
+
return f"{self.__class__.__name__}({mapping})"
|
|
24
|
+
|
|
25
|
+
__class_getitem__ = classmethod(type([]))
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class KeysView(MappingView):
|
|
29
|
+
def __iter__(self) -> Iterator[str]:
|
|
30
|
+
yield from self._keys()
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ValuesView(MappingView):
|
|
34
|
+
def __iter__(self) -> Iterator[Any]:
|
|
35
|
+
for key in self._keys():
|
|
36
|
+
yield self._mapping[key]
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class ItemsView(MappingView):
|
|
40
|
+
def __iter__(self) -> Iterator[Tuple[str, Any]]:
|
|
41
|
+
for key in self._keys():
|
|
42
|
+
yield (key, self._mapping[key])
|
|
43
|
+
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
from .karate import KarateClub
|
|
2
|
+
from .fake import FakeDataset, FakeHeteroDataset
|
|
3
|
+
from .planetoid import Planetoid
|
|
4
|
+
from .tu_dataset import TUDataset
|
|
5
|
+
from .citation_full import CitationFull, CoraFull
|
|
6
|
+
from .amazon import Amazon
|
|
7
|
+
from .coauthor import Coauthor
|
|
8
|
+
from .wikics import WikiCS
|
|
9
|
+
from .webkb import WebKB
|
|
10
|
+
from .actor import Actor
|
|
11
|
+
from .polblogs import PolBlogs
|
|
12
|
+
from .airports import Airports
|
|
13
|
+
from .email_eu_core import EmailEUCore
|
|
14
|
+
from .github import GitHub
|
|
15
|
+
from .facebook import FacebookPagePage
|
|
16
|
+
from .lastfm_asia import LastFMAsia
|
|
17
|
+
from .twitch import Twitch
|
|
18
|
+
from .ba_shapes import BAShapes
|
|
19
|
+
from .ba2motif_dataset import BA2MotifDataset
|
|
20
|
+
from .sbm_dataset import StochasticBlockModelDataset, RandomPartitionGraphDataset
|
|
21
|
+
from .explainer_dataset import ExplainerDataset
|
|
22
|
+
from .entities import Entities
|
|
23
|
+
from .word_net import WordNet18, WordNet18RR
|
|
24
|
+
from .freebase import FB15k_237
|
|
25
|
+
from .dblp import DBLP
|
|
26
|
+
from .imdb import IMDB
|
|
27
|
+
from .qm7 import QM7b
|
|
28
|
+
from .molecule_net import MoleculeNet
|
|
29
|
+
from .ppi import PPI
|
|
30
|
+
from .reddit import Reddit
|
|
31
|
+
from .digits import Digits
|
|
32
|
+
from .seal import SEALDataset
|
|
33
|
+
from .bitcoin_otc import BitcoinOTC
|
|
34
|
+
from .geometric_shapes import GeometricShapes
|
|
35
|
+
from .shape_scenes import ShapeScenes
|
|
36
|
+
from .mesh_correspondence import MeshCorrespondence
|
|
37
|
+
from .movielens import MovieLens100K
|
|
38
|
+
from .qm9 import QM9
|
|
39
|
+
from .jodie import JODIEDataset
|
|
40
|
+
from .icews import ICEWS18
|
|
41
|
+
|
|
42
|
+
__all__ = [
|
|
43
|
+
"Digits",
|
|
44
|
+
"SEALDataset",
|
|
45
|
+
"BitcoinOTC",
|
|
46
|
+
"GeometricShapes",
|
|
47
|
+
"ShapeScenes",
|
|
48
|
+
"MeshCorrespondence",
|
|
49
|
+
"MovieLens100K",
|
|
50
|
+
"QM9",
|
|
51
|
+
"JODIEDataset",
|
|
52
|
+
"ICEWS18",
|
|
53
|
+
"KarateClub",
|
|
54
|
+
"FakeDataset",
|
|
55
|
+
"FakeHeteroDataset",
|
|
56
|
+
"Planetoid",
|
|
57
|
+
"TUDataset",
|
|
58
|
+
"CitationFull",
|
|
59
|
+
"CoraFull",
|
|
60
|
+
"Amazon",
|
|
61
|
+
"Coauthor",
|
|
62
|
+
"WikiCS",
|
|
63
|
+
"WebKB",
|
|
64
|
+
"Actor",
|
|
65
|
+
"PolBlogs",
|
|
66
|
+
"Airports",
|
|
67
|
+
"EmailEUCore",
|
|
68
|
+
"GitHub",
|
|
69
|
+
"FacebookPagePage",
|
|
70
|
+
"LastFMAsia",
|
|
71
|
+
"Twitch",
|
|
72
|
+
"BAShapes",
|
|
73
|
+
"BA2MotifDataset",
|
|
74
|
+
"StochasticBlockModelDataset",
|
|
75
|
+
"RandomPartitionGraphDataset",
|
|
76
|
+
"ExplainerDataset",
|
|
77
|
+
"Entities",
|
|
78
|
+
"WordNet18",
|
|
79
|
+
"WordNet18RR",
|
|
80
|
+
"FB15k_237",
|
|
81
|
+
"DBLP",
|
|
82
|
+
"IMDB",
|
|
83
|
+
"QM7b",
|
|
84
|
+
"MoleculeNet",
|
|
85
|
+
"PPI",
|
|
86
|
+
"Reddit",
|
|
87
|
+
]
|
|
88
|
+
|