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,334 @@
|
|
|
1
|
+
"""Makes graphs and graph loaders usable directly with ``model.fit`` / ``evaluate`` / ``predict``.
|
|
2
|
+
|
|
3
|
+
A model receives each batch as a :class:`GraphBatch`: a named tuple with one field per graph
|
|
4
|
+
attribute (``data.x``, ``data.edge_index``, ``data.batch``, ...), so its ``call`` reads like
|
|
5
|
+
PyG's ``forward``. The target (``y`` or ``edge_label``) and sample weights are passed to Keras
|
|
6
|
+
separately.
|
|
7
|
+
"""
|
|
8
|
+
import collections
|
|
9
|
+
from typing import Any, Dict, Optional, Sequence, Tuple
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
import keras
|
|
13
|
+
from keras import ops
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class _GraphBatchMixin:
|
|
17
|
+
__slots__ = ()
|
|
18
|
+
|
|
19
|
+
@property
|
|
20
|
+
def num_graphs(self) -> Optional[int]:
|
|
21
|
+
"""Number of graphs in the batch.
|
|
22
|
+
|
|
23
|
+
It is a Python ``int`` whenever the batch shape is known, including inside compiled
|
|
24
|
+
(``jax.jit`` / XLA) training steps, so it can be passed as ``size`` to the global pooling
|
|
25
|
+
functions: ``global_add_pool(x, data.batch, data.num_graphs)``.
|
|
26
|
+
"""
|
|
27
|
+
ptr = getattr(self, "ptr", None)
|
|
28
|
+
if ptr is not None:
|
|
29
|
+
n = ptr.shape[0]
|
|
30
|
+
return None if n is None else int(n) - 1
|
|
31
|
+
return 1 if getattr(self, "batch", None) is None else None
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def num_nodes(self) -> Optional[int]:
|
|
35
|
+
x = getattr(self, "x", None)
|
|
36
|
+
return None if x is None else x.shape[0]
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
_TYPES: Dict[Tuple[str, ...], type] = {}
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _graph_batch_type(fields: Tuple[str, ...]) -> type:
|
|
43
|
+
if fields not in _TYPES:
|
|
44
|
+
base = collections.namedtuple("GraphBatch", fields)
|
|
45
|
+
_TYPES[fields] = type("GraphBatch", (base, _GraphBatchMixin), {"__slots__": ()})
|
|
46
|
+
return _TYPES[fields]
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _as_array(value) -> Optional[np.ndarray]:
|
|
50
|
+
if isinstance(value, (bool, int, float, str)) or value is None:
|
|
51
|
+
return None
|
|
52
|
+
try:
|
|
53
|
+
array = np.asarray(ops.convert_to_numpy(value))
|
|
54
|
+
except Exception:
|
|
55
|
+
return None
|
|
56
|
+
if array.dtype.kind not in "biuf" or array.ndim == 0: # skip strings (e.g. SMILES) and objects
|
|
57
|
+
return None
|
|
58
|
+
if array.dtype == np.float64:
|
|
59
|
+
array = array.astype(np.float32)
|
|
60
|
+
elif array.dtype == np.int64:
|
|
61
|
+
array = array.astype(np.int32)
|
|
62
|
+
return array
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
HeteroGraphBatch = collections.namedtuple("HeteroGraphBatch", ["x_dict", "edge_index_dict"])
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _hetero_keras_batch(data, target, mask, normalize_mask, node_type):
|
|
69
|
+
x_dict = {k: _as_array(v) for k, v in data.collect("x").items()}
|
|
70
|
+
edge_index_dict = {k: _as_array(v) for k, v in data.collect("edge_index").items()}
|
|
71
|
+
inputs = HeteroGraphBatch(x_dict, edge_index_dict)
|
|
72
|
+
if node_type is None:
|
|
73
|
+
return (inputs,)
|
|
74
|
+
store = data[node_type]
|
|
75
|
+
y = _as_array(getattr(store, target or "y"))
|
|
76
|
+
if mask is None:
|
|
77
|
+
return inputs, y
|
|
78
|
+
weight = np.asarray(ops.convert_to_numpy(getattr(store, mask))).astype(np.float32)
|
|
79
|
+
if normalize_mask and weight.sum() > 0:
|
|
80
|
+
weight = weight * (weight.shape[0] / weight.sum())
|
|
81
|
+
return inputs, y, weight
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def to_keras_batch(
|
|
85
|
+
data: Any,
|
|
86
|
+
target: Optional[str] = None,
|
|
87
|
+
mask: Optional[str] = None,
|
|
88
|
+
normalize_mask: bool = True,
|
|
89
|
+
index: Optional[str] = None,
|
|
90
|
+
node_type: Optional[str] = None,
|
|
91
|
+
):
|
|
92
|
+
r"""Converts a :class:`~k3_node.data.Data` / :class:`~k3_node.data.Batch` into Keras' format.
|
|
93
|
+
|
|
94
|
+
For a :class:`~k3_node.data.HeteroData` graph, ``inputs`` is a :class:`HeteroGraphBatch`
|
|
95
|
+
(``x_dict`` and ``edge_index_dict``) and ``target`` / ``mask`` are read from ``node_type``.
|
|
96
|
+
|
|
97
|
+
Returns ``(inputs,)``, ``(inputs, y)`` or ``(inputs, y, sample_weight)``, where ``inputs`` is a
|
|
98
|
+
:class:`GraphBatch` holding every array attribute except the target and ``*_mask`` attributes.
|
|
99
|
+
|
|
100
|
+
Args:
|
|
101
|
+
data: The graph (or mini-batch of graphs).
|
|
102
|
+
target (str, optional): Attribute to predict. Defaults to ``"edge_label"`` if present
|
|
103
|
+
(link prediction), else ``"y"``.
|
|
104
|
+
mask (str, optional): Node mask attribute (e.g. ``"train_mask"``) used as sample weights.
|
|
105
|
+
normalize_mask (bool): Scale sample weights so the loss is the mean over the weighted
|
|
106
|
+
nodes (as in PyG), not Keras' sum divided by the number of all nodes.
|
|
107
|
+
index (str, optional): Node index attribute (e.g. ``"train_idx"``) for datasets that store
|
|
108
|
+
splits as node indices, with ``target`` holding the labels of exactly those nodes
|
|
109
|
+
(e.g. ``"train_y"``). The labels are placed at their nodes and only these nodes count.
|
|
110
|
+
"""
|
|
111
|
+
if hasattr(data, "node_types") and hasattr(data, "edge_types"):
|
|
112
|
+
return _hetero_keras_batch(data, target, mask, normalize_mask, node_type)
|
|
113
|
+
attrs = data.to_dict() if hasattr(data, "to_dict") else dict(vars(data))
|
|
114
|
+
if target is None:
|
|
115
|
+
target = "edge_label" if attrs.get("edge_label") is not None else "y"
|
|
116
|
+
|
|
117
|
+
fields = {}
|
|
118
|
+
for key in sorted(attrs):
|
|
119
|
+
if key == target or key.endswith(("_mask", "_idx")):
|
|
120
|
+
continue
|
|
121
|
+
array = _as_array(attrs[key])
|
|
122
|
+
if array is not None:
|
|
123
|
+
fields[key] = array
|
|
124
|
+
for key in ("batch", "ptr"): # Batch exposes these as properties
|
|
125
|
+
value = getattr(data, key, None)
|
|
126
|
+
if key not in fields and value is not None and _as_array(value) is not None:
|
|
127
|
+
fields[key] = _as_array(value)
|
|
128
|
+
if "batch" in fields and fields["batch"].shape[0] == 0:
|
|
129
|
+
# Graphs without nodes (e.g. batches of events): an empty `batch` would come first and
|
|
130
|
+
# make Keras weight the reported loss by 0.
|
|
131
|
+
fields.pop("batch")
|
|
132
|
+
fields.pop("ptr", None)
|
|
133
|
+
inputs = _graph_batch_type(tuple(fields))(**fields)
|
|
134
|
+
|
|
135
|
+
y = _as_array(attrs.get(target))
|
|
136
|
+
if y is None:
|
|
137
|
+
return (inputs,)
|
|
138
|
+
|
|
139
|
+
weight = None
|
|
140
|
+
if index is not None:
|
|
141
|
+
if attrs.get(index) is None:
|
|
142
|
+
raise ValueError(f"The graph has no index attribute '{index}'")
|
|
143
|
+
idx = np.asarray(ops.convert_to_numpy(attrs[index])).astype(np.int64)
|
|
144
|
+
num_nodes = attrs.get("num_nodes") or data.num_nodes
|
|
145
|
+
y_full = np.zeros((num_nodes,) + y.shape[1:], dtype=y.dtype)
|
|
146
|
+
y_full[idx] = y
|
|
147
|
+
y, weight = y_full, np.zeros(num_nodes, dtype=np.float32)
|
|
148
|
+
weight[idx] = 1.0
|
|
149
|
+
elif mask is not None:
|
|
150
|
+
if attrs.get(mask) is None:
|
|
151
|
+
raise ValueError(f"The graph has no mask attribute '{mask}'")
|
|
152
|
+
weight = np.asarray(ops.convert_to_numpy(attrs[mask])).astype(np.float32)
|
|
153
|
+
elif isinstance(attrs.get("batch_size"), int) and y.shape[0] == attrs.get("num_nodes", y.shape[0]):
|
|
154
|
+
# NeighborLoader: only the first `batch_size` (seed) nodes are supervised
|
|
155
|
+
weight = np.zeros(y.shape[0], dtype=np.float32)
|
|
156
|
+
weight[: attrs["batch_size"]] = 1.0
|
|
157
|
+
if weight is None:
|
|
158
|
+
return inputs, y
|
|
159
|
+
if normalize_mask and weight.sum() > 0:
|
|
160
|
+
weight = weight * (weight.shape[0] / weight.sum())
|
|
161
|
+
return inputs, y, weight
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
class KerasLoaderMixin(keras.utils.PyDataset):
|
|
165
|
+
r"""Lets a graph loader be passed straight to ``model.fit`` / ``evaluate`` / ``predict``.
|
|
166
|
+
|
|
167
|
+
Iterating the loader still yields :class:`~k3_node.data.Batch` objects; Keras instead reads
|
|
168
|
+
batches through ``__getitem__`` in the format produced by :func:`to_keras_batch`.
|
|
169
|
+
"""
|
|
170
|
+
|
|
171
|
+
# Class-level defaults stand in for PyDataset.__init__, which loader constructors don't call.
|
|
172
|
+
_workers = 1
|
|
173
|
+
_use_multiprocessing = False
|
|
174
|
+
_max_queue_size = 10
|
|
175
|
+
keras_target: Optional[str] = None
|
|
176
|
+
keras_mask: Optional[str] = None
|
|
177
|
+
|
|
178
|
+
def _keras_index_batches(self):
|
|
179
|
+
if getattr(self, "_keras_batches", None) is None:
|
|
180
|
+
batch_sampler = getattr(self, "batch_sampler", None)
|
|
181
|
+
if batch_sampler is not None and getattr(self, "batch_size", None) is not None:
|
|
182
|
+
self._keras_batches = list(iter(batch_sampler))
|
|
183
|
+
else:
|
|
184
|
+
self._keras_batches = False # no random access: stream from the iterator
|
|
185
|
+
return self._keras_batches
|
|
186
|
+
|
|
187
|
+
def __getitem__(self, index):
|
|
188
|
+
index_batches = self._keras_index_batches()
|
|
189
|
+
if index_batches:
|
|
190
|
+
batch = self.collate_fn([self.dataset[i] for i in index_batches[index]])
|
|
191
|
+
else:
|
|
192
|
+
if index == 0 or getattr(self, "_keras_iter", None) is None:
|
|
193
|
+
self._keras_iter = iter(self)
|
|
194
|
+
try:
|
|
195
|
+
batch = next(self._keras_iter)
|
|
196
|
+
except StopIteration:
|
|
197
|
+
self._keras_iter = iter(self)
|
|
198
|
+
batch = next(self._keras_iter)
|
|
199
|
+
return to_keras_batch(batch, target=self.keras_target, mask=self.keras_mask)
|
|
200
|
+
|
|
201
|
+
def with_mask(self, mask: str):
|
|
202
|
+
r"""Only the nodes in the node mask attribute ``mask`` (e.g. ``"train_mask"``) of each batch
|
|
203
|
+
count in the loss and in ``weighted_metrics``. Returns the loader, for chaining."""
|
|
204
|
+
self.keras_mask = mask
|
|
205
|
+
return self
|
|
206
|
+
|
|
207
|
+
def with_target(self, target: str):
|
|
208
|
+
r"""Sets the attribute Keras predicts (default: ``"edge_label"`` if present, else ``"y"``).
|
|
209
|
+
Returns the loader, for chaining."""
|
|
210
|
+
self.keras_target = target
|
|
211
|
+
return self
|
|
212
|
+
|
|
213
|
+
@property
|
|
214
|
+
def num_batches(self):
|
|
215
|
+
return len(self)
|
|
216
|
+
|
|
217
|
+
def on_epoch_end(self):
|
|
218
|
+
self._keras_batches = None # reshuffle next epoch
|
|
219
|
+
self._keras_iter = None
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _patch_tf_signature():
|
|
223
|
+
"""On TensorFlow, Keras fixes every tensor size that is the same in the first few batches it
|
|
224
|
+
reads (except the first axis). A sampling loader with fewer batches than that (e.g. a single
|
|
225
|
+
validation batch) returns different subgraphs on every call, so Keras would fix sizes that
|
|
226
|
+
change later. For K3-Node's loaders, the first batches are therefore read at least twice:
|
|
227
|
+
sizes that vary between calls become variable, and deterministic batches keep static shapes.
|
|
228
|
+
Falls back to Keras' behavior if its internals change."""
|
|
229
|
+
try:
|
|
230
|
+
from keras.src.trainers.data_adapters import data_adapter_utils
|
|
231
|
+
from keras.src.trainers.data_adapters.py_dataset_adapter import PyDatasetAdapter
|
|
232
|
+
except ImportError:
|
|
233
|
+
return
|
|
234
|
+
if getattr(PyDatasetAdapter, "_k3_node_patched", False):
|
|
235
|
+
return
|
|
236
|
+
original = PyDatasetAdapter.get_tf_dataset
|
|
237
|
+
|
|
238
|
+
def get_tf_dataset(self):
|
|
239
|
+
dataset = getattr(self, "py_dataset", None)
|
|
240
|
+
if getattr(self, "_output_signature", "missing") is None and isinstance(dataset, KerasLoaderMixin):
|
|
241
|
+
try:
|
|
242
|
+
num_samples = max(data_adapter_utils.NUM_BATCHES_FOR_TENSOR_SPEC, 2)
|
|
243
|
+
num_batches = dataset.num_batches or num_samples
|
|
244
|
+
# e.g. 3 samples of a 1-batch loader read batch 0 three times
|
|
245
|
+
batches = [self._standardize_batch(dataset[i % num_batches]) for i in range(num_samples)]
|
|
246
|
+
self._output_signature = data_adapter_utils.get_tensor_spec(batches)
|
|
247
|
+
except Exception:
|
|
248
|
+
self._output_signature = None
|
|
249
|
+
return original(self)
|
|
250
|
+
|
|
251
|
+
PyDatasetAdapter.get_tf_dataset = get_tf_dataset
|
|
252
|
+
PyDatasetAdapter._k3_node_patched = True
|
|
253
|
+
|
|
254
|
+
|
|
255
|
+
class FullGraphDataset(keras.utils.PyDataset):
|
|
256
|
+
r"""Feeds one whole graph to ``model.fit`` / ``evaluate`` / ``predict`` as a single batch.
|
|
257
|
+
|
|
258
|
+
Passing node arrays together with ``edge_index`` directly to ``fit`` fails, because Keras
|
|
259
|
+
slices every input along its first axis and ``edge_index`` has shape ``[2, num_edges]``.
|
|
260
|
+
This dataset yields the full graph as one batch per epoch instead. The model receives a
|
|
261
|
+
:class:`GraphBatch` (``data.x``, ``data.edge_index``, ...).
|
|
262
|
+
|
|
263
|
+
Args:
|
|
264
|
+
data: A :class:`~k3_node.data.Data` object.
|
|
265
|
+
mask (str, optional): Node mask attribute (e.g. ``"train_mask"``) whose nodes are used in
|
|
266
|
+
the loss and in ``weighted_metrics``. (default: :obj:`None`, all nodes)
|
|
267
|
+
target (str, optional): Attribute to predict. (default: ``"y"``)
|
|
268
|
+
index (str, optional): For splits stored as node indices: the index attribute (e.g.
|
|
269
|
+
``"train_idx"``), with ``target`` giving the labels of those nodes (e.g. ``"train_y"``).
|
|
270
|
+
neg_sampling_ratio (float, optional): For link prediction: adds this many random
|
|
271
|
+
non-edges per labeled edge (label 0), sampled anew every epoch.
|
|
272
|
+
node_type (str, optional): For a heterogeneous graph: the node type whose ``target`` is
|
|
273
|
+
predicted and whose ``mask`` selects the nodes. The model receives ``data.x_dict``
|
|
274
|
+
and ``data.edge_index_dict``.
|
|
275
|
+
|
|
276
|
+
Example:
|
|
277
|
+
```python
|
|
278
|
+
model.compile(optimizer="adam", loss=..., weighted_metrics=["accuracy"])
|
|
279
|
+
model.fit(FullGraphDataset(data, mask="train_mask"), epochs=200)
|
|
280
|
+
model.evaluate(FullGraphDataset(data, mask="test_mask"))
|
|
281
|
+
```
|
|
282
|
+
"""
|
|
283
|
+
|
|
284
|
+
def __init__(self, data, mask: Optional[str] = None, target: Optional[str] = None,
|
|
285
|
+
index: Optional[str] = None, neg_sampling_ratio: Optional[float] = None,
|
|
286
|
+
node_type: Optional[str] = None, **kwargs):
|
|
287
|
+
super().__init__(**kwargs)
|
|
288
|
+
self._data, self._neg_sampling_ratio = data, neg_sampling_ratio
|
|
289
|
+
self._kwargs = dict(target=target, mask=mask, index=index)
|
|
290
|
+
if node_type is not None:
|
|
291
|
+
self._kwargs["node_type"] = node_type
|
|
292
|
+
self._batch = None if neg_sampling_ratio else to_keras_batch(data, **self._kwargs)
|
|
293
|
+
|
|
294
|
+
def __len__(self):
|
|
295
|
+
return 1
|
|
296
|
+
|
|
297
|
+
def __getitem__(self, index):
|
|
298
|
+
if index != 0:
|
|
299
|
+
raise IndexError(index)
|
|
300
|
+
if self._neg_sampling_ratio: # fresh negative edges every epoch
|
|
301
|
+
return to_keras_batch(add_negative_edges(self._data, self._neg_sampling_ratio), **self._kwargs)
|
|
302
|
+
return self._batch
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
def add_negative_edges(data, ratio: float = 1.0):
|
|
306
|
+
r"""Returns a copy of a link prediction graph with random non-edges added to its labeled edges.
|
|
307
|
+
|
|
308
|
+
``ratio`` negatives are sampled per labeled edge in ``edge_label_index`` (or per edge of
|
|
309
|
+
``edge_index`` if there are no labeled edges), avoiding the edges of ``edge_index``. They are
|
|
310
|
+
appended to ``edge_label_index`` with label 0 in ``edge_label``.
|
|
311
|
+
"""
|
|
312
|
+
import copy
|
|
313
|
+
|
|
314
|
+
from k3_node.models.utils import negative_sampling
|
|
315
|
+
|
|
316
|
+
pos = getattr(data, "edge_label_index", None)
|
|
317
|
+
pos = np.asarray(ops.convert_to_numpy(data.edge_index if pos is None else pos))
|
|
318
|
+
label = getattr(data, "edge_label", None)
|
|
319
|
+
label = np.ones(pos.shape[1], np.float32) if label is None else np.asarray(ops.convert_to_numpy(label))
|
|
320
|
+
neg = np.asarray(ops.convert_to_numpy(negative_sampling(
|
|
321
|
+
data.edge_index, data.num_nodes, num_neg_samples=int(round(ratio * pos.shape[1])))))
|
|
322
|
+
out = copy.copy(data)
|
|
323
|
+
out.edge_label_index = np.concatenate([pos, neg.astype(pos.dtype)], axis=1)
|
|
324
|
+
out.edge_label = np.concatenate([label, np.zeros(neg.shape[1], label.dtype)])
|
|
325
|
+
return out
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def loader_bases(base: type) -> tuple:
|
|
329
|
+
"""Base classes for a graph loader: its data-loader base plus :class:`KerasLoaderMixin`."""
|
|
330
|
+
return (base, KerasLoaderMixin) if base is not object else (KerasLoaderMixin,)
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
if keras.config.backend() == "tensorflow":
|
|
334
|
+
_patch_tf_signature()
|
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
import torch
|
|
8
|
+
from torch import Tensor
|
|
9
|
+
except ImportError:
|
|
10
|
+
torch = None
|
|
11
|
+
Tensor = type(None)
|
|
12
|
+
|
|
13
|
+
from k3_node.data import Data, HeteroData
|
|
14
|
+
from k3_node.loader.base import BaseDataLoader, DataLoaderIterator
|
|
15
|
+
from k3_node.loader.mixin import AffinityMixin, LogMemoryMixin, MultithreadingMixin
|
|
16
|
+
from k3_node.loader.node_loader import HeteroSamplerOutput, SamplerOutput
|
|
17
|
+
from k3_node.loader.utils import (
|
|
18
|
+
filter_data,
|
|
19
|
+
filter_hetero_data,
|
|
20
|
+
get_edge_label_index,
|
|
21
|
+
infer_filter_per_worker,
|
|
22
|
+
)
|
|
23
|
+
from k3_node.loader.keras_dataset import loader_bases
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@dataclass
|
|
27
|
+
class EdgeSamplerInput:
|
|
28
|
+
input_id: Optional[Any]
|
|
29
|
+
row: Any
|
|
30
|
+
col: Any
|
|
31
|
+
label: Optional[Any] = None
|
|
32
|
+
time: Optional[Any] = None
|
|
33
|
+
input_type: Optional[Tuple[str, str, str]] = None
|
|
34
|
+
|
|
35
|
+
def __getitem__(self, index: Any) -> 'EdgeSamplerInput':
|
|
36
|
+
if torch is not None and not isinstance(index, Tensor):
|
|
37
|
+
index = torch.as_tensor(index, dtype=torch.long)
|
|
38
|
+
return EdgeSamplerInput(
|
|
39
|
+
input_id=self.input_id[index] if self.input_id is not None else index,
|
|
40
|
+
row=self.row[index],
|
|
41
|
+
col=self.col[index],
|
|
42
|
+
label=self.label[index] if self.label is not None else None,
|
|
43
|
+
time=self.time[index] if self.time is not None else None,
|
|
44
|
+
input_type=self.input_type,
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class LinkLoader(*loader_bases(BaseDataLoader), AffinityMixin, MultithreadingMixin, LogMemoryMixin):
|
|
49
|
+
r"""A data loader that performs mini-batch sampling from link information."""
|
|
50
|
+
def __init__(
|
|
51
|
+
self,
|
|
52
|
+
data: Union[Data, HeteroData],
|
|
53
|
+
link_sampler: Any,
|
|
54
|
+
edge_label_index: Any = None,
|
|
55
|
+
edge_label: Optional[Any] = None,
|
|
56
|
+
edge_label_time: Optional[Any] = None,
|
|
57
|
+
neg_sampling: Optional[Any] = None,
|
|
58
|
+
neg_sampling_ratio: Optional[Union[int, float]] = None,
|
|
59
|
+
transform: Optional[Callable] = None,
|
|
60
|
+
transform_sampler_output: Optional[Callable] = None,
|
|
61
|
+
filter_per_worker: Optional[bool] = None,
|
|
62
|
+
custom_cls: Optional[Any] = None,
|
|
63
|
+
input_id: Optional[Any] = None,
|
|
64
|
+
**kwargs,
|
|
65
|
+
):
|
|
66
|
+
if filter_per_worker is None:
|
|
67
|
+
filter_per_worker = infer_filter_per_worker(data)
|
|
68
|
+
|
|
69
|
+
kwargs.pop('dataset', None)
|
|
70
|
+
kwargs.pop('collate_fn', None)
|
|
71
|
+
|
|
72
|
+
input_type, edge_label_index = get_edge_label_index(data, edge_label_index)
|
|
73
|
+
|
|
74
|
+
self.data = data
|
|
75
|
+
self.link_sampler = link_sampler
|
|
76
|
+
self.neg_sampling = neg_sampling
|
|
77
|
+
self.neg_sampling_ratio = neg_sampling_ratio
|
|
78
|
+
self.transform = transform
|
|
79
|
+
self.transform_sampler_output = transform_sampler_output
|
|
80
|
+
self.filter_per_worker = filter_per_worker
|
|
81
|
+
self.custom_cls = custom_cls
|
|
82
|
+
|
|
83
|
+
if torch is not None and isinstance(edge_label_index, Tensor):
|
|
84
|
+
row = edge_label_index[0]
|
|
85
|
+
col = edge_label_index[1]
|
|
86
|
+
num_edges = edge_label_index.size(1)
|
|
87
|
+
else:
|
|
88
|
+
np_edges = np.asarray(edge_label_index)
|
|
89
|
+
row = np_edges[0]
|
|
90
|
+
col = np_edges[1]
|
|
91
|
+
num_edges = np_edges.shape[1]
|
|
92
|
+
|
|
93
|
+
self.input_data = EdgeSamplerInput(
|
|
94
|
+
input_id=input_id,
|
|
95
|
+
row=row,
|
|
96
|
+
col=col,
|
|
97
|
+
label=edge_label,
|
|
98
|
+
time=edge_label_time,
|
|
99
|
+
input_type=input_type,
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
iterator = range(num_edges)
|
|
103
|
+
|
|
104
|
+
if torch is not None:
|
|
105
|
+
super().__init__(iterator, collate_fn=self.collate_fn, **kwargs)
|
|
106
|
+
else:
|
|
107
|
+
self.dataset = iterator
|
|
108
|
+
self.collate_fn = self.collate_fn
|
|
109
|
+
|
|
110
|
+
def __call__(self, index: Any) -> Union[Data, HeteroData]:
|
|
111
|
+
out = self.collate_fn(index)
|
|
112
|
+
if not self.filter_per_worker:
|
|
113
|
+
out = self.filter_fn(out)
|
|
114
|
+
return out
|
|
115
|
+
|
|
116
|
+
def collate_fn(self, index: Any) -> Any:
|
|
117
|
+
input_data = self.input_data[index]
|
|
118
|
+
out = self.link_sampler.sample_from_edges(input_data)
|
|
119
|
+
if self.filter_per_worker:
|
|
120
|
+
out = self.filter_fn(out)
|
|
121
|
+
return out
|
|
122
|
+
|
|
123
|
+
def filter_fn(self, out: Any) -> Union[Data, HeteroData]:
|
|
124
|
+
if self.transform_sampler_output:
|
|
125
|
+
out = self.transform_sampler_output(out)
|
|
126
|
+
|
|
127
|
+
if isinstance(out, SamplerOutput):
|
|
128
|
+
perm = getattr(self.link_sampler, 'edge_permutation', None)
|
|
129
|
+
data = filter_data(self.data, out.node, out.row, out.col, out.edge, perm)
|
|
130
|
+
|
|
131
|
+
data.n_id = out.node
|
|
132
|
+
if out.edge is not None:
|
|
133
|
+
data.e_id = out.edge
|
|
134
|
+
data.batch = out.batch
|
|
135
|
+
data.num_sampled_nodes = out.num_sampled_nodes
|
|
136
|
+
data.num_sampled_edges = out.num_sampled_edges
|
|
137
|
+
|
|
138
|
+
meta = out.metadata or (None, None, None)
|
|
139
|
+
data.input_id = meta[0]
|
|
140
|
+
data.edge_label_index = meta[1]
|
|
141
|
+
data.edge_label = meta[2]
|
|
142
|
+
|
|
143
|
+
elif isinstance(out, HeteroSamplerOutput):
|
|
144
|
+
perm = getattr(self.link_sampler, 'edge_permutation', None)
|
|
145
|
+
data = filter_hetero_data(self.data, out.node, out.row, out.col, out.edge, perm)
|
|
146
|
+
|
|
147
|
+
for key, node in out.node.items():
|
|
148
|
+
data[key].n_id = node
|
|
149
|
+
|
|
150
|
+
for key, edge in (out.edge or {}).items():
|
|
151
|
+
if edge is not None:
|
|
152
|
+
data[key].e_id = edge
|
|
153
|
+
|
|
154
|
+
if out.batch is not None:
|
|
155
|
+
data.set_value_dict('batch', out.batch)
|
|
156
|
+
if out.num_sampled_nodes is not None:
|
|
157
|
+
data.set_value_dict('num_sampled_nodes', out.num_sampled_nodes)
|
|
158
|
+
if out.num_sampled_edges is not None:
|
|
159
|
+
data.set_value_dict('num_sampled_edges', out.num_sampled_edges)
|
|
160
|
+
|
|
161
|
+
input_type = self.input_data.input_type
|
|
162
|
+
meta = out.metadata or (None, None, None)
|
|
163
|
+
if input_type is not None:
|
|
164
|
+
data[input_type].input_id = meta[0]
|
|
165
|
+
data[input_type].edge_label_index = meta[1]
|
|
166
|
+
data[input_type].edge_label = meta[2]
|
|
167
|
+
else:
|
|
168
|
+
data = out
|
|
169
|
+
|
|
170
|
+
return data if self.transform is None else self.transform(data)
|
|
171
|
+
|
|
172
|
+
def _get_iterator(self) -> Iterator:
|
|
173
|
+
if self.filter_per_worker:
|
|
174
|
+
return super()._get_iterator()
|
|
175
|
+
return DataLoaderIterator(super()._get_iterator(), self.filter_fn)
|
|
176
|
+
|
|
177
|
+
def __repr__(self) -> str:
|
|
178
|
+
return f'{self.__class__.__name__}()'
|
|
179
|
+
|