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
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
from typing import List
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
import torch.utils.data
|
|
8
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
9
|
+
except ImportError:
|
|
10
|
+
torch = None
|
|
11
|
+
BaseDataLoader = object
|
|
12
|
+
|
|
13
|
+
from k3_node.data import TemporalData
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class TemporalDataLoader(BaseDataLoader):
|
|
17
|
+
r"""A data loader which merges successive events of a
|
|
18
|
+
:class:`k3_node.data.TemporalData` to a mini-batch.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
data (TemporalData): The :obj:`~k3_node.data.TemporalData` from which to load.
|
|
22
|
+
batch_size (int, optional): How many samples per batch to load. (default: :obj:`1`)
|
|
23
|
+
neg_sampling_ratio (float, optional): The ratio of sampled negative
|
|
24
|
+
destination nodes to the number of positive destination nodes. (default: :obj:`0.0`)
|
|
25
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
data: TemporalData,
|
|
30
|
+
batch_size: int = 1,
|
|
31
|
+
neg_sampling_ratio: float = 0.0,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
kwargs.pop('dataset', None)
|
|
35
|
+
kwargs.pop('collate_fn', None)
|
|
36
|
+
kwargs.pop('shuffle', None)
|
|
37
|
+
|
|
38
|
+
self.data = data
|
|
39
|
+
self.events_per_batch = batch_size
|
|
40
|
+
self.neg_sampling_ratio = neg_sampling_ratio
|
|
41
|
+
|
|
42
|
+
if neg_sampling_ratio > 0:
|
|
43
|
+
dst = data.dst
|
|
44
|
+
if torch is not None and isinstance(dst, torch.Tensor):
|
|
45
|
+
self.min_dst = int(dst.min())
|
|
46
|
+
self.max_dst = int(dst.max())
|
|
47
|
+
else:
|
|
48
|
+
self.min_dst = int(np.min(np.asarray(dst)))
|
|
49
|
+
self.max_dst = int(np.max(np.asarray(dst)))
|
|
50
|
+
|
|
51
|
+
if kwargs.get('drop_last', False) and len(data) % batch_size != 0:
|
|
52
|
+
arange = list(range(0, len(data) - batch_size, batch_size))
|
|
53
|
+
else:
|
|
54
|
+
arange = list(range(0, len(data), batch_size))
|
|
55
|
+
|
|
56
|
+
if torch is not None:
|
|
57
|
+
super().__init__(arange, 1, shuffle=False, collate_fn=self, **kwargs)
|
|
58
|
+
else:
|
|
59
|
+
self.dataset = arange
|
|
60
|
+
self.batch_size = 1
|
|
61
|
+
self.shuffle = False
|
|
62
|
+
self.collate_fn = self
|
|
63
|
+
|
|
64
|
+
def __call__(self, arange: List[int]) -> TemporalData:
|
|
65
|
+
start = arange[0]
|
|
66
|
+
end = start + self.events_per_batch
|
|
67
|
+
batch = self.data[start:end]
|
|
68
|
+
|
|
69
|
+
is_torch = torch is not None and isinstance(batch.dst, torch.Tensor)
|
|
70
|
+
|
|
71
|
+
n_ids = [batch.src, batch.dst]
|
|
72
|
+
|
|
73
|
+
if self.neg_sampling_ratio > 0:
|
|
74
|
+
num_neg = round(self.neg_sampling_ratio * (batch.dst.size(0) if is_torch else len(batch.dst)))
|
|
75
|
+
if is_torch:
|
|
76
|
+
batch.neg_dst = torch.randint(
|
|
77
|
+
low=self.min_dst,
|
|
78
|
+
high=self.max_dst + 1,
|
|
79
|
+
size=(num_neg,),
|
|
80
|
+
dtype=batch.dst.dtype,
|
|
81
|
+
device=batch.dst.device,
|
|
82
|
+
)
|
|
83
|
+
else:
|
|
84
|
+
batch.neg_dst = np.random.randint(
|
|
85
|
+
low=self.min_dst,
|
|
86
|
+
high=self.max_dst + 1,
|
|
87
|
+
size=(num_neg,),
|
|
88
|
+
dtype=np.int64,
|
|
89
|
+
)
|
|
90
|
+
n_ids.append(batch.neg_dst)
|
|
91
|
+
|
|
92
|
+
if is_torch:
|
|
93
|
+
batch.n_id = torch.cat(n_ids, dim=0).unique()
|
|
94
|
+
else:
|
|
95
|
+
batch.n_id = np.unique(np.concatenate([np.asarray(x) for x in n_ids], axis=0))
|
|
96
|
+
|
|
97
|
+
return batch
|
|
98
|
+
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
|
|
4
|
+
try:
|
|
5
|
+
import torch
|
|
6
|
+
except ImportError:
|
|
7
|
+
torch = None
|
|
8
|
+
|
|
9
|
+
from k3_node.data import Data, HeteroData, TemporalData
|
|
10
|
+
from k3_node.loader import (
|
|
11
|
+
DataListLoader,
|
|
12
|
+
DataLoader,
|
|
13
|
+
DenseDataLoader,
|
|
14
|
+
RandomNodeLoader,
|
|
15
|
+
TemporalDataLoader,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def create_dummy_data(num_nodes=4, num_features=8):
|
|
20
|
+
x = np.random.randn(num_nodes, num_features).astype(np.float32)
|
|
21
|
+
edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int64)
|
|
22
|
+
y = np.array([0, 1, 0, 1], dtype=np.int64)
|
|
23
|
+
if torch is not None:
|
|
24
|
+
x = torch.from_numpy(x)
|
|
25
|
+
edge_index = torch.from_numpy(edge_index)
|
|
26
|
+
y = torch.from_numpy(y)
|
|
27
|
+
return Data(x=x, edge_index=edge_index, y=y)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def test_data_loader_basic():
|
|
31
|
+
dataset = [create_dummy_data(num_nodes=i + 2) for i in range(4)]
|
|
32
|
+
loader = DataLoader(dataset, batch_size=2, shuffle=False)
|
|
33
|
+
|
|
34
|
+
batches = list(loader)
|
|
35
|
+
assert len(batches) == 2
|
|
36
|
+
|
|
37
|
+
batch0 = batches[0]
|
|
38
|
+
assert batch0.num_graphs == 2
|
|
39
|
+
assert hasattr(batch0, 'batch')
|
|
40
|
+
assert hasattr(batch0, 'ptr')
|
|
41
|
+
assert batch0.num_nodes == dataset[0].num_nodes + dataset[1].num_nodes
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def test_data_loader_follow_batch_and_exclude():
|
|
45
|
+
dataset = [create_dummy_data(num_nodes=3) for _ in range(3)]
|
|
46
|
+
loader = DataLoader(dataset, batch_size=2, follow_batch=['y'], exclude_keys=['edge_index'])
|
|
47
|
+
|
|
48
|
+
batch = next(iter(loader))
|
|
49
|
+
assert 'edge_index' not in batch
|
|
50
|
+
assert hasattr(batch, 'y_batch')
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def test_data_list_loader():
|
|
54
|
+
dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
|
|
55
|
+
loader = DataListLoader(dataset, batch_size=2, shuffle=False)
|
|
56
|
+
|
|
57
|
+
batches = list(loader)
|
|
58
|
+
assert len(batches) == 2
|
|
59
|
+
assert isinstance(batches[0], list)
|
|
60
|
+
assert len(batches[0]) == 2
|
|
61
|
+
assert isinstance(batches[0][0], Data)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def test_dense_data_loader():
|
|
65
|
+
def create_dense_graph(num_nodes=4, num_features=6):
|
|
66
|
+
x = np.random.randn(num_nodes, num_features).astype(np.float32)
|
|
67
|
+
adj = np.random.randn(num_nodes, num_nodes).astype(np.float32)
|
|
68
|
+
if torch is not None:
|
|
69
|
+
x = torch.from_numpy(x)
|
|
70
|
+
adj = torch.from_numpy(adj)
|
|
71
|
+
return Data(x=x, adj=adj)
|
|
72
|
+
|
|
73
|
+
dataset = [create_dense_graph() for _ in range(4)]
|
|
74
|
+
loader = DenseDataLoader(dataset, batch_size=2, shuffle=False)
|
|
75
|
+
|
|
76
|
+
batch = next(iter(loader))
|
|
77
|
+
assert batch.x.shape == (2, 4, 6)
|
|
78
|
+
assert batch.adj.shape == (2, 4, 4)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def test_temporal_data_loader():
|
|
82
|
+
src = np.array([0, 1, 0, 2, 1, 3], dtype=np.int64)
|
|
83
|
+
dst = np.array([1, 2, 2, 3, 3, 0], dtype=np.int64)
|
|
84
|
+
t = np.array([1, 2, 3, 4, 5, 6], dtype=np.int64)
|
|
85
|
+
msg = np.random.randn(6, 4).astype(np.float32)
|
|
86
|
+
|
|
87
|
+
if torch is not None:
|
|
88
|
+
src = torch.from_numpy(src)
|
|
89
|
+
dst = torch.from_numpy(dst)
|
|
90
|
+
t = torch.from_numpy(t)
|
|
91
|
+
msg = torch.from_numpy(msg)
|
|
92
|
+
|
|
93
|
+
data = TemporalData(src=src, dst=dst, t=t, msg=msg)
|
|
94
|
+
loader = TemporalDataLoader(data, batch_size=3, neg_sampling_ratio=1.0)
|
|
95
|
+
|
|
96
|
+
batches = list(loader)
|
|
97
|
+
assert len(batches) == 2
|
|
98
|
+
assert len(batches[0]) == 3
|
|
99
|
+
assert hasattr(batches[0], 'neg_dst')
|
|
100
|
+
assert hasattr(batches[0], 'n_id')
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def test_random_node_loader():
|
|
104
|
+
data = create_dummy_data(num_nodes=10)
|
|
105
|
+
loader = RandomNodeLoader(data, num_parts=2)
|
|
106
|
+
|
|
107
|
+
parts = list(loader)
|
|
108
|
+
assert len(parts) == 2
|
|
109
|
+
for part in parts:
|
|
110
|
+
assert isinstance(part, Data)
|
|
111
|
+
assert part.num_nodes <= 10
|
|
112
|
+
assert hasattr(part, 'x')
|
|
113
|
+
|
|
@@ -0,0 +1,221 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data
|
|
6
|
+
from k3_node.layers import GCNConv
|
|
7
|
+
from k3_node.loader import DataLoader, FullGraphDataset, NeighborLoader
|
|
8
|
+
from k3_node.loader.keras_dataset import to_keras_batch
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _graph(num_nodes=30, num_classes=3):
|
|
12
|
+
rng = np.random.default_rng(0)
|
|
13
|
+
y = rng.integers(0, num_classes, num_nodes)
|
|
14
|
+
x = (np.eye(num_classes)[y] + 0.3 * rng.standard_normal((num_nodes, num_classes))).astype("float32")
|
|
15
|
+
edge_index = rng.integers(0, num_nodes, (2, 90)).astype("int32")
|
|
16
|
+
train_mask = np.zeros(num_nodes, dtype=bool)
|
|
17
|
+
train_mask[:10] = True
|
|
18
|
+
return Data(x=x, edge_index=edge_index, y=y.astype("int32"), train_mask=train_mask, test_mask=~train_mask)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class SmallGCN(keras.Model):
|
|
22
|
+
def __init__(self):
|
|
23
|
+
super().__init__()
|
|
24
|
+
self.conv = GCNConv(3, 3)
|
|
25
|
+
|
|
26
|
+
def call(self, data):
|
|
27
|
+
return self.conv(data.x, data.edge_index)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def test_fit_and_evaluate_on_full_graph():
|
|
31
|
+
data = _graph()
|
|
32
|
+
model = SmallGCN()
|
|
33
|
+
model.compile(keras.optimizers.Adam(0.05), keras.losses.SparseCategoricalCrossentropy(from_logits=True))
|
|
34
|
+
history = model.fit(FullGraphDataset(data, mask="train_mask"), epochs=30, verbose=0)
|
|
35
|
+
assert history.history["loss"][-1] < history.history["loss"][0]
|
|
36
|
+
model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
|
|
37
|
+
preds = model.predict(FullGraphDataset(data), verbose=0)
|
|
38
|
+
assert preds.shape == (30, 3)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def test_normalized_mask_gives_mean_over_masked_nodes():
|
|
42
|
+
data = _graph()
|
|
43
|
+
model = SmallGCN()
|
|
44
|
+
loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction=None)
|
|
45
|
+
model.compile(loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True))
|
|
46
|
+
reported = model.evaluate(FullGraphDataset(data, mask="train_mask"), verbose=0)
|
|
47
|
+
per_node = ops.convert_to_numpy(loss_fn(data.y, model(to_keras_batch(data)[0])))
|
|
48
|
+
np.testing.assert_allclose(reported, per_node[: 10].mean(), rtol=1e-5)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def test_graph_loader_works_with_fit_evaluate_predict():
|
|
52
|
+
from k3_node.datasets import FakeDataset
|
|
53
|
+
from k3_node.layers import GCNConv, global_mean_pool
|
|
54
|
+
|
|
55
|
+
class GraphClassifier(keras.Model):
|
|
56
|
+
def __init__(self):
|
|
57
|
+
super().__init__()
|
|
58
|
+
self.conv = GCNConv(8, 16)
|
|
59
|
+
self.head = keras.layers.Dense(3)
|
|
60
|
+
|
|
61
|
+
def call(self, data):
|
|
62
|
+
x = self.conv(data.x, data.edge_index)
|
|
63
|
+
return self.head(global_mean_pool(x, data.batch, data.num_graphs))
|
|
64
|
+
|
|
65
|
+
dataset = FakeDataset(num_graphs=20, avg_num_nodes=8, num_channels=8, num_classes=3)
|
|
66
|
+
loader = DataLoader(dataset, batch_size=6, shuffle=True)
|
|
67
|
+
model = GraphClassifier()
|
|
68
|
+
model.compile(optimizer="adam", loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=["accuracy"])
|
|
69
|
+
history = model.fit(loader, epochs=2, verbose=0) # compiled on TF/JAX: num_graphs must be static
|
|
70
|
+
assert len(history.history["loss"]) == 2
|
|
71
|
+
model.evaluate(DataLoader(dataset, batch_size=6), verbose=0)
|
|
72
|
+
assert model.predict(DataLoader(dataset, batch_size=6), verbose=0).shape == (20, 3)
|
|
73
|
+
# Iterating the loader directly still yields Batch objects
|
|
74
|
+
batch = next(iter(loader))
|
|
75
|
+
assert hasattr(batch, "edge_index") and batch.num_graphs == 6
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def test_neighbor_loader_weights_seed_nodes_only():
|
|
79
|
+
data = _graph()
|
|
80
|
+
loader = NeighborLoader(data, num_neighbors=[3], batch_size=4, input_nodes=np.arange(10))
|
|
81
|
+
inputs, y, weight = loader[0]
|
|
82
|
+
assert y.shape[0] == inputs.x.shape[0]
|
|
83
|
+
assert (weight[:4] > 0).all() and (weight[4:] == 0).all()
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def test_full_graph_dataset_index_split():
|
|
87
|
+
import numpy as np
|
|
88
|
+
from k3_node.data import Data
|
|
89
|
+
from k3_node.loader import FullGraphDataset
|
|
90
|
+
|
|
91
|
+
data = Data(x=np.ones((5, 2), "float32"), edge_index=np.array([[0, 1], [1, 2]]),
|
|
92
|
+
train_idx=np.array([3, 1]), train_y=np.array([2, 1]), num_nodes=5)
|
|
93
|
+
inputs, y, weight = FullGraphDataset(data, index="train_idx", target="train_y")[0]
|
|
94
|
+
assert "train_idx" not in inputs._fields and "train_y" not in inputs._fields
|
|
95
|
+
np.testing.assert_array_equal(y, [0, 1, 0, 2, 0])
|
|
96
|
+
np.testing.assert_allclose(weight, [0, 2.5, 0, 2.5, 0]) # mean over the 2 indexed nodes
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def test_full_graph_dataset_fresh_negatives_every_epoch():
|
|
100
|
+
import numpy as np
|
|
101
|
+
from k3_node.data import Data
|
|
102
|
+
from k3_node.loader import FullGraphDataset
|
|
103
|
+
|
|
104
|
+
edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 4]])
|
|
105
|
+
data = Data(x=np.ones((50, 2), "float32"), edge_index=edge_index, edge_label_index=edge_index,
|
|
106
|
+
edge_label=np.ones(4, "float32"), num_nodes=50)
|
|
107
|
+
dataset = FullGraphDataset(data, neg_sampling_ratio=2.0)
|
|
108
|
+
(inputs1, y1), (inputs2, _) = dataset[0], dataset[0]
|
|
109
|
+
assert inputs1.edge_label_index.shape == (2, 12)
|
|
110
|
+
np.testing.assert_array_equal(y1, [1] * 4 + [0] * 8)
|
|
111
|
+
assert not np.array_equal(inputs1.edge_label_index[:, 4:], inputs2.edge_label_index[:, 4:])
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def test_link_neighbor_loader_uses_local_ids():
|
|
115
|
+
import numpy as np
|
|
116
|
+
from k3_node.data import Data
|
|
117
|
+
from k3_node.loader import LinkNeighborLoader
|
|
118
|
+
|
|
119
|
+
rng = np.random.default_rng(0)
|
|
120
|
+
edge_index = rng.integers(0, 200, size=(2, 600))
|
|
121
|
+
data = Data(x=np.arange(200, dtype="float32")[:, None], edge_index=edge_index, num_nodes=200)
|
|
122
|
+
loader = LinkNeighborLoader(data, num_neighbors=[5], batch_size=16, neg_sampling_ratio=1.0)
|
|
123
|
+
batch = next(iter(loader))
|
|
124
|
+
eli = np.asarray(batch.edge_label_index)
|
|
125
|
+
n_id = np.asarray(batch.n_id)
|
|
126
|
+
assert eli.max() < n_id.shape[0] # local ids into the sampled subgraph
|
|
127
|
+
label = np.asarray(batch.edge_label)
|
|
128
|
+
pos = n_id[eli[:, label == 1]] # back to global ids: must be real edges
|
|
129
|
+
real = set(map(tuple, edge_index.T.tolist()))
|
|
130
|
+
assert all(tuple(e) in real for e in pos.T.tolist())
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def test_loader_with_mask():
|
|
134
|
+
import numpy as np
|
|
135
|
+
from k3_node.data import Data
|
|
136
|
+
from k3_node.loader import ClusterData, ClusterLoader
|
|
137
|
+
|
|
138
|
+
rng = np.random.default_rng(0)
|
|
139
|
+
data = Data(x=rng.random((60, 3)).astype("float32"), edge_index=rng.integers(0, 60, (2, 200)),
|
|
140
|
+
y=rng.integers(0, 3, 60), train_mask=np.arange(60) < 30, num_nodes=60)
|
|
141
|
+
loader = ClusterLoader(ClusterData(data, num_parts=4), batch_size=2).with_mask("train_mask")
|
|
142
|
+
inputs, y, weight = loader[0]
|
|
143
|
+
assert weight.shape == y.shape and "train_mask" not in inputs._fields
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def test_shadow_roots_and_labels():
|
|
147
|
+
import numpy as np
|
|
148
|
+
from k3_node.data import Data
|
|
149
|
+
from k3_node.loader import ShaDowKHopSampler
|
|
150
|
+
|
|
151
|
+
rng = np.random.default_rng(0)
|
|
152
|
+
data = Data(x=np.arange(50, dtype="float32")[:, None], edge_index=rng.integers(0, 50, (2, 200)),
|
|
153
|
+
y=np.arange(50), num_nodes=50)
|
|
154
|
+
loader = ShaDowKHopSampler(data, depth=2, num_neighbors=3, node_idx=np.arange(10, 20), batch_size=5)
|
|
155
|
+
batch = next(iter(loader))
|
|
156
|
+
root_x = np.asarray(batch.x)[np.asarray(batch.root_n_id), 0]
|
|
157
|
+
np.testing.assert_array_equal(root_x, np.arange(10, 15)) # the roots are the seed nodes
|
|
158
|
+
np.testing.assert_array_equal(np.asarray(batch.y), np.arange(10, 15)) # one label per subgraph
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
def test_graph_saint_samplers():
|
|
162
|
+
import numpy as np
|
|
163
|
+
from k3_node.data import Data
|
|
164
|
+
from k3_node.loader import GraphSAINTEdgeSampler, GraphSAINTNodeSampler, GraphSAINTRandomWalkSampler
|
|
165
|
+
|
|
166
|
+
rng = np.random.default_rng(0)
|
|
167
|
+
edge_index = rng.integers(0, 100, (2, 400))
|
|
168
|
+
data = Data(x=np.arange(100, dtype="float32")[:, None], edge_index=edge_index,
|
|
169
|
+
edge_attr=np.arange(400, dtype="float32"), y=np.arange(100), num_nodes=100)
|
|
170
|
+
for loader in [GraphSAINTNodeSampler(data, batch_size=30, num_steps=3, sample_coverage=5),
|
|
171
|
+
GraphSAINTEdgeSampler(data, batch_size=20, num_steps=3, sample_coverage=5),
|
|
172
|
+
GraphSAINTRandomWalkSampler(data, batch_size=10, walk_length=2, num_steps=3, sample_coverage=5)]:
|
|
173
|
+
batches = list(loader)
|
|
174
|
+
assert len(batches) == 3
|
|
175
|
+
batch = batches[0]
|
|
176
|
+
x = np.asarray(batch.x)[:, 0].astype(int)
|
|
177
|
+
ei = np.asarray(batch.edge_index)
|
|
178
|
+
# every subgraph edge is a real edge between sampled nodes, carrying its own attribute
|
|
179
|
+
real = {tuple(e): i for i, e in enumerate(edge_index.T.tolist())}
|
|
180
|
+
for (a, b), attr in zip(ei.T.tolist(), np.asarray(batch.edge_attr).astype(int)):
|
|
181
|
+
assert (x[a], x[b]) in real and tuple(edge_index[:, attr]) == (x[a], x[b])
|
|
182
|
+
assert np.asarray(batch.node_norm).shape == (batch.num_nodes,)
|
|
183
|
+
assert np.asarray(batch.edge_norm).shape == (ei.shape[1],)
|
|
184
|
+
inputs, y = loader[0] # Keras batch
|
|
185
|
+
assert inputs.x.shape[0] == y.shape[0]
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def test_neighbor_loader_disjoint():
|
|
189
|
+
import numpy as np
|
|
190
|
+
from k3_node.data import Data
|
|
191
|
+
from k3_node.loader import NeighborLoader
|
|
192
|
+
|
|
193
|
+
rng = np.random.default_rng(0)
|
|
194
|
+
edge_index = rng.integers(0, 30, (2, 150))
|
|
195
|
+
data = Data(x=np.arange(30, dtype="float32")[:, None], edge_index=edge_index, num_nodes=30)
|
|
196
|
+
batch = next(iter(NeighborLoader(data, num_neighbors=[3, 2], batch_size=8, disjoint=True)))
|
|
197
|
+
b, n_id, ei = np.asarray(batch.batch), np.asarray(batch.n_id), np.asarray(batch.edge_index)
|
|
198
|
+
np.testing.assert_array_equal(b[:8], np.arange(8)) # seeds first, one subgraph each
|
|
199
|
+
assert np.all(b[ei[0]] == b[ei[1]]) # edges never cross subgraphs
|
|
200
|
+
real = set(map(tuple, edge_index.T.tolist()))
|
|
201
|
+
assert all((n_id[s], n_id[t]) in real for s, t in ei.T.tolist())
|
|
202
|
+
for g in range(8): # no node appears twice within one subgraph
|
|
203
|
+
assert len(set(n_id[b == g].tolist())) == int((b == g).sum())
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def test_full_graph_dataset_hetero():
|
|
207
|
+
import numpy as np
|
|
208
|
+
from k3_node.data import HeteroData
|
|
209
|
+
from k3_node.loader import FullGraphDataset
|
|
210
|
+
|
|
211
|
+
data = HeteroData()
|
|
212
|
+
data["user"].x = np.ones((4, 2), "float32")
|
|
213
|
+
data["user"].y = np.array([0, 1, 0, 1])
|
|
214
|
+
data["user"].train_mask = np.array([True, True, False, False])
|
|
215
|
+
data["item"].x = np.ones((3, 5), "float32")
|
|
216
|
+
data["user", "buys", "item"].edge_index = np.array([[0, 1, 3], [0, 2, 1]])
|
|
217
|
+
inputs, y, weight = FullGraphDataset(data, node_type="user", mask="train_mask")[0]
|
|
218
|
+
assert set(inputs.x_dict) == {"user", "item"}
|
|
219
|
+
assert ("user", "buys", "item") in inputs.edge_index_dict
|
|
220
|
+
np.testing.assert_array_equal(y, [0, 1, 0, 1])
|
|
221
|
+
np.testing.assert_allclose(weight, [2, 2, 0, 0])
|
|
@@ -0,0 +1,122 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
try:
|
|
4
|
+
import torch
|
|
5
|
+
except ImportError:
|
|
6
|
+
torch = None
|
|
7
|
+
|
|
8
|
+
from k3_node.data import Data, HeteroData
|
|
9
|
+
from k3_node.loader import HGTLoader, LinkNeighborLoader, NeighborLoader
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def get_homo_graph():
|
|
13
|
+
# 6 nodes connected in a ring: 0->1->2->3->4->5->0 and some cross edges
|
|
14
|
+
edge_index = np.array([
|
|
15
|
+
[0, 1, 2, 3, 4, 5, 0, 2],
|
|
16
|
+
[1, 2, 3, 4, 5, 0, 3, 5],
|
|
17
|
+
], dtype=np.int64)
|
|
18
|
+
x = np.random.randn(6, 16).astype(np.float32)
|
|
19
|
+
y = np.array([0, 1, 0, 1, 0, 1], dtype=np.int64)
|
|
20
|
+
|
|
21
|
+
if torch is not None:
|
|
22
|
+
edge_index = torch.from_numpy(edge_index)
|
|
23
|
+
x = torch.from_numpy(x)
|
|
24
|
+
y = torch.from_numpy(y)
|
|
25
|
+
|
|
26
|
+
return Data(x=x, edge_index=edge_index, y=y)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def get_hetero_graph():
|
|
30
|
+
data = HeteroData()
|
|
31
|
+
data['paper'].x = np.random.randn(10, 8).astype(np.float32)
|
|
32
|
+
data['author'].x = np.random.randn(5, 8).astype(np.float32)
|
|
33
|
+
|
|
34
|
+
data['author', 'writes', 'paper'].edge_index = np.array([
|
|
35
|
+
[0, 1, 2, 3, 4, 0, 1],
|
|
36
|
+
[0, 1, 2, 3, 4, 5, 6],
|
|
37
|
+
], dtype=np.int64)
|
|
38
|
+
data['paper', 'cites', 'paper'].edge_index = np.array([
|
|
39
|
+
[0, 1, 2, 3],
|
|
40
|
+
[1, 2, 3, 4],
|
|
41
|
+
], dtype=np.int64)
|
|
42
|
+
|
|
43
|
+
if torch is not None:
|
|
44
|
+
data['paper'].x = torch.from_numpy(data['paper'].x)
|
|
45
|
+
data['author'].x = torch.from_numpy(data['author'].x)
|
|
46
|
+
data['author', 'writes', 'paper'].edge_index = torch.from_numpy(data['author', 'writes', 'paper'].edge_index)
|
|
47
|
+
data['paper', 'cites', 'paper'].edge_index = torch.from_numpy(data['paper', 'cites', 'paper'].edge_index)
|
|
48
|
+
|
|
49
|
+
return data
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def test_homo_neighbor_loader():
|
|
53
|
+
data = get_homo_graph()
|
|
54
|
+
input_nodes = [0, 1]
|
|
55
|
+
if torch is not None:
|
|
56
|
+
input_nodes = torch.tensor(input_nodes, dtype=torch.long)
|
|
57
|
+
|
|
58
|
+
loader = NeighborLoader(
|
|
59
|
+
data,
|
|
60
|
+
num_neighbors=[2, 2],
|
|
61
|
+
batch_size=2,
|
|
62
|
+
input_nodes=input_nodes,
|
|
63
|
+
shuffle=False,
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
batch = next(iter(loader))
|
|
67
|
+
assert batch.batch_size == 2
|
|
68
|
+
assert hasattr(batch, 'n_id')
|
|
69
|
+
assert hasattr(batch, 'e_id')
|
|
70
|
+
assert hasattr(batch, 'num_sampled_nodes')
|
|
71
|
+
assert hasattr(batch, 'num_sampled_edges')
|
|
72
|
+
assert batch.x.shape[0] == len(batch.n_id)
|
|
73
|
+
assert batch.edge_index.shape[0] == 2
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def test_hetero_neighbor_loader():
|
|
77
|
+
data = get_hetero_graph()
|
|
78
|
+
loader = NeighborLoader(
|
|
79
|
+
data,
|
|
80
|
+
num_neighbors=[2, 2],
|
|
81
|
+
batch_size=2,
|
|
82
|
+
input_nodes=('paper', [0, 1]),
|
|
83
|
+
shuffle=False,
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
batch = next(iter(loader))
|
|
87
|
+
assert batch['paper'].batch_size == 2
|
|
88
|
+
assert hasattr(batch['paper'], 'n_id')
|
|
89
|
+
assert hasattr(batch['author', 'writes', 'paper'], 'edge_index')
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def test_link_neighbor_loader():
|
|
93
|
+
data = get_homo_graph()
|
|
94
|
+
loader = LinkNeighborLoader(
|
|
95
|
+
data,
|
|
96
|
+
num_neighbors=[2, 2],
|
|
97
|
+
batch_size=2,
|
|
98
|
+
neg_sampling_ratio=1.0,
|
|
99
|
+
shuffle=False,
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
batch = next(iter(loader))
|
|
103
|
+
assert hasattr(batch, 'edge_label_index')
|
|
104
|
+
assert hasattr(batch, 'edge_label')
|
|
105
|
+
assert hasattr(batch, 'n_id')
|
|
106
|
+
assert batch.edge_label.shape[0] == 4 # 2 positive + 2 negative
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def test_hgt_loader():
|
|
110
|
+
data = get_hetero_graph()
|
|
111
|
+
loader = HGTLoader(
|
|
112
|
+
data,
|
|
113
|
+
num_samples=[4, 4],
|
|
114
|
+
input_nodes=('paper', [0, 1]),
|
|
115
|
+
batch_size=2,
|
|
116
|
+
shuffle=False,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
batch = next(iter(loader))
|
|
120
|
+
assert batch['paper'].batch_size == 2
|
|
121
|
+
assert hasattr(batch['paper'], 'n_id')
|
|
122
|
+
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
|
|
4
|
+
from k3_node.loader.sampler_utils import FastGraph, sample_neighbors_homo
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def _reference_full_sampling(edge_index, seeds, num_hops, subgraph_type):
|
|
8
|
+
"""The original per-node Python algorithm, restricted to k=-1 where it is deterministic."""
|
|
9
|
+
graph = FastGraph(edge_index)
|
|
10
|
+
nodes, visited = [], {}
|
|
11
|
+
for s in seeds:
|
|
12
|
+
if int(s) not in visited:
|
|
13
|
+
visited[int(s)] = len(nodes)
|
|
14
|
+
nodes.append(int(s))
|
|
15
|
+
frontier, sampled = list(nodes), []
|
|
16
|
+
for _ in range(num_hops):
|
|
17
|
+
next_frontier = []
|
|
18
|
+
for target in frontier:
|
|
19
|
+
srcs, e_ids = graph.get_neighbors(target)
|
|
20
|
+
for s, e in zip(srcs, e_ids):
|
|
21
|
+
sampled.append((int(s), target, int(e)))
|
|
22
|
+
if int(s) not in visited:
|
|
23
|
+
visited[int(s)] = len(nodes)
|
|
24
|
+
nodes.append(int(s))
|
|
25
|
+
next_frontier.append(int(s))
|
|
26
|
+
frontier = next_frontier
|
|
27
|
+
if subgraph_type == "induced":
|
|
28
|
+
edges = [(visited[u], visited[v], e) for e, (u, v) in enumerate(edge_index.T) if u in visited and v in visited]
|
|
29
|
+
elif subgraph_type == "bidirectional":
|
|
30
|
+
edges = [x for u, v, e in sampled for x in ((visited[u], visited[v], e), (visited[v], visited[u], e))]
|
|
31
|
+
else:
|
|
32
|
+
edges = [(visited[u], visited[v], e) for u, v, e in sampled]
|
|
33
|
+
edges = np.array(edges, dtype=np.int64).reshape(-1, 3)
|
|
34
|
+
return np.array(nodes), edges[:, 0], edges[:, 1], edges[:, 2]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def _random_graph(num_nodes=40, num_edges=160, seed=0):
|
|
38
|
+
rng = np.random.default_rng(seed)
|
|
39
|
+
return rng.integers(0, num_nodes, (2, num_edges)).astype(np.int64)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
@pytest.mark.parametrize("subgraph_type", ["directional", "bidirectional", "induced"])
|
|
43
|
+
def test_full_neighborhood_matches_reference(subgraph_type):
|
|
44
|
+
edge_index = _random_graph()
|
|
45
|
+
seeds = np.array([3, 7, 3, 11])
|
|
46
|
+
expected = _reference_full_sampling(edge_index, seeds, 2, subgraph_type)
|
|
47
|
+
nodes, row, col, edge, _, _ = sample_neighbors_homo(edge_index, seeds, [-1, -1], subgraph_type=subgraph_type)
|
|
48
|
+
for got, want in zip((nodes, row, col, edge), expected):
|
|
49
|
+
np.testing.assert_array_equal(got, want)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@pytest.mark.parametrize("replace", [False, True])
|
|
53
|
+
def test_sampled_neighborhood_properties(replace):
|
|
54
|
+
edge_index = _random_graph()
|
|
55
|
+
graph = FastGraph(edge_index)
|
|
56
|
+
seeds = np.array([0, 5, 9])
|
|
57
|
+
nodes, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
|
|
58
|
+
edge_index, seeds, [3, 2], replace=replace, graph=graph
|
|
59
|
+
)
|
|
60
|
+
assert len(np.unique(nodes)) == len(nodes)
|
|
61
|
+
np.testing.assert_array_equal(nodes[:3], seeds)
|
|
62
|
+
assert sum(n_counts) == len(nodes) and sum(e_counts) == len(edge)
|
|
63
|
+
# Every sampled edge is a real edge, with endpoints mapped to their local ids.
|
|
64
|
+
np.testing.assert_array_equal(nodes[row], edge_index[0, edge])
|
|
65
|
+
np.testing.assert_array_equal(nodes[col], edge_index[1, edge])
|
|
66
|
+
# First hop: each seed gets min(3, in-degree) edges, distinct when sampling without replacement.
|
|
67
|
+
in_degree = np.bincount(edge_index[1], minlength=40)
|
|
68
|
+
first_hop = slice(0, e_counts[0])
|
|
69
|
+
for i, s in enumerate(seeds):
|
|
70
|
+
picked = edge[first_hop][col[first_hop] == i]
|
|
71
|
+
assert len(picked) == min(3, in_degree[s])
|
|
72
|
+
if not replace:
|
|
73
|
+
assert len(np.unique(picked)) == len(picked)
|
|
74
|
+
# The reusable id map is left clean for the next batch.
|
|
75
|
+
assert (graph.local_map(40) == -1).all()
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def test_disjoint_and_out_of_range_seeds():
|
|
79
|
+
edge_index = np.array([[1, 2], [0, 0]])
|
|
80
|
+
nodes, row, col, edge, _, _ = sample_neighbors_homo(edge_index, np.array([0, 5]), [-1], num_nodes=3, disjoint=True)
|
|
81
|
+
np.testing.assert_array_equal(nodes, [0, 5, 1, 2])
|
|
82
|
+
np.testing.assert_array_equal(edge, [0, 1])
|