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,202 @@
|
|
|
1
|
+
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
from torch import Tensor
|
|
8
|
+
except ImportError:
|
|
9
|
+
torch = None
|
|
10
|
+
Tensor = type(None)
|
|
11
|
+
|
|
12
|
+
from k3_node.data import Data, HeteroData
|
|
13
|
+
from k3_node.loader.link_loader import EdgeSamplerInput, LinkLoader
|
|
14
|
+
from k3_node.loader.node_loader import HeteroSamplerOutput, SamplerOutput
|
|
15
|
+
from k3_node.loader.sampler_utils import FastGraph, sample_neighbors_hetero, sample_neighbors_homo
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class InternalLinkNeighborSampler:
|
|
19
|
+
r"""Pure Python/NumPy link neighborhood sampler."""
|
|
20
|
+
def __init__(
|
|
21
|
+
self,
|
|
22
|
+
data: Union[Data, HeteroData],
|
|
23
|
+
num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
|
|
24
|
+
replace: bool = False,
|
|
25
|
+
subgraph_type: str = 'directional',
|
|
26
|
+
disjoint: bool = False,
|
|
27
|
+
neg_sampling_ratio: float = 0.0,
|
|
28
|
+
):
|
|
29
|
+
self.data = data
|
|
30
|
+
self.num_neighbors = num_neighbors
|
|
31
|
+
self.replace = replace
|
|
32
|
+
self.subgraph_type = subgraph_type
|
|
33
|
+
self.disjoint = disjoint
|
|
34
|
+
self.neg_sampling_ratio = neg_sampling_ratio
|
|
35
|
+
self.edge_permutation = None
|
|
36
|
+
|
|
37
|
+
def _cached_graph(self):
|
|
38
|
+
# The CSR index depends only on the graph; build it once instead of per batch.
|
|
39
|
+
if getattr(self, "_graph", None) is None:
|
|
40
|
+
self._graph = FastGraph(self.data.edge_index, num_nodes=self.data.num_nodes)
|
|
41
|
+
return self._graph
|
|
42
|
+
|
|
43
|
+
def sample_from_edges(self, input_data: EdgeSamplerInput) -> Union[SamplerOutput, HeteroSamplerOutput]:
|
|
44
|
+
is_torch = torch is not None and isinstance(input_data.row, Tensor)
|
|
45
|
+
|
|
46
|
+
pos_row = input_data.row
|
|
47
|
+
pos_col = input_data.col
|
|
48
|
+
pos_len = pos_row.size(0) if is_torch else len(pos_row)
|
|
49
|
+
|
|
50
|
+
if self.neg_sampling_ratio > 0:
|
|
51
|
+
num_neg = round(self.neg_sampling_ratio * pos_len)
|
|
52
|
+
num_total_nodes = self.data.num_nodes
|
|
53
|
+
if is_torch:
|
|
54
|
+
# As in PyG, both endpoints of a negative are random nodes
|
|
55
|
+
neg_row = torch.randint(0, num_total_nodes, (num_neg,), dtype=torch.long, device=pos_row.device)
|
|
56
|
+
neg_col = torch.randint(0, num_total_nodes, (num_neg,), dtype=torch.long, device=pos_row.device)
|
|
57
|
+
total_row = torch.cat([pos_row, neg_row], dim=0)
|
|
58
|
+
total_col = torch.cat([pos_col, neg_col], dim=0)
|
|
59
|
+
edge_label = torch.cat([torch.ones(pos_len, dtype=torch.float, device=pos_row.device),
|
|
60
|
+
torch.zeros(num_neg, dtype=torch.float, device=pos_row.device)], dim=0)
|
|
61
|
+
else:
|
|
62
|
+
neg_row = np.random.randint(0, num_total_nodes, size=(num_neg,), dtype=np.int64)
|
|
63
|
+
neg_col = np.random.randint(0, num_total_nodes, size=(num_neg,), dtype=np.int64)
|
|
64
|
+
total_row = np.concatenate([pos_row, neg_row], axis=0)
|
|
65
|
+
total_col = np.concatenate([pos_col, neg_col], axis=0)
|
|
66
|
+
edge_label = np.concatenate([np.ones(pos_len, dtype=np.float32), np.zeros(num_neg, dtype=np.float32)], axis=0)
|
|
67
|
+
else:
|
|
68
|
+
total_row = pos_row
|
|
69
|
+
total_col = pos_col
|
|
70
|
+
edge_label = input_data.label
|
|
71
|
+
|
|
72
|
+
# The sampled subgraph starts with the (sorted, unique) seed nodes, so `edge_label_index`
|
|
73
|
+
# is relabeled to their positions, i.e. to local node ids, as in PyG.
|
|
74
|
+
def unique_inverse(values):
|
|
75
|
+
if is_torch:
|
|
76
|
+
return torch.unique(values, return_inverse=True)
|
|
77
|
+
return np.unique(np.asarray(values), return_inverse=True)
|
|
78
|
+
|
|
79
|
+
def stack(a, b):
|
|
80
|
+
return torch.stack([a, b], dim=0) if is_torch else np.stack([a, b], axis=0)
|
|
81
|
+
|
|
82
|
+
if is_torch:
|
|
83
|
+
seed_nodes, inverse = unique_inverse(torch.cat([total_row, total_col], dim=0))
|
|
84
|
+
else:
|
|
85
|
+
seed_nodes, inverse = unique_inverse(np.concatenate([total_row, total_col], axis=0))
|
|
86
|
+
edge_label_index = stack(inverse[:len(total_row)], inverse[len(total_row):])
|
|
87
|
+
|
|
88
|
+
if isinstance(self.data, Data):
|
|
89
|
+
node, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
|
|
90
|
+
edge_index=self.data.edge_index,
|
|
91
|
+
seed_nodes=seed_nodes,
|
|
92
|
+
num_neighbors=self.num_neighbors,
|
|
93
|
+
num_nodes=self.data.num_nodes,
|
|
94
|
+
graph=self._cached_graph(),
|
|
95
|
+
replace=self.replace,
|
|
96
|
+
subgraph_type=self.subgraph_type,
|
|
97
|
+
disjoint=self.disjoint,
|
|
98
|
+
)
|
|
99
|
+
return SamplerOutput(
|
|
100
|
+
node=node,
|
|
101
|
+
row=row,
|
|
102
|
+
col=col,
|
|
103
|
+
edge=edge,
|
|
104
|
+
num_sampled_nodes=n_counts,
|
|
105
|
+
num_sampled_edges=e_counts,
|
|
106
|
+
metadata=(input_data.input_id, edge_label_index, edge_label),
|
|
107
|
+
)
|
|
108
|
+
elif isinstance(self.data, HeteroData):
|
|
109
|
+
edge_index_dict = {k: self.data[k].edge_index for k in self.data.edge_types}
|
|
110
|
+
input_type = input_data.input_type or self.data.edge_types[0]
|
|
111
|
+
src_type, _, dst_type = input_type
|
|
112
|
+
|
|
113
|
+
seed_dict = {k: None for k in self.data.node_types}
|
|
114
|
+
if src_type == dst_type:
|
|
115
|
+
seed_dict[src_type] = seed_nodes
|
|
116
|
+
else:
|
|
117
|
+
seed_dict[src_type], src_inverse = unique_inverse(total_row)
|
|
118
|
+
seed_dict[dst_type], dst_inverse = unique_inverse(total_col)
|
|
119
|
+
edge_label_index = stack(src_inverse, dst_inverse)
|
|
120
|
+
|
|
121
|
+
norm_num_neighbors = self.num_neighbors
|
|
122
|
+
if isinstance(norm_num_neighbors, dict):
|
|
123
|
+
norm_num_neighbors = {self.data._to_canonical(*k): v for k, v in norm_num_neighbors.items()}
|
|
124
|
+
|
|
125
|
+
node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
|
|
126
|
+
edge_index_dict=edge_index_dict,
|
|
127
|
+
seed_nodes_dict=seed_dict,
|
|
128
|
+
num_neighbors=norm_num_neighbors,
|
|
129
|
+
replace=self.replace,
|
|
130
|
+
subgraph_type=self.subgraph_type,
|
|
131
|
+
)
|
|
132
|
+
return HeteroSamplerOutput(
|
|
133
|
+
node=node_dict,
|
|
134
|
+
row=row_dict,
|
|
135
|
+
col=col_dict,
|
|
136
|
+
edge=edge_dict,
|
|
137
|
+
num_sampled_nodes=n_counts,
|
|
138
|
+
num_sampled_edges=e_counts,
|
|
139
|
+
metadata=(input_data.input_id, edge_label_index, edge_label),
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
raise TypeError(f"Invalid data type: {type(self.data)}")
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
class LinkNeighborLoader(LinkLoader):
|
|
146
|
+
r"""A link-based data loader derived as an extension of NeighborLoader."""
|
|
147
|
+
def __init__(
|
|
148
|
+
self,
|
|
149
|
+
data: Union[Data, HeteroData],
|
|
150
|
+
num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
|
|
151
|
+
edge_label_index: Any = None,
|
|
152
|
+
edge_label: Optional[Any] = None,
|
|
153
|
+
edge_label_time: Optional[Any] = None,
|
|
154
|
+
replace: bool = False,
|
|
155
|
+
subgraph_type: str = 'directional',
|
|
156
|
+
disjoint: bool = False,
|
|
157
|
+
temporal_strategy: str = 'uniform',
|
|
158
|
+
neg_sampling: Optional[Any] = None,
|
|
159
|
+
neg_sampling_ratio: Optional[Union[int, float]] = None,
|
|
160
|
+
time_attr: Optional[str] = None,
|
|
161
|
+
weight_attr: Optional[str] = None,
|
|
162
|
+
transform: Optional[Callable] = None,
|
|
163
|
+
transform_sampler_output: Optional[Callable] = None,
|
|
164
|
+
is_sorted: bool = False,
|
|
165
|
+
filter_per_worker: Optional[bool] = None,
|
|
166
|
+
neighbor_sampler: Optional[Any] = None,
|
|
167
|
+
directed: bool = True,
|
|
168
|
+
**kwargs,
|
|
169
|
+
):
|
|
170
|
+
if not directed:
|
|
171
|
+
subgraph_type = 'induced'
|
|
172
|
+
|
|
173
|
+
ratio = 0.0
|
|
174
|
+
if neg_sampling_ratio is not None:
|
|
175
|
+
ratio = float(neg_sampling_ratio)
|
|
176
|
+
elif neg_sampling is not None:
|
|
177
|
+
ratio = float(getattr(neg_sampling, 'amount', 1.0)) if hasattr(neg_sampling, 'amount') else 1.0
|
|
178
|
+
|
|
179
|
+
if neighbor_sampler is None:
|
|
180
|
+
neighbor_sampler = InternalLinkNeighborSampler(
|
|
181
|
+
data,
|
|
182
|
+
num_neighbors=num_neighbors,
|
|
183
|
+
replace=replace,
|
|
184
|
+
subgraph_type=str(getattr(subgraph_type, 'value', subgraph_type)),
|
|
185
|
+
disjoint=disjoint,
|
|
186
|
+
neg_sampling_ratio=ratio,
|
|
187
|
+
)
|
|
188
|
+
|
|
189
|
+
super().__init__(
|
|
190
|
+
data=data,
|
|
191
|
+
link_sampler=neighbor_sampler,
|
|
192
|
+
edge_label_index=edge_label_index,
|
|
193
|
+
edge_label=edge_label,
|
|
194
|
+
edge_label_time=edge_label_time,
|
|
195
|
+
neg_sampling=neg_sampling,
|
|
196
|
+
neg_sampling_ratio=neg_sampling_ratio,
|
|
197
|
+
transform=transform,
|
|
198
|
+
transform_sampler_output=transform_sampler_output,
|
|
199
|
+
filter_per_worker=filter_per_worker,
|
|
200
|
+
**kwargs,
|
|
201
|
+
)
|
|
202
|
+
|
k3_node/loader/mixin.py
ADDED
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
import glob
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
import os.path as osp
|
|
5
|
+
import warnings
|
|
6
|
+
from contextlib import contextmanager
|
|
7
|
+
from typing import Any, Callable, Dict, List, Optional, Union
|
|
8
|
+
|
|
9
|
+
try:
|
|
10
|
+
import psutil
|
|
11
|
+
except ImportError:
|
|
12
|
+
psutil = None
|
|
13
|
+
|
|
14
|
+
try:
|
|
15
|
+
import torch
|
|
16
|
+
except ImportError:
|
|
17
|
+
torch = None
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def get_numa_nodes_cores() -> Dict[str, Any]:
|
|
21
|
+
"""Parses numa nodes information into a dictionary."""
|
|
22
|
+
numa_node_paths = glob.glob('/sys/devices/system/node/node[0-9]*')
|
|
23
|
+
if not numa_node_paths:
|
|
24
|
+
return {}
|
|
25
|
+
|
|
26
|
+
nodes = {}
|
|
27
|
+
try:
|
|
28
|
+
for node_path in numa_node_paths:
|
|
29
|
+
numa_node_id = int(osp.basename(node_path)[4:])
|
|
30
|
+
thread_siblings = {}
|
|
31
|
+
for cpu_dir in glob.glob(osp.join(node_path, 'cpu[0-9]*')):
|
|
32
|
+
cpu_id = int(osp.basename(cpu_dir)[3:])
|
|
33
|
+
if cpu_id > 0:
|
|
34
|
+
with open(osp.join(cpu_dir, 'online')) as core_online_file:
|
|
35
|
+
core_online = int(core_online_file.read().splitlines()[0])
|
|
36
|
+
else:
|
|
37
|
+
core_online = 1 # cpu0 is always online
|
|
38
|
+
if core_online == 1:
|
|
39
|
+
with open(osp.join(cpu_dir, 'topology', 'core_id')) as core_id_file:
|
|
40
|
+
core_id = int(core_id_file.read().strip())
|
|
41
|
+
if core_id in thread_siblings:
|
|
42
|
+
thread_siblings[core_id].append(cpu_id)
|
|
43
|
+
else:
|
|
44
|
+
thread_siblings[core_id] = [cpu_id]
|
|
45
|
+
|
|
46
|
+
nodes[numa_node_id] = sorted([(k, sorted(v)) for k, v in thread_siblings.items()])
|
|
47
|
+
except (OSError, ValueError, IndexError):
|
|
48
|
+
warnings.warn('Failed to read NUMA info')
|
|
49
|
+
return {}
|
|
50
|
+
|
|
51
|
+
return nodes
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class WorkerInitWrapper:
|
|
55
|
+
r"""Wraps the :attr:`worker_init_fn` argument for DataLoader workers."""
|
|
56
|
+
def __init__(self, func: Optional[Callable]) -> None:
|
|
57
|
+
self.func = func
|
|
58
|
+
|
|
59
|
+
def __call__(self, worker_id: int) -> None:
|
|
60
|
+
if self.func is not None:
|
|
61
|
+
self.func(worker_id)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class LogMemoryMixin:
|
|
65
|
+
r"""A context manager to enable logging of memory consumption in
|
|
66
|
+
DataLoader workers.
|
|
67
|
+
"""
|
|
68
|
+
def _mem_init_fn(self, worker_id: int) -> None:
|
|
69
|
+
if psutil is not None:
|
|
70
|
+
proc = psutil.Process(os.getpid())
|
|
71
|
+
memory = proc.memory_info().rss / (1024 * 1024)
|
|
72
|
+
logging.debug(f"Worker {worker_id} @ PID {proc.pid}: {memory:.2f} MB")
|
|
73
|
+
self._old_worker_init_fn(worker_id)
|
|
74
|
+
|
|
75
|
+
@contextmanager
|
|
76
|
+
def enable_memory_log(self):
|
|
77
|
+
self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
|
|
78
|
+
try:
|
|
79
|
+
self.worker_init_fn = self._mem_init_fn
|
|
80
|
+
yield
|
|
81
|
+
finally:
|
|
82
|
+
self.worker_init_fn = self._old_worker_init_fn
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class MultithreadingMixin:
|
|
86
|
+
r"""A context manager to enable multi-threading in DataLoader workers."""
|
|
87
|
+
def _mt_init_fn(self, worker_id: int) -> None:
|
|
88
|
+
if torch is not None:
|
|
89
|
+
try:
|
|
90
|
+
torch.set_num_threads(int(self._worker_threads))
|
|
91
|
+
except IndexError as e:
|
|
92
|
+
raise ValueError(f"Cannot set {self._worker_threads} threads in worker {worker_id}") from e
|
|
93
|
+
self._old_worker_init_fn(worker_id)
|
|
94
|
+
|
|
95
|
+
@contextmanager
|
|
96
|
+
def enable_multithreading(self, worker_threads: Optional[int] = None):
|
|
97
|
+
num_workers = getattr(self, 'num_workers', 0)
|
|
98
|
+
if not num_workers > 0:
|
|
99
|
+
raise ValueError(f"'enable_multithreading' needs to be performed with at least one worker (got {num_workers})")
|
|
100
|
+
|
|
101
|
+
if torch is not None:
|
|
102
|
+
if worker_threads is None:
|
|
103
|
+
worker_threads = torch.get_num_threads() // num_workers
|
|
104
|
+
if worker_threads > torch.get_num_threads():
|
|
105
|
+
raise ValueError(
|
|
106
|
+
f"'worker_threads' should be smaller than total available threads {torch.get_num_threads()} (got {worker_threads})"
|
|
107
|
+
)
|
|
108
|
+
context = torch.multiprocessing.get_context()._name
|
|
109
|
+
if context != 'spawn':
|
|
110
|
+
raise ValueError(f"'enable_multithreading' can only be used with 'spawn' multiprocessing context (got {context})")
|
|
111
|
+
else:
|
|
112
|
+
if worker_threads is None:
|
|
113
|
+
worker_threads = 1
|
|
114
|
+
|
|
115
|
+
self._worker_threads = worker_threads
|
|
116
|
+
self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
|
|
117
|
+
try:
|
|
118
|
+
logging.debug(f"Using {worker_threads} threads in each worker")
|
|
119
|
+
self.worker_init_fn = self._mt_init_fn
|
|
120
|
+
yield
|
|
121
|
+
finally:
|
|
122
|
+
self.worker_init_fn = self._old_worker_init_fn
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class AffinityMixin:
|
|
126
|
+
r"""A context manager to enable CPU affinity for data loader workers."""
|
|
127
|
+
def _aff_init_fn(self, worker_id: int) -> None:
|
|
128
|
+
try:
|
|
129
|
+
worker_cores = self.loader_cores[worker_id]
|
|
130
|
+
if not isinstance(worker_cores, list):
|
|
131
|
+
worker_cores = [worker_cores]
|
|
132
|
+
|
|
133
|
+
if torch is not None and torch.multiprocessing.get_context()._name == 'spawn':
|
|
134
|
+
torch.set_num_threads(len(worker_cores))
|
|
135
|
+
|
|
136
|
+
if psutil is not None:
|
|
137
|
+
psutil.Process().cpu_affinity(worker_cores)
|
|
138
|
+
except IndexError as e:
|
|
139
|
+
raise ValueError(f"Cannot use CPU affinity for worker ID {worker_id} on CPU {self.loader_cores}") from e
|
|
140
|
+
|
|
141
|
+
self._old_worker_init_fn(worker_id)
|
|
142
|
+
|
|
143
|
+
@contextmanager
|
|
144
|
+
def enable_cpu_affinity(self, loader_cores: Optional[Union[List[List[int]], List[int]]] = None):
|
|
145
|
+
num_workers = getattr(self, 'num_workers', 0)
|
|
146
|
+
if not num_workers > 0:
|
|
147
|
+
raise ValueError(f"'enable_cpu_affinity' should be used with at least one worker (got {num_workers})")
|
|
148
|
+
if loader_cores and len(loader_cores) != num_workers:
|
|
149
|
+
raise ValueError(
|
|
150
|
+
f"The number of loader cores ({len(loader_cores)}) in 'enable_cpu_affinity' should match number of workers ({num_workers})"
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
from k3_node.data import HeteroData
|
|
154
|
+
if hasattr(self, 'data') and isinstance(self.data, HeteroData):
|
|
155
|
+
warnings.warn(
|
|
156
|
+
"Due to conflicting parallelization methods it is not advised to use affinitization with 'HeteroData' datasets.",
|
|
157
|
+
stacklevel=2,
|
|
158
|
+
)
|
|
159
|
+
|
|
160
|
+
self.loader_cores = loader_cores[:] if loader_cores else None
|
|
161
|
+
if self.loader_cores is None:
|
|
162
|
+
numa_info = get_numa_nodes_cores()
|
|
163
|
+
if numa_info and len(numa_info.get(0, [])) > num_workers:
|
|
164
|
+
node0_cores = [cpus[0] for core_id, cpus in numa_info[0]]
|
|
165
|
+
node0_cores.sort()
|
|
166
|
+
elif psutil is not None:
|
|
167
|
+
node0_cores = list(range(psutil.cpu_count(logical=False) or 1))
|
|
168
|
+
else:
|
|
169
|
+
node0_cores = list(range(os.cpu_count() or 1))
|
|
170
|
+
|
|
171
|
+
if len(node0_cores) < num_workers:
|
|
172
|
+
raise ValueError(f"More workers ({num_workers}) than available cores ({len(node0_cores)})")
|
|
173
|
+
|
|
174
|
+
if torch is not None and torch.multiprocessing.get_context()._name == 'spawn':
|
|
175
|
+
work_thread_pool = int(len(node0_cores) / num_workers)
|
|
176
|
+
self.loader_cores = [
|
|
177
|
+
list(range(work_thread_pool * i, work_thread_pool * (i + 1)))
|
|
178
|
+
for i in range(num_workers)
|
|
179
|
+
]
|
|
180
|
+
else:
|
|
181
|
+
self.loader_cores = node0_cores[:num_workers]
|
|
182
|
+
|
|
183
|
+
self._old_worker_init_fn = WorkerInitWrapper(getattr(self, 'worker_init_fn', None))
|
|
184
|
+
try:
|
|
185
|
+
logging.debug(f"{num_workers} data loader workers assigned to CPUs {self.loader_cores}")
|
|
186
|
+
self.worker_init_fn = self._aff_init_fn
|
|
187
|
+
yield
|
|
188
|
+
finally:
|
|
189
|
+
self.worker_init_fn = self._old_worker_init_fn
|
|
190
|
+
|
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
2
|
+
|
|
3
|
+
from k3_node.data import Data, HeteroData
|
|
4
|
+
from k3_node.loader.node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput, SamplerOutput
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from k3_node.loader.sampler_utils import (FastGraph, sample_neighbors_disjoint, sample_neighbors_hetero,
|
|
8
|
+
sample_neighbors_homo)
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class InternalNeighborSampler:
|
|
12
|
+
r"""Pure Python/NumPy neighborhood sampling engine."""
|
|
13
|
+
def __init__(
|
|
14
|
+
self,
|
|
15
|
+
data: Union[Data, HeteroData],
|
|
16
|
+
num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
|
|
17
|
+
replace: bool = False,
|
|
18
|
+
subgraph_type: str = 'directional',
|
|
19
|
+
disjoint: bool = False,
|
|
20
|
+
):
|
|
21
|
+
self.data = data
|
|
22
|
+
self.num_neighbors = num_neighbors
|
|
23
|
+
self.replace = replace
|
|
24
|
+
self.subgraph_type = subgraph_type
|
|
25
|
+
self.disjoint = disjoint
|
|
26
|
+
self.edge_permutation = None
|
|
27
|
+
self._graph = None # CSR index, built once and reused for every batch
|
|
28
|
+
|
|
29
|
+
def sample_from_nodes(self, input_data: NodeSamplerInput) -> Union[SamplerOutput, HeteroSamplerOutput]:
|
|
30
|
+
if isinstance(self.data, Data):
|
|
31
|
+
if self._graph is None:
|
|
32
|
+
self._graph = FastGraph(self.data.edge_index, num_nodes=self.data.num_nodes)
|
|
33
|
+
if self.disjoint: # one separate subgraph per seed node
|
|
34
|
+
seeds = input_data.node
|
|
35
|
+
is_torch = hasattr(seeds, "detach")
|
|
36
|
+
out = sample_neighbors_disjoint(self._graph, seeds.detach().cpu().numpy() if is_torch else seeds,
|
|
37
|
+
self.num_neighbors, replace=self.replace)
|
|
38
|
+
node, row, col, edge, batch, n_counts, e_counts = out
|
|
39
|
+
if is_torch:
|
|
40
|
+
import torch
|
|
41
|
+
node, row, col, edge, batch = (torch.from_numpy(a) for a in (node, row, col, edge, batch))
|
|
42
|
+
return SamplerOutput(node=node, row=row, col=col, edge=edge, batch=batch,
|
|
43
|
+
num_sampled_nodes=n_counts, num_sampled_edges=e_counts,
|
|
44
|
+
metadata=(input_data.input_id, input_data.time))
|
|
45
|
+
node, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
|
|
46
|
+
edge_index=self.data.edge_index,
|
|
47
|
+
seed_nodes=input_data.node,
|
|
48
|
+
num_neighbors=self.num_neighbors,
|
|
49
|
+
num_nodes=self.data.num_nodes,
|
|
50
|
+
replace=self.replace,
|
|
51
|
+
subgraph_type=self.subgraph_type,
|
|
52
|
+
disjoint=self.disjoint,
|
|
53
|
+
graph=self._graph,
|
|
54
|
+
)
|
|
55
|
+
return SamplerOutput(
|
|
56
|
+
node=node,
|
|
57
|
+
row=row,
|
|
58
|
+
col=col,
|
|
59
|
+
edge=edge,
|
|
60
|
+
num_sampled_nodes=n_counts,
|
|
61
|
+
num_sampled_edges=e_counts,
|
|
62
|
+
metadata=(input_data.input_id, input_data.time),
|
|
63
|
+
)
|
|
64
|
+
elif isinstance(self.data, HeteroData):
|
|
65
|
+
edge_index_dict = {}
|
|
66
|
+
for edge_type in self.data.edge_types:
|
|
67
|
+
canonical = self.data._to_canonical(*edge_type) if hasattr(self.data, '_to_canonical') else edge_type
|
|
68
|
+
edge_index_dict[canonical] = self.data[edge_type].edge_index
|
|
69
|
+
|
|
70
|
+
node_type = input_data.input_type or self.data.node_types[0]
|
|
71
|
+
seed_dict = {k: None for k in self.data.node_types}
|
|
72
|
+
seed_dict[node_type] = input_data.node
|
|
73
|
+
|
|
74
|
+
# Normalize num_neighbors for hetero
|
|
75
|
+
if isinstance(self.num_neighbors, dict):
|
|
76
|
+
norm_num_neighbors = {}
|
|
77
|
+
for k, v in self.num_neighbors.items():
|
|
78
|
+
can = self.data._to_canonical(*k) if hasattr(self.data, '_to_canonical') else k
|
|
79
|
+
norm_num_neighbors[can] = v
|
|
80
|
+
else:
|
|
81
|
+
norm_num_neighbors = self.num_neighbors
|
|
82
|
+
|
|
83
|
+
node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
|
|
84
|
+
edge_index_dict=edge_index_dict,
|
|
85
|
+
seed_nodes_dict=seed_dict,
|
|
86
|
+
num_neighbors=norm_num_neighbors,
|
|
87
|
+
replace=self.replace,
|
|
88
|
+
subgraph_type=self.subgraph_type,
|
|
89
|
+
)
|
|
90
|
+
return HeteroSamplerOutput(
|
|
91
|
+
node=node_dict,
|
|
92
|
+
row=row_dict,
|
|
93
|
+
col=col_dict,
|
|
94
|
+
edge=edge_dict,
|
|
95
|
+
num_sampled_nodes=n_counts,
|
|
96
|
+
num_sampled_edges=e_counts,
|
|
97
|
+
metadata=(input_data.input_id, input_data.time),
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
raise TypeError(f"Invalid data type for sampling: {type(self.data)}")
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class NeighborLoader(NodeLoader):
|
|
104
|
+
r"""A data loader that performs neighbor sampling as introduced in
|
|
105
|
+
"Inductive Representation Learning on Large Graphs".
|
|
106
|
+
|
|
107
|
+
Args:
|
|
108
|
+
data (Data or HeteroData): The graph data object.
|
|
109
|
+
num_neighbors (List[int] or Dict[EdgeType, List[int]]): Number of neighbors to sample per iteration.
|
|
110
|
+
input_nodes (Tensor or str or Tuple[str, Tensor], optional): Seed nodes. (default: :obj:`None`)
|
|
111
|
+
replace (bool, optional): Sample with replacement. (default: :obj:`False`)
|
|
112
|
+
subgraph_type (str, optional): :obj:`"directional"`, :obj:`"bidirectional"`, or :obj:`"induced"`.
|
|
113
|
+
(default: :obj:`"directional"`)
|
|
114
|
+
disjoint (bool, optional): If :obj:`True`, creates disjoint subgraphs per seed node. (default: :obj:`False`)
|
|
115
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
116
|
+
"""
|
|
117
|
+
def __init__(
|
|
118
|
+
self,
|
|
119
|
+
data: Union[Data, HeteroData],
|
|
120
|
+
num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
|
|
121
|
+
input_nodes: Any = None,
|
|
122
|
+
input_time: Optional[Any] = None,
|
|
123
|
+
replace: bool = False,
|
|
124
|
+
subgraph_type: str = 'directional',
|
|
125
|
+
disjoint: bool = False,
|
|
126
|
+
temporal_strategy: str = 'uniform',
|
|
127
|
+
time_attr: Optional[str] = None,
|
|
128
|
+
weight_attr: Optional[str] = None,
|
|
129
|
+
transform: Optional[Callable] = None,
|
|
130
|
+
transform_sampler_output: Optional[Callable] = None,
|
|
131
|
+
is_sorted: bool = False,
|
|
132
|
+
filter_per_worker: Optional[bool] = None,
|
|
133
|
+
neighbor_sampler: Optional[Any] = None,
|
|
134
|
+
directed: bool = True,
|
|
135
|
+
**kwargs,
|
|
136
|
+
):
|
|
137
|
+
if not directed:
|
|
138
|
+
subgraph_type = 'induced'
|
|
139
|
+
|
|
140
|
+
if neighbor_sampler is None:
|
|
141
|
+
neighbor_sampler = InternalNeighborSampler(
|
|
142
|
+
data,
|
|
143
|
+
num_neighbors=num_neighbors,
|
|
144
|
+
replace=replace,
|
|
145
|
+
subgraph_type=str(getattr(subgraph_type, 'value', subgraph_type)),
|
|
146
|
+
disjoint=disjoint,
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
super().__init__(
|
|
150
|
+
data=data,
|
|
151
|
+
node_sampler=neighbor_sampler,
|
|
152
|
+
input_nodes=input_nodes,
|
|
153
|
+
input_time=input_time,
|
|
154
|
+
transform=transform,
|
|
155
|
+
transform_sampler_output=transform_sampler_output,
|
|
156
|
+
filter_per_worker=filter_per_worker,
|
|
157
|
+
**kwargs,
|
|
158
|
+
)
|
|
159
|
+
|