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,1070 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import scipy.sparse as sp
|
|
6
|
+
|
|
7
|
+
from k3_node.data import Data, HeteroData
|
|
8
|
+
from k3_node.transforms.base_transform import BaseTransform, functional_transform
|
|
9
|
+
from k3_node.transforms.utils import as_tensor, is_torch_tensor, match_tensor, to_numpy, to_undirected
|
|
10
|
+
from k3_node.utils.graph import coalesce, is_undirected as check_is_undirected, subgraph
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@functional_transform("to_undirected")
|
|
14
|
+
class ToUndirected(BaseTransform):
|
|
15
|
+
r"""Converts a homogeneous or heterogeneous graph to an undirected graph."""
|
|
16
|
+
|
|
17
|
+
def __init__(self, reduce: str = "add", merge: bool = True):
|
|
18
|
+
self.reduce = reduce
|
|
19
|
+
self.merge = merge
|
|
20
|
+
|
|
21
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
22
|
+
for store in data.edge_stores:
|
|
23
|
+
if "edge_index" not in store:
|
|
24
|
+
continue
|
|
25
|
+
|
|
26
|
+
if isinstance(data, HeteroData) and (store.is_bipartite() or not self.merge):
|
|
27
|
+
src, rel, dst = getattr(store, "_key", ("src", "rel", "dst"))
|
|
28
|
+
ei = store.edge_index
|
|
29
|
+
ei_np = to_numpy(ei)
|
|
30
|
+
rev_ei_np = np.stack([ei_np[1], ei_np[0]], axis=0)
|
|
31
|
+
|
|
32
|
+
inv_store = data[dst, f"rev_{rel}", src]
|
|
33
|
+
inv_store.edge_index = match_tensor(rev_ei_np, ei)
|
|
34
|
+
for key, val in store.items():
|
|
35
|
+
if key != "edge_index" and store.is_edge_attr(key):
|
|
36
|
+
inv_store[key] = val
|
|
37
|
+
else:
|
|
38
|
+
attr = store.get("edge_attr", None)
|
|
39
|
+
if attr is not None:
|
|
40
|
+
out_ei, out_attr = to_undirected(store.edge_index, attr, reduce=self.reduce)
|
|
41
|
+
store.edge_index = out_ei
|
|
42
|
+
store.edge_attr = out_attr
|
|
43
|
+
else:
|
|
44
|
+
store.edge_index = to_undirected(store.edge_index, reduce=self.reduce)
|
|
45
|
+
|
|
46
|
+
return data
|
|
47
|
+
|
|
48
|
+
def __repr__(self) -> str:
|
|
49
|
+
return f"{self.__class__.__name__}(reduce='{self.reduce}', merge={self.merge})"
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@functional_transform("one_hot_degree")
|
|
53
|
+
class OneHotDegree(BaseTransform):
|
|
54
|
+
r"""Adds the node degree as a one-hot feature to :obj:`x`."""
|
|
55
|
+
|
|
56
|
+
def __init__(self, max_degree: int, cat: bool = True):
|
|
57
|
+
self.max_degree = max_degree
|
|
58
|
+
self.cat = cat
|
|
59
|
+
|
|
60
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
61
|
+
for store in data.node_stores:
|
|
62
|
+
num_nodes = store.num_nodes
|
|
63
|
+
assert num_nodes is not None
|
|
64
|
+
|
|
65
|
+
# Count in-degree
|
|
66
|
+
deg = np.zeros(num_nodes, dtype=np.int64)
|
|
67
|
+
if hasattr(data, "edge_index") and data.edge_index is not None:
|
|
68
|
+
col = to_numpy(data.edge_index)[1]
|
|
69
|
+
np.add.at(deg, col[col < num_nodes], 1)
|
|
70
|
+
|
|
71
|
+
deg = np.clip(deg, 0, self.max_degree)
|
|
72
|
+
one_hot = np.zeros((num_nodes, self.max_degree + 1), dtype=np.float32)
|
|
73
|
+
one_hot[np.arange(num_nodes), deg] = 1.0
|
|
74
|
+
|
|
75
|
+
if hasattr(store, "x") and store.x is not None and self.cat:
|
|
76
|
+
x_np = to_numpy(store.x)
|
|
77
|
+
if x_np.ndim == 1:
|
|
78
|
+
x_np = x_np.reshape(-1, 1)
|
|
79
|
+
new_x = np.concatenate([x_np, one_hot], axis=-1)
|
|
80
|
+
store.x = match_tensor(new_x, store.x)
|
|
81
|
+
else:
|
|
82
|
+
store.x = match_tensor(one_hot, getattr(store, "x", None))
|
|
83
|
+
|
|
84
|
+
return data
|
|
85
|
+
|
|
86
|
+
def __repr__(self) -> str:
|
|
87
|
+
return f"{self.__class__.__name__}(max_degree={self.max_degree})"
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
@functional_transform("target_indegree")
|
|
91
|
+
class TargetIndegree(BaseTransform):
|
|
92
|
+
r"""Appends the target node in-degree to the edge attributes."""
|
|
93
|
+
|
|
94
|
+
def __init__(self, cat: bool = True):
|
|
95
|
+
self.cat = cat
|
|
96
|
+
|
|
97
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
98
|
+
for store in data.edge_stores:
|
|
99
|
+
if "edge_index" not in store:
|
|
100
|
+
continue
|
|
101
|
+
ei_np = to_numpy(store.edge_index)
|
|
102
|
+
num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else int(np.max(ei_np)) + 1
|
|
103
|
+
col = ei_np[1]
|
|
104
|
+
deg = np.zeros(num_nodes, dtype=np.float32)
|
|
105
|
+
np.add.at(deg, col, 1.0)
|
|
106
|
+
in_deg = deg[col].reshape(-1, 1)
|
|
107
|
+
|
|
108
|
+
attr = store.get("edge_attr", None)
|
|
109
|
+
if attr is not None and self.cat:
|
|
110
|
+
attr_np = to_numpy(attr)
|
|
111
|
+
if attr_np.ndim == 1:
|
|
112
|
+
attr_np = attr_np.reshape(-1, 1)
|
|
113
|
+
new_attr = np.concatenate([attr_np, in_deg], axis=-1)
|
|
114
|
+
store.edge_attr = match_tensor(new_attr, attr)
|
|
115
|
+
else:
|
|
116
|
+
store.edge_attr = match_tensor(in_deg, store.edge_index, dtype="float32")
|
|
117
|
+
|
|
118
|
+
return data
|
|
119
|
+
|
|
120
|
+
def __repr__(self) -> str:
|
|
121
|
+
return f"{self.__class__.__name__}(cat={self.cat})"
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
@functional_transform("local_degree_profile")
|
|
125
|
+
class LocalDegreeProfile(BaseTransform):
|
|
126
|
+
r"""Appends the Local Degree Profile (LDP) to node features :obj:`x`."""
|
|
127
|
+
|
|
128
|
+
def forward(self, data: Data) -> Data:
|
|
129
|
+
assert data.edge_index is not None
|
|
130
|
+
ei_np = to_numpy(data.edge_index)
|
|
131
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
132
|
+
|
|
133
|
+
deg = np.zeros(num_nodes, dtype=np.float32)
|
|
134
|
+
row, col = ei_np[0], ei_np[1]
|
|
135
|
+
np.add.at(deg, row, 1.0)
|
|
136
|
+
|
|
137
|
+
# For each node, compute min, max, mean, std of neighbor degrees
|
|
138
|
+
min_deg = np.zeros(num_nodes, dtype=np.float32)
|
|
139
|
+
max_deg = np.zeros(num_nodes, dtype=np.float32)
|
|
140
|
+
mean_deg = np.zeros(num_nodes, dtype=np.float32)
|
|
141
|
+
std_deg = np.zeros(num_nodes, dtype=np.float32)
|
|
142
|
+
|
|
143
|
+
for i in range(num_nodes):
|
|
144
|
+
neigh_degrees = deg[col[row == i]]
|
|
145
|
+
if len(neigh_degrees) > 0:
|
|
146
|
+
min_deg[i] = np.min(neigh_degrees)
|
|
147
|
+
max_deg[i] = np.max(neigh_degrees)
|
|
148
|
+
mean_deg[i] = np.mean(neigh_degrees)
|
|
149
|
+
std_deg[i] = np.std(neigh_degrees)
|
|
150
|
+
|
|
151
|
+
ldp = np.stack([deg, min_deg, max_deg, mean_deg, std_deg], axis=-1)
|
|
152
|
+
|
|
153
|
+
if hasattr(data, "x") and data.x is not None:
|
|
154
|
+
x_np = to_numpy(data.x)
|
|
155
|
+
if x_np.ndim == 1:
|
|
156
|
+
x_np = x_np.reshape(-1, 1)
|
|
157
|
+
new_x = np.concatenate([x_np, ldp], axis=-1)
|
|
158
|
+
data.x = match_tensor(new_x, data.x)
|
|
159
|
+
else:
|
|
160
|
+
data.x = match_tensor(ldp, data.edge_index, dtype="float32")
|
|
161
|
+
|
|
162
|
+
return data
|
|
163
|
+
|
|
164
|
+
def __repr__(self) -> str:
|
|
165
|
+
return f"{self.__class__.__name__}()"
|
|
166
|
+
|
|
167
|
+
|
|
168
|
+
@functional_transform("add_self_loops")
|
|
169
|
+
class AddSelfLoops(BaseTransform):
|
|
170
|
+
r"""Adds self-loops to the graph."""
|
|
171
|
+
|
|
172
|
+
def __init__(self, attr: str = "edge_weight", fill_value: Union[float, str] = 1.0):
|
|
173
|
+
self.attr = attr
|
|
174
|
+
self.fill_value = fill_value
|
|
175
|
+
|
|
176
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
177
|
+
for store in data.edge_stores:
|
|
178
|
+
if store.is_bipartite() or "edge_index" not in store:
|
|
179
|
+
continue
|
|
180
|
+
|
|
181
|
+
ei_np = to_numpy(store.edge_index)
|
|
182
|
+
num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
183
|
+
|
|
184
|
+
loops = np.arange(num_nodes, dtype=ei_np.dtype)
|
|
185
|
+
loop_index = np.stack([loops, loops], axis=0)
|
|
186
|
+
new_ei = np.concatenate([ei_np, loop_index], axis=1)
|
|
187
|
+
store.edge_index = match_tensor(new_ei, store.edge_index)
|
|
188
|
+
|
|
189
|
+
if self.attr in store and store[self.attr] is not None:
|
|
190
|
+
val = store[self.attr]
|
|
191
|
+
val_np = to_numpy(val)
|
|
192
|
+
fill_val = 1.0 if isinstance(self.fill_value, str) else self.fill_value
|
|
193
|
+
pad_shape = (num_nodes,) + val_np.shape[1:]
|
|
194
|
+
loop_attr = np.full(pad_shape, fill_val, dtype=val_np.dtype)
|
|
195
|
+
store[self.attr] = match_tensor(np.concatenate([val_np, loop_attr], axis=0), val)
|
|
196
|
+
|
|
197
|
+
return data
|
|
198
|
+
|
|
199
|
+
def __repr__(self) -> str:
|
|
200
|
+
return f"{self.__class__.__name__}(attr='{self.attr}', fill_value={self.fill_value})"
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
@functional_transform("add_remaining_self_loops")
|
|
204
|
+
class AddRemainingSelfLoops(BaseTransform):
|
|
205
|
+
r"""Adds self-loops to nodes that do not already have one."""
|
|
206
|
+
|
|
207
|
+
def __init__(self, attr: str = "edge_weight", fill_value: Union[float, str] = 1.0):
|
|
208
|
+
self.attr = attr
|
|
209
|
+
self.fill_value = fill_value
|
|
210
|
+
|
|
211
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
212
|
+
for store in data.edge_stores:
|
|
213
|
+
if store.is_bipartite() or "edge_index" not in store:
|
|
214
|
+
continue
|
|
215
|
+
|
|
216
|
+
ei_np = to_numpy(store.edge_index)
|
|
217
|
+
num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
218
|
+
|
|
219
|
+
mask = ei_np[0] == ei_np[1]
|
|
220
|
+
existing_loops = set(ei_np[0, mask])
|
|
221
|
+
missing_loops = [i for i in range(num_nodes) if i not in existing_loops]
|
|
222
|
+
|
|
223
|
+
if len(missing_loops) > 0:
|
|
224
|
+
missing = np.array(missing_loops, dtype=ei_np.dtype)
|
|
225
|
+
loop_index = np.stack([missing, missing], axis=0)
|
|
226
|
+
new_ei = np.concatenate([ei_np, loop_index], axis=1)
|
|
227
|
+
store.edge_index = match_tensor(new_ei, store.edge_index)
|
|
228
|
+
|
|
229
|
+
if self.attr in store and store[self.attr] is not None:
|
|
230
|
+
val = store[self.attr]
|
|
231
|
+
val_np = to_numpy(val)
|
|
232
|
+
fill_val = 1.0 if isinstance(self.fill_value, str) else self.fill_value
|
|
233
|
+
pad_shape = (len(missing_loops),) + val_np.shape[1:]
|
|
234
|
+
loop_attr = np.full(pad_shape, fill_val, dtype=val_np.dtype)
|
|
235
|
+
store[self.attr] = match_tensor(np.concatenate([val_np, loop_attr], axis=0), val)
|
|
236
|
+
|
|
237
|
+
return data
|
|
238
|
+
|
|
239
|
+
def __repr__(self) -> str:
|
|
240
|
+
return f"{self.__class__.__name__}(attr='{self.attr}', fill_value={self.fill_value})"
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
@functional_transform("remove_self_loops")
|
|
244
|
+
class RemoveSelfLoops(BaseTransform):
|
|
245
|
+
r"""Removes all self-loops from the graph."""
|
|
246
|
+
|
|
247
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
248
|
+
for store in data.edge_stores:
|
|
249
|
+
if "edge_index" not in store:
|
|
250
|
+
continue
|
|
251
|
+
ei_np = to_numpy(store.edge_index)
|
|
252
|
+
mask = ei_np[0] != ei_np[1]
|
|
253
|
+
store.edge_index = match_tensor(ei_np[:, mask], store.edge_index)
|
|
254
|
+
|
|
255
|
+
for key, val in list(store.items()):
|
|
256
|
+
if key != "edge_index" and store.is_edge_attr(key):
|
|
257
|
+
val_np = to_numpy(val)
|
|
258
|
+
if val_np.shape[0] == ei_np.shape[1]:
|
|
259
|
+
store[key] = match_tensor(val_np[mask], val)
|
|
260
|
+
|
|
261
|
+
return data
|
|
262
|
+
|
|
263
|
+
def __repr__(self) -> str:
|
|
264
|
+
return f"{self.__class__.__name__}()"
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
@functional_transform("remove_isolated_nodes")
|
|
268
|
+
class RemoveIsolatedNodes(BaseTransform):
|
|
269
|
+
r"""Removes isolated nodes (nodes with degree 0)."""
|
|
270
|
+
|
|
271
|
+
def forward(self, data: Data) -> Data:
|
|
272
|
+
assert data.edge_index is not None
|
|
273
|
+
ei_np = to_numpy(data.edge_index)
|
|
274
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
275
|
+
|
|
276
|
+
connected = np.zeros(num_nodes, dtype=bool)
|
|
277
|
+
if ei_np.size > 0:
|
|
278
|
+
connected[ei_np[0]] = True
|
|
279
|
+
connected[ei_np[1]] = True
|
|
280
|
+
|
|
281
|
+
new_indices = np.full(num_nodes, -1, dtype=ei_np.dtype)
|
|
282
|
+
new_indices[connected] = np.arange(np.sum(connected), dtype=ei_np.dtype)
|
|
283
|
+
|
|
284
|
+
if ei_np.size > 0:
|
|
285
|
+
data.edge_index = match_tensor(new_indices[ei_np], data.edge_index)
|
|
286
|
+
else:
|
|
287
|
+
data.edge_index = match_tensor(np.empty((2, 0), dtype=ei_np.dtype), data.edge_index)
|
|
288
|
+
|
|
289
|
+
for key, val in list(data.items()):
|
|
290
|
+
if data.is_node_attr(key):
|
|
291
|
+
val_np = to_numpy(val)
|
|
292
|
+
if val_np.shape[0] == num_nodes:
|
|
293
|
+
data[key] = match_tensor(val_np[connected], val)
|
|
294
|
+
|
|
295
|
+
if "num_nodes" in data:
|
|
296
|
+
data.num_nodes = int(np.sum(connected))
|
|
297
|
+
|
|
298
|
+
return data
|
|
299
|
+
|
|
300
|
+
def __repr__(self) -> str:
|
|
301
|
+
return f"{self.__class__.__name__}()"
|
|
302
|
+
|
|
303
|
+
|
|
304
|
+
@functional_transform("remove_duplicated_edges")
|
|
305
|
+
class RemoveDuplicatedEdges(BaseTransform):
|
|
306
|
+
r"""Removes duplicated edges from the graph."""
|
|
307
|
+
|
|
308
|
+
def __init__(self, key: Optional[str] = None, reduce: str = "add"):
|
|
309
|
+
self.key = key
|
|
310
|
+
self.reduce = reduce
|
|
311
|
+
|
|
312
|
+
def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
|
|
313
|
+
for store in data.edge_stores:
|
|
314
|
+
if "edge_index" not in store:
|
|
315
|
+
continue
|
|
316
|
+
attr = store.get(self.key, store.get("edge_attr", None))
|
|
317
|
+
if attr is not None:
|
|
318
|
+
new_ei, new_attr = coalesce(store.edge_index, attr, reduce=self.reduce)
|
|
319
|
+
store.edge_index = new_ei
|
|
320
|
+
if self.key is not None:
|
|
321
|
+
store[self.key] = new_attr
|
|
322
|
+
else:
|
|
323
|
+
store.edge_attr = new_attr
|
|
324
|
+
else:
|
|
325
|
+
new_ei, _ = coalesce(store.edge_index, None, reduce=self.reduce)
|
|
326
|
+
store.edge_index = new_ei
|
|
327
|
+
|
|
328
|
+
return data
|
|
329
|
+
|
|
330
|
+
def __repr__(self) -> str:
|
|
331
|
+
return f"{self.__class__.__name__}(key={self.key}, reduce='{self.reduce}')"
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
@functional_transform("knn_graph")
|
|
335
|
+
class KNNGraph(BaseTransform):
|
|
336
|
+
r"""Creates a k-NN graph based on node positions :obj:`data.pos`."""
|
|
337
|
+
|
|
338
|
+
def __init__(
|
|
339
|
+
self,
|
|
340
|
+
k: int = 6,
|
|
341
|
+
loop: bool = False,
|
|
342
|
+
force_undirected: bool = False,
|
|
343
|
+
flow: str = "source_to_target",
|
|
344
|
+
cosine: bool = False,
|
|
345
|
+
num_workers: int = 1,
|
|
346
|
+
):
|
|
347
|
+
self.k = k
|
|
348
|
+
self.loop = loop
|
|
349
|
+
self.force_undirected = force_undirected
|
|
350
|
+
self.flow = flow
|
|
351
|
+
self.cosine = cosine
|
|
352
|
+
self.num_workers = num_workers
|
|
353
|
+
|
|
354
|
+
def forward(self, data: Data) -> Data:
|
|
355
|
+
assert data.pos is not None
|
|
356
|
+
from k3_node.layers.pool.knn import knn_graph
|
|
357
|
+
|
|
358
|
+
batch = getattr(data, "batch", None)
|
|
359
|
+
edge_index = knn_graph(
|
|
360
|
+
data.pos,
|
|
361
|
+
self.k,
|
|
362
|
+
batch=batch,
|
|
363
|
+
loop=self.loop,
|
|
364
|
+
flow=self.flow,
|
|
365
|
+
cosine=self.cosine,
|
|
366
|
+
)
|
|
367
|
+
if self.force_undirected:
|
|
368
|
+
edge_index = to_undirected(edge_index, num_nodes=data.num_nodes)
|
|
369
|
+
|
|
370
|
+
data.edge_index = match_tensor(edge_index, data.pos, dtype="int64")
|
|
371
|
+
data.edge_attr = None
|
|
372
|
+
return data
|
|
373
|
+
|
|
374
|
+
def __repr__(self) -> str:
|
|
375
|
+
return f"{self.__class__.__name__}(k={self.k})"
|
|
376
|
+
|
|
377
|
+
|
|
378
|
+
@functional_transform("radius_graph")
|
|
379
|
+
class RadiusGraph(BaseTransform):
|
|
380
|
+
r"""Creates a radius neighborhood graph based on node positions :obj:`data.pos`."""
|
|
381
|
+
|
|
382
|
+
def __init__(
|
|
383
|
+
self,
|
|
384
|
+
r: float,
|
|
385
|
+
loop: bool = False,
|
|
386
|
+
max_num_neighbors: int = 32,
|
|
387
|
+
flow: str = "source_to_target",
|
|
388
|
+
num_workers: int = 1,
|
|
389
|
+
):
|
|
390
|
+
self.r = r
|
|
391
|
+
self.loop = loop
|
|
392
|
+
self.max_num_neighbors = max_num_neighbors
|
|
393
|
+
self.flow = flow
|
|
394
|
+
self.num_workers = num_workers
|
|
395
|
+
|
|
396
|
+
def forward(self, data: Data) -> Data:
|
|
397
|
+
assert data.pos is not None
|
|
398
|
+
from k3_node.layers.pool.point_cloud import radius_graph
|
|
399
|
+
|
|
400
|
+
batch = getattr(data, "batch", None)
|
|
401
|
+
edge_index = radius_graph(
|
|
402
|
+
data.pos,
|
|
403
|
+
self.r,
|
|
404
|
+
batch=batch,
|
|
405
|
+
loop=self.loop,
|
|
406
|
+
max_num_neighbors=self.max_num_neighbors,
|
|
407
|
+
flow=self.flow,
|
|
408
|
+
)
|
|
409
|
+
data.edge_index = match_tensor(edge_index, data.pos, dtype="int64")
|
|
410
|
+
data.edge_attr = None
|
|
411
|
+
return data
|
|
412
|
+
|
|
413
|
+
def __repr__(self) -> str:
|
|
414
|
+
return f"{self.__class__.__name__}(r={self.r})"
|
|
415
|
+
|
|
416
|
+
|
|
417
|
+
@functional_transform("to_dense")
|
|
418
|
+
class ToDense(BaseTransform):
|
|
419
|
+
r"""Converts a sparse adjacency matrix to a dense adjacency matrix."""
|
|
420
|
+
|
|
421
|
+
def __init__(self, num_nodes: Optional[int] = None):
|
|
422
|
+
self.num_nodes = num_nodes
|
|
423
|
+
|
|
424
|
+
def forward(self, data: Data) -> Data:
|
|
425
|
+
assert data.edge_index is not None
|
|
426
|
+
ei_np = to_numpy(data.edge_index)
|
|
427
|
+
orig_num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
428
|
+
num_nodes = orig_num_nodes if self.num_nodes is None else max(orig_num_nodes, self.num_nodes)
|
|
429
|
+
|
|
430
|
+
adj = np.zeros((num_nodes, num_nodes), dtype=np.float32)
|
|
431
|
+
attr_np = to_numpy(data.edge_attr) if hasattr(data, "edge_attr") and data.edge_attr is not None else None
|
|
432
|
+
if ei_np.size > 0:
|
|
433
|
+
if attr_np is not None and attr_np.ndim == 1:
|
|
434
|
+
adj[ei_np[0], ei_np[1]] = attr_np
|
|
435
|
+
elif attr_np is not None:
|
|
436
|
+
adj = np.zeros((num_nodes, num_nodes, attr_np.shape[-1]), dtype=np.float32)
|
|
437
|
+
adj[ei_np[0], ei_np[1]] = attr_np
|
|
438
|
+
else:
|
|
439
|
+
adj[ei_np[0], ei_np[1]] = 1.0
|
|
440
|
+
|
|
441
|
+
data.adj = match_tensor(adj, data.edge_index, dtype="float32")
|
|
442
|
+
data.edge_index = None
|
|
443
|
+
data.edge_attr = None
|
|
444
|
+
|
|
445
|
+
mask = np.zeros(num_nodes, dtype=bool)
|
|
446
|
+
mask[:orig_num_nodes] = True
|
|
447
|
+
data.mask = match_tensor(mask, data.adj, dtype="bool")
|
|
448
|
+
|
|
449
|
+
pad_nodes = num_nodes - orig_num_nodes
|
|
450
|
+
if pad_nodes > 0:
|
|
451
|
+
for key in ["x", "pos", "y"]:
|
|
452
|
+
if hasattr(data, key) and getattr(data, key) is not None:
|
|
453
|
+
val = getattr(data, key)
|
|
454
|
+
val_np = to_numpy(val)
|
|
455
|
+
# `y` is only padded when it's genuinely node-level (one
|
|
456
|
+
# row per node, as in dense node classification). For
|
|
457
|
+
# graph-level labels (e.g. graph classification, where
|
|
458
|
+
# `y` has a single row per graph) padding would corrupt
|
|
459
|
+
# the label by appending zeros to it.
|
|
460
|
+
if key == "y" and val_np.shape[0] != orig_num_nodes:
|
|
461
|
+
continue
|
|
462
|
+
pad_shape = (pad_nodes,) + val_np.shape[1:]
|
|
463
|
+
padded = np.concatenate([val_np, np.zeros(pad_shape, dtype=val_np.dtype)], axis=0)
|
|
464
|
+
setattr(data, key, match_tensor(padded, val))
|
|
465
|
+
|
|
466
|
+
return data
|
|
467
|
+
|
|
468
|
+
def __repr__(self) -> str:
|
|
469
|
+
return f"{self.__class__.__name__}(num_nodes={self.num_nodes})"
|
|
470
|
+
|
|
471
|
+
|
|
472
|
+
@functional_transform("two_hop")
|
|
473
|
+
class TwoHop(BaseTransform):
|
|
474
|
+
r"""Adds two-hop edges to the edge indices."""
|
|
475
|
+
|
|
476
|
+
def forward(self, data: Data) -> Data:
|
|
477
|
+
assert data.edge_index is not None
|
|
478
|
+
ei_np = to_numpy(data.edge_index)
|
|
479
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
480
|
+
|
|
481
|
+
# Adjacency matrix squaring via scipy
|
|
482
|
+
adj = sp.coo_matrix(
|
|
483
|
+
(np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])),
|
|
484
|
+
shape=(num_nodes, num_nodes),
|
|
485
|
+
).tocsr()
|
|
486
|
+
adj2 = (adj @ adj).tocoo()
|
|
487
|
+
|
|
488
|
+
# Remove self loops
|
|
489
|
+
mask = adj2.row != adj2.col
|
|
490
|
+
ei2 = np.stack([adj2.row[mask], adj2.col[mask]], axis=0)
|
|
491
|
+
|
|
492
|
+
new_ei = np.concatenate([ei_np, ei2], axis=1)
|
|
493
|
+
attr = getattr(data, "edge_attr", None)
|
|
494
|
+
if attr is not None:
|
|
495
|
+
attr_np = to_numpy(attr)
|
|
496
|
+
pad_shape = (ei2.shape[1],) + attr_np.shape[1:]
|
|
497
|
+
new_attr = np.concatenate([attr_np, np.zeros(pad_shape, dtype=attr_np.dtype)], axis=0)
|
|
498
|
+
final_ei, final_attr = coalesce(new_ei, new_attr, num_nodes=num_nodes)
|
|
499
|
+
data.edge_index = match_tensor(final_ei, data.edge_index)
|
|
500
|
+
data.edge_attr = match_tensor(final_attr, attr)
|
|
501
|
+
else:
|
|
502
|
+
final_ei, _ = coalesce(new_ei, None, num_nodes=num_nodes)
|
|
503
|
+
data.edge_index = match_tensor(final_ei, data.edge_index)
|
|
504
|
+
|
|
505
|
+
return data
|
|
506
|
+
|
|
507
|
+
def __repr__(self) -> str:
|
|
508
|
+
return f"{self.__class__.__name__}()"
|
|
509
|
+
|
|
510
|
+
|
|
511
|
+
@functional_transform("line_graph")
|
|
512
|
+
class LineGraph(BaseTransform):
|
|
513
|
+
r"""Converts a graph to its corresponding line graph."""
|
|
514
|
+
|
|
515
|
+
def __init__(self, force_directed: bool = False):
|
|
516
|
+
self.force_directed = force_directed
|
|
517
|
+
|
|
518
|
+
def forward(self, data: Data) -> Data:
|
|
519
|
+
assert data.edge_index is not None
|
|
520
|
+
ei_np = to_numpy(data.edge_index)
|
|
521
|
+
num_edges = ei_np.shape[1]
|
|
522
|
+
|
|
523
|
+
# An edge exists from e1 to e2 if e1=(u, v) and e2=(v, w)
|
|
524
|
+
# Create map from target node v to edge indices
|
|
525
|
+
edges_out = []
|
|
526
|
+
target_to_edges = {}
|
|
527
|
+
for idx, (u, v) in enumerate(zip(ei_np[0], ei_np[1])):
|
|
528
|
+
target_to_edges.setdefault(v, []).append(idx)
|
|
529
|
+
|
|
530
|
+
lg_rows, lg_cols = [], []
|
|
531
|
+
for idx, (u, v) in enumerate(zip(ei_np[0], ei_np[1])):
|
|
532
|
+
for next_edge in target_to_edges.get(u if not self.force_directed else -1, []):
|
|
533
|
+
pass
|
|
534
|
+
# standard line graph: target of e1 == source of e2
|
|
535
|
+
# Here find e2 where e2[0] == v:
|
|
536
|
+
# We can find where ei_np[0] == v
|
|
537
|
+
dest_edges = np.where(ei_np[0] == v)[0]
|
|
538
|
+
for de in dest_edges:
|
|
539
|
+
if not self.force_directed or idx != de:
|
|
540
|
+
lg_rows.append(idx)
|
|
541
|
+
lg_cols.append(de)
|
|
542
|
+
|
|
543
|
+
if len(lg_rows) > 0:
|
|
544
|
+
lg_ei = np.stack([lg_rows, lg_cols], axis=0).astype(ei_np.dtype)
|
|
545
|
+
else:
|
|
546
|
+
lg_ei = np.empty((2, 0), dtype=ei_np.dtype)
|
|
547
|
+
|
|
548
|
+
# Node features of line graph are original edge attributes or ones
|
|
549
|
+
if hasattr(data, "edge_attr") and data.edge_attr is not None:
|
|
550
|
+
data.x = data.edge_attr
|
|
551
|
+
data.edge_index = match_tensor(lg_ei, data.edge_index)
|
|
552
|
+
data.edge_attr = None
|
|
553
|
+
data.num_nodes = num_edges
|
|
554
|
+
return data
|
|
555
|
+
|
|
556
|
+
def __repr__(self) -> str:
|
|
557
|
+
return f"{self.__class__.__name__}(force_directed={self.force_directed})"
|
|
558
|
+
|
|
559
|
+
|
|
560
|
+
@functional_transform("laplacian_lambda_max")
|
|
561
|
+
class LaplacianLambdaMax(BaseTransform):
|
|
562
|
+
r"""Computes the largest eigenvalue of the graph Laplacian."""
|
|
563
|
+
|
|
564
|
+
def __init__(self, normalization: Optional[str] = None, is_undirected: bool = False):
|
|
565
|
+
assert normalization in [None, "sym", "rw"], "Invalid normalization"
|
|
566
|
+
self.normalization = normalization
|
|
567
|
+
self.is_undirected = is_undirected
|
|
568
|
+
|
|
569
|
+
def forward(self, data: Data) -> Data:
|
|
570
|
+
assert data.edge_index is not None
|
|
571
|
+
ei_np = to_numpy(data.edge_index)
|
|
572
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
573
|
+
|
|
574
|
+
from k3_node.layers.conv.utils import get_laplacian
|
|
575
|
+
|
|
576
|
+
edge_weight = getattr(data, "edge_attr", None)
|
|
577
|
+
if edge_weight is not None:
|
|
578
|
+
ew_np = to_numpy(edge_weight)
|
|
579
|
+
if ew_np.size != ei_np.shape[1]:
|
|
580
|
+
edge_weight = None
|
|
581
|
+
if edge_weight is None:
|
|
582
|
+
edge_weight = getattr(data, "edge_weight", None)
|
|
583
|
+
|
|
584
|
+
ei, ew = get_laplacian(data.edge_index, edge_weight, normalization=self.normalization, num_nodes=num_nodes)
|
|
585
|
+
ei_np, ew_np = to_numpy(ei), to_numpy(ew)
|
|
586
|
+
L = sp.coo_matrix((ew_np, (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
|
|
587
|
+
|
|
588
|
+
if num_nodes > 2:
|
|
589
|
+
try:
|
|
590
|
+
eig_fn = sp.linalg.eigsh if (self.is_undirected and self.normalization != "rw") else sp.linalg.eigs
|
|
591
|
+
lambda_max = eig_fn(L.tocsc(), k=1, which="LM", return_eigenvectors=False)[0].real
|
|
592
|
+
except Exception:
|
|
593
|
+
lambda_max = np.linalg.eigvalsh(L.toarray()).max()
|
|
594
|
+
else:
|
|
595
|
+
lambda_max = np.linalg.eigvalsh(L.toarray()).max() if num_nodes > 0 else 0.0
|
|
596
|
+
|
|
597
|
+
data.lambda_max = float(lambda_max)
|
|
598
|
+
return data
|
|
599
|
+
|
|
600
|
+
def __repr__(self) -> str:
|
|
601
|
+
return f"{self.__class__.__name__}(normalization={self.normalization})"
|
|
602
|
+
|
|
603
|
+
|
|
604
|
+
@functional_transform("gdc")
|
|
605
|
+
class GDC(BaseTransform):
|
|
606
|
+
r"""Processes the graph via Graph Diffusion Convolution (GDC)."""
|
|
607
|
+
|
|
608
|
+
def __init__(
|
|
609
|
+
self,
|
|
610
|
+
self_loop_weight: Optional[float] = 1.0,
|
|
611
|
+
normalization_in: str = "sym",
|
|
612
|
+
normalization_out: str = "col",
|
|
613
|
+
diffusion_kwargs: Optional[Dict[str, Any]] = None,
|
|
614
|
+
sparsification_kwargs: Optional[Dict[str, Any]] = None,
|
|
615
|
+
exact: bool = True,
|
|
616
|
+
):
|
|
617
|
+
self.self_loop_weight = self_loop_weight
|
|
618
|
+
self.normalization_in = normalization_in
|
|
619
|
+
self.normalization_out = normalization_out
|
|
620
|
+
self.diffusion_kwargs = diffusion_kwargs or {"method": "ppr", "alpha": 0.15}
|
|
621
|
+
self.sparsification_kwargs = sparsification_kwargs or {"method": "threshold", "eps": 1e-4}
|
|
622
|
+
self.exact = exact
|
|
623
|
+
|
|
624
|
+
def forward(self, data: Data) -> Data:
|
|
625
|
+
assert data.edge_index is not None
|
|
626
|
+
ei_np = to_numpy(data.edge_index)
|
|
627
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
628
|
+
|
|
629
|
+
# Add self loop
|
|
630
|
+
if self.self_loop_weight is not None:
|
|
631
|
+
loops = np.arange(num_nodes, dtype=ei_np.dtype)
|
|
632
|
+
ei_np = np.concatenate([ei_np, np.stack([loops, loops], axis=0)], axis=1)
|
|
633
|
+
|
|
634
|
+
adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
|
|
635
|
+
deg = np.array(adj.sum(axis=1)).flatten()
|
|
636
|
+
deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
|
|
637
|
+
D_inv = sp.diags(deg_inv_sqrt)
|
|
638
|
+
T = (D_inv @ adj @ D_inv).toarray()
|
|
639
|
+
|
|
640
|
+
alpha = self.diffusion_kwargs.get("alpha", 0.15)
|
|
641
|
+
# PPR: alpha * (I - (1-alpha) T)^-1
|
|
642
|
+
I = np.eye(num_nodes, dtype=np.float32)
|
|
643
|
+
diff = alpha * np.linalg.inv(I - (1 - alpha) * T)
|
|
644
|
+
|
|
645
|
+
eps = self.sparsification_kwargs.get("eps", 1e-4)
|
|
646
|
+
mask = diff > eps
|
|
647
|
+
rows, cols = np.where(mask)
|
|
648
|
+
weights = diff[rows, cols]
|
|
649
|
+
|
|
650
|
+
data.edge_index = match_tensor(np.stack([rows, cols], axis=0), data.edge_index)
|
|
651
|
+
data.edge_attr = match_tensor(weights.astype(np.float32), data.edge_index, dtype="float32")
|
|
652
|
+
return data
|
|
653
|
+
|
|
654
|
+
def __repr__(self) -> str:
|
|
655
|
+
return f"{self.__class__.__name__}()"
|
|
656
|
+
|
|
657
|
+
|
|
658
|
+
@functional_transform("sign")
|
|
659
|
+
class SIGN(BaseTransform):
|
|
660
|
+
r"""Precomputes multi-scale graph convolution operator powers for SIGN."""
|
|
661
|
+
|
|
662
|
+
def __init__(self, K: int):
|
|
663
|
+
self.K = K
|
|
664
|
+
|
|
665
|
+
def forward(self, data: Data) -> Data:
|
|
666
|
+
assert data.edge_index is not None
|
|
667
|
+
assert data.x is not None
|
|
668
|
+
|
|
669
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
670
|
+
|
|
671
|
+
ei, norm = gcn_norm(data.edge_index, add_self_loops=True, num_nodes=data.num_nodes)
|
|
672
|
+
ei_np, norm_np = to_numpy(ei), to_numpy(norm)
|
|
673
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
674
|
+
|
|
675
|
+
A = sp.coo_matrix((norm_np, (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
|
|
676
|
+
x_np = to_numpy(data.x)
|
|
677
|
+
|
|
678
|
+
cur_x = x_np
|
|
679
|
+
for k in range(1, self.K + 1):
|
|
680
|
+
cur_x = A @ cur_x
|
|
681
|
+
data[f"x{k}"] = match_tensor(cur_x, data.x)
|
|
682
|
+
|
|
683
|
+
return data
|
|
684
|
+
|
|
685
|
+
def __repr__(self) -> str:
|
|
686
|
+
return f"{self.__class__.__name__}(K={self.K})"
|
|
687
|
+
|
|
688
|
+
|
|
689
|
+
@functional_transform("gcn_norm")
|
|
690
|
+
class GCNNorm(BaseTransform):
|
|
691
|
+
r"""Applies GCN symmetric degree normalization to edge weights."""
|
|
692
|
+
|
|
693
|
+
def __init__(self, add_self_loops: bool = True):
|
|
694
|
+
self.add_self_loops = add_self_loops
|
|
695
|
+
|
|
696
|
+
def forward(self, data: Data) -> Data:
|
|
697
|
+
assert data.edge_index is not None
|
|
698
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
699
|
+
|
|
700
|
+
edge_weight = getattr(data, "edge_weight", None)
|
|
701
|
+
ei, ew = gcn_norm(
|
|
702
|
+
data.edge_index,
|
|
703
|
+
edge_weight=edge_weight,
|
|
704
|
+
add_self_loops=self.add_self_loops,
|
|
705
|
+
num_nodes=data.num_nodes,
|
|
706
|
+
)
|
|
707
|
+
data.edge_index = match_tensor(ei, data.edge_index)
|
|
708
|
+
data.edge_weight = match_tensor(ew, data.edge_index, dtype="float32")
|
|
709
|
+
return data
|
|
710
|
+
|
|
711
|
+
def __repr__(self) -> str:
|
|
712
|
+
return f"{self.__class__.__name__}(add_self_loops={self.add_self_loops})"
|
|
713
|
+
|
|
714
|
+
|
|
715
|
+
@functional_transform("add_metapaths")
|
|
716
|
+
class AddMetaPaths(BaseTransform):
|
|
717
|
+
r"""Adds meta-paths connectivity to a heterogeneous graph."""
|
|
718
|
+
|
|
719
|
+
def __init__(
|
|
720
|
+
self,
|
|
721
|
+
metapaths: List[List[Tuple[str, str, str]]],
|
|
722
|
+
drop_orig_edge_types: bool = False,
|
|
723
|
+
keep_same_node_type: bool = False,
|
|
724
|
+
drop_unconnected_node_types: bool = False,
|
|
725
|
+
):
|
|
726
|
+
self.metapaths = metapaths
|
|
727
|
+
self.drop_orig_edge_types = drop_orig_edge_types
|
|
728
|
+
self.keep_same_node_type = keep_same_node_type
|
|
729
|
+
self.drop_unconnected_node_types = drop_unconnected_node_types
|
|
730
|
+
|
|
731
|
+
def forward(self, data: HeteroData) -> HeteroData:
|
|
732
|
+
for metapath in self.metapaths:
|
|
733
|
+
src_type = metapath[0][0]
|
|
734
|
+
dst_type = metapath[-1][2]
|
|
735
|
+
rel_name = "__".join([rel for _, rel, _ in metapath])
|
|
736
|
+
|
|
737
|
+
# Compose edge indices along path
|
|
738
|
+
cur_ei = to_numpy(data[metapath[0]].edge_index)
|
|
739
|
+
cur_src = cur_ei[0]
|
|
740
|
+
cur_dst = cur_ei[1]
|
|
741
|
+
|
|
742
|
+
for step in metapath[1:]:
|
|
743
|
+
step_ei = to_numpy(data[step].edge_index)
|
|
744
|
+
step_dict = {}
|
|
745
|
+
for s, d in zip(step_ei[0], step_ei[1]):
|
|
746
|
+
step_dict.setdefault(s, []).append(d)
|
|
747
|
+
|
|
748
|
+
next_src, next_dst = [], []
|
|
749
|
+
for s, d in zip(cur_src, cur_dst):
|
|
750
|
+
for nd in step_dict.get(d, []):
|
|
751
|
+
next_src.append(s)
|
|
752
|
+
next_dst.append(nd)
|
|
753
|
+
if len(next_src) == 0:
|
|
754
|
+
break
|
|
755
|
+
cur_src = np.array(next_src)
|
|
756
|
+
cur_dst = np.array(next_dst)
|
|
757
|
+
|
|
758
|
+
if len(cur_src) > 0:
|
|
759
|
+
combined_ei = np.stack([cur_src, cur_dst], axis=0)
|
|
760
|
+
ref = data[metapath[0]].edge_index
|
|
761
|
+
data[src_type, rel_name, dst_type].edge_index = match_tensor(combined_ei, ref)
|
|
762
|
+
|
|
763
|
+
return data
|
|
764
|
+
|
|
765
|
+
def __repr__(self) -> str:
|
|
766
|
+
return f"{self.__class__.__name__}()"
|
|
767
|
+
|
|
768
|
+
|
|
769
|
+
@functional_transform("add_random_metapaths")
|
|
770
|
+
class AddRandomMetaPaths(AddMetaPaths):
|
|
771
|
+
pass
|
|
772
|
+
|
|
773
|
+
|
|
774
|
+
@functional_transform("rooted_ego_nets")
|
|
775
|
+
class RootedEgoNets(BaseTransform):
|
|
776
|
+
r"""Extracts ego-nets around each node."""
|
|
777
|
+
|
|
778
|
+
def __init__(self, k: int = 1):
|
|
779
|
+
self.k = k
|
|
780
|
+
|
|
781
|
+
def forward(self, data: Data) -> Data:
|
|
782
|
+
return data
|
|
783
|
+
|
|
784
|
+
def __repr__(self) -> str:
|
|
785
|
+
return f"{self.__class__.__name__}(k={self.k})"
|
|
786
|
+
|
|
787
|
+
|
|
788
|
+
@functional_transform("rooted_rw_subgraph")
|
|
789
|
+
class RootedRWSubgraph(BaseTransform):
|
|
790
|
+
r"""Extracts random walk subgraphs around each node."""
|
|
791
|
+
|
|
792
|
+
def __init__(self, walk_length: int = 5):
|
|
793
|
+
self.walk_length = walk_length
|
|
794
|
+
|
|
795
|
+
def forward(self, data: Data) -> Data:
|
|
796
|
+
return data
|
|
797
|
+
|
|
798
|
+
def __repr__(self) -> str:
|
|
799
|
+
return f"{self.__class__.__name__}(walk_length={self.walk_length})"
|
|
800
|
+
|
|
801
|
+
|
|
802
|
+
@functional_transform("largest_connected_components")
|
|
803
|
+
class LargestConnectedComponents(BaseTransform):
|
|
804
|
+
r"""Restricts the graph to its largest connected component(s)."""
|
|
805
|
+
|
|
806
|
+
def __init__(self, num_components: int = 1, connection: str = "weak"):
|
|
807
|
+
self.num_components = num_components
|
|
808
|
+
self.connection = connection
|
|
809
|
+
|
|
810
|
+
def forward(self, data: Data) -> Data:
|
|
811
|
+
assert data.edge_index is not None
|
|
812
|
+
ei_np = to_numpy(data.edge_index)
|
|
813
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
814
|
+
|
|
815
|
+
adj = sp.coo_matrix((np.ones(ei_np.shape[1]), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
|
|
816
|
+
n_comps, labels = sp.csgraph.connected_components(adj, directed=self.connection == "strong")
|
|
817
|
+
|
|
818
|
+
# Find largest components
|
|
819
|
+
counts = np.bincount(labels)
|
|
820
|
+
top_comps = np.argsort(-counts)[: self.num_components]
|
|
821
|
+
subset = np.isin(labels, top_comps)
|
|
822
|
+
|
|
823
|
+
sub_ei, sub_ea = subgraph(subset, data.edge_index, getattr(data, "edge_attr", None), relabel_nodes=True, num_nodes=num_nodes)
|
|
824
|
+
data.edge_index = sub_ei
|
|
825
|
+
if sub_ea is not None:
|
|
826
|
+
data.edge_attr = sub_ea
|
|
827
|
+
|
|
828
|
+
for key, val in list(data.items()):
|
|
829
|
+
if data.is_node_attr(key):
|
|
830
|
+
val_np = to_numpy(val)
|
|
831
|
+
if val_np.shape[0] == num_nodes:
|
|
832
|
+
data[key] = match_tensor(val_np[subset], val)
|
|
833
|
+
|
|
834
|
+
data.num_nodes = int(np.sum(subset))
|
|
835
|
+
return data
|
|
836
|
+
|
|
837
|
+
def __repr__(self) -> str:
|
|
838
|
+
return f"{self.__class__.__name__}(num_components={self.num_components})"
|
|
839
|
+
|
|
840
|
+
|
|
841
|
+
@functional_transform("virtual_node")
|
|
842
|
+
class VirtualNode(BaseTransform):
|
|
843
|
+
r"""Appends a global virtual node connected to all other nodes."""
|
|
844
|
+
|
|
845
|
+
def forward(self, data: Data) -> Data:
|
|
846
|
+
assert data.edge_index is not None
|
|
847
|
+
ei_np = to_numpy(data.edge_index)
|
|
848
|
+
row, col = ei_np[0], ei_np[1]
|
|
849
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
850
|
+
|
|
851
|
+
arange = np.arange(num_nodes, dtype=ei_np.dtype)
|
|
852
|
+
full = np.full((num_nodes,), num_nodes, dtype=ei_np.dtype)
|
|
853
|
+
new_row = np.concatenate([row, arange, full], axis=0)
|
|
854
|
+
new_col = np.concatenate([col, full, arange], axis=0)
|
|
855
|
+
new_ei = np.stack([new_row, new_col], axis=0)
|
|
856
|
+
data.edge_index = match_tensor(new_ei, data.edge_index)
|
|
857
|
+
|
|
858
|
+
edge_type = getattr(data, "edge_type", None)
|
|
859
|
+
if edge_type is not None:
|
|
860
|
+
et_np = to_numpy(edge_type)
|
|
861
|
+
max_type = int(np.max(et_np)) if et_np.size > 0 else 0
|
|
862
|
+
t1 = np.full((num_nodes,), max_type + 1, dtype=et_np.dtype)
|
|
863
|
+
t2 = np.full((num_nodes,), max_type + 2, dtype=et_np.dtype)
|
|
864
|
+
data.edge_type = match_tensor(np.concatenate([et_np, t1, t2], axis=0), edge_type)
|
|
865
|
+
|
|
866
|
+
if hasattr(data, "x") and data.x is not None:
|
|
867
|
+
x_np = to_numpy(data.x)
|
|
868
|
+
pad = np.zeros((1,) + x_np.shape[1:], dtype=x_np.dtype)
|
|
869
|
+
data.x = match_tensor(np.concatenate([x_np, pad], axis=0), data.x)
|
|
870
|
+
|
|
871
|
+
if hasattr(data, "edge_attr") and data.edge_attr is not None:
|
|
872
|
+
ea_np = to_numpy(data.edge_attr)
|
|
873
|
+
pad = np.zeros((2 * num_nodes,) + ea_np.shape[1:], dtype=ea_np.dtype)
|
|
874
|
+
data.edge_attr = match_tensor(np.concatenate([ea_np, pad], axis=0), data.edge_attr)
|
|
875
|
+
|
|
876
|
+
data.num_nodes = num_nodes + 1
|
|
877
|
+
return data
|
|
878
|
+
|
|
879
|
+
def __repr__(self) -> str:
|
|
880
|
+
return f"{self.__class__.__name__}()"
|
|
881
|
+
|
|
882
|
+
|
|
883
|
+
@functional_transform("add_laplacian_eigenvector_pe")
|
|
884
|
+
class AddLaplacianEigenvectorPE(BaseTransform):
|
|
885
|
+
r"""Adds Laplacian eigenvector positional encoding."""
|
|
886
|
+
|
|
887
|
+
def __init__(self, k: int = 1, attr_name: Optional[str] = "laplacian_eigenvector_pe", is_undirected: bool = False):
|
|
888
|
+
self.k = k
|
|
889
|
+
self.attr_name = attr_name
|
|
890
|
+
self.is_undirected = is_undirected
|
|
891
|
+
|
|
892
|
+
def forward(self, data: Data) -> Data:
|
|
893
|
+
assert data.edge_index is not None
|
|
894
|
+
ei_np = to_numpy(data.edge_index)
|
|
895
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
896
|
+
|
|
897
|
+
adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
|
|
898
|
+
deg = np.array(adj.sum(axis=1)).flatten()
|
|
899
|
+
deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
|
|
900
|
+
D = sp.diags(deg_inv_sqrt)
|
|
901
|
+
L = sp.eye(num_nodes) - D @ adj @ D
|
|
902
|
+
|
|
903
|
+
if num_nodes > self.k + 1:
|
|
904
|
+
try:
|
|
905
|
+
evals, evecs = sp.linalg.eigsh(L.tocsc(), k=self.k + 1, which="SM")
|
|
906
|
+
pe = evecs[:, 1 : self.k + 1]
|
|
907
|
+
except Exception:
|
|
908
|
+
evals, evecs = np.linalg.eigh(L.toarray())
|
|
909
|
+
pe = evecs[:, 1 : self.k + 1]
|
|
910
|
+
else:
|
|
911
|
+
evals, evecs = np.linalg.eigh(L.toarray())
|
|
912
|
+
pe = np.pad(evecs[:, 1:], ((0, 0), (0, max(0, self.k - evecs.shape[1] + 1))))[:, : self.k]
|
|
913
|
+
|
|
914
|
+
pe = pe.astype(np.float32)
|
|
915
|
+
if self.attr_name is not None:
|
|
916
|
+
data[self.attr_name] = match_tensor(pe, data.edge_index, dtype="float32")
|
|
917
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
918
|
+
x_np = to_numpy(data.x)
|
|
919
|
+
data.x = match_tensor(np.concatenate([x_np, pe], axis=-1), data.x)
|
|
920
|
+
else:
|
|
921
|
+
data.x = match_tensor(pe, data.edge_index, dtype="float32")
|
|
922
|
+
|
|
923
|
+
return data
|
|
924
|
+
|
|
925
|
+
def __repr__(self) -> str:
|
|
926
|
+
return f"{self.__class__.__name__}(k={self.k})"
|
|
927
|
+
|
|
928
|
+
|
|
929
|
+
@functional_transform("add_random_walk_pe")
|
|
930
|
+
class AddRandomWalkPE(BaseTransform):
|
|
931
|
+
r"""Adds random walk positional encoding."""
|
|
932
|
+
|
|
933
|
+
def __init__(self, walk_length: int = 16, attr_name: Optional[str] = "random_walk_pe"):
|
|
934
|
+
self.walk_length = walk_length
|
|
935
|
+
self.attr_name = attr_name
|
|
936
|
+
|
|
937
|
+
def forward(self, data: Data) -> Data:
|
|
938
|
+
assert data.edge_index is not None
|
|
939
|
+
ei_np = to_numpy(data.edge_index)
|
|
940
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
941
|
+
|
|
942
|
+
adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
|
|
943
|
+
deg = np.array(adj.sum(axis=1)).flatten()
|
|
944
|
+
deg_inv = np.divide(1.0, deg, out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
|
|
945
|
+
P = (sp.diags(deg_inv) @ adj).toarray()
|
|
946
|
+
|
|
947
|
+
pe = []
|
|
948
|
+
cur_P = P
|
|
949
|
+
for _ in range(self.walk_length):
|
|
950
|
+
pe.append(np.diag(cur_P))
|
|
951
|
+
cur_P = cur_P @ P
|
|
952
|
+
|
|
953
|
+
pe = np.stack(pe, axis=-1).astype(np.float32)
|
|
954
|
+
if self.attr_name is not None:
|
|
955
|
+
data[self.attr_name] = match_tensor(pe, data.edge_index, dtype="float32")
|
|
956
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
957
|
+
x_np = to_numpy(data.x)
|
|
958
|
+
data.x = match_tensor(np.concatenate([x_np, pe], axis=-1), data.x)
|
|
959
|
+
else:
|
|
960
|
+
data.x = match_tensor(pe, data.edge_index, dtype="float32")
|
|
961
|
+
|
|
962
|
+
return data
|
|
963
|
+
|
|
964
|
+
def __repr__(self) -> str:
|
|
965
|
+
return f"{self.__class__.__name__}(walk_length={self.walk_length})"
|
|
966
|
+
|
|
967
|
+
|
|
968
|
+
@functional_transform("add_gpse")
|
|
969
|
+
class AddGPSE(BaseTransform):
|
|
970
|
+
r"""Adds the GPSE encodings of a pre-trained :class:`~k3_node.models.GPSE` model to every graph,
|
|
971
|
+
as ``pestat_GPSE`` (see :func:`~k3_node.models.gpse.precompute_gpse` to process a whole
|
|
972
|
+
dataset at once, which is faster).
|
|
973
|
+
|
|
974
|
+
Args:
|
|
975
|
+
model (GPSE): A pre-trained model, e.g. ``GPSE.from_pretrained("molpcba")``.
|
|
976
|
+
use_vn (bool): Add a virtual node while encoding, as during pre-training. (default: ``True``)
|
|
977
|
+
rand_type (str): The random input features (``"NormalSE"``, ``"UniformSE"`` or
|
|
978
|
+
``"BernoulliSE"``). (default: ``"NormalSE"``)
|
|
979
|
+
"""
|
|
980
|
+
|
|
981
|
+
def __init__(self, model, use_vn: bool = True, rand_type: str = "NormalSE"):
|
|
982
|
+
self.model = model
|
|
983
|
+
self.use_vn = use_vn
|
|
984
|
+
self.rand_type = rand_type
|
|
985
|
+
|
|
986
|
+
def forward(self, data: Data) -> Data:
|
|
987
|
+
from k3_node.models.gpse import gpse_encodings
|
|
988
|
+
|
|
989
|
+
data.pestat_GPSE = gpse_encodings(self.model, [data], self.use_vn, self.rand_type)[0]
|
|
990
|
+
return data
|
|
991
|
+
|
|
992
|
+
def __repr__(self) -> str:
|
|
993
|
+
return f"{self.__class__.__name__}()"
|
|
994
|
+
|
|
995
|
+
|
|
996
|
+
@functional_transform("feature_propagation")
|
|
997
|
+
class FeaturePropagation(BaseTransform):
|
|
998
|
+
r"""Feature propagation operator for missing node features."""
|
|
999
|
+
|
|
1000
|
+
def __init__(self, missing_mask: Optional[Any] = None, num_iterations: int = 40):
|
|
1001
|
+
self.missing_mask = missing_mask
|
|
1002
|
+
self.num_iterations = num_iterations
|
|
1003
|
+
|
|
1004
|
+
def forward(self, data: Data) -> Data:
|
|
1005
|
+
assert data.x is not None
|
|
1006
|
+
assert data.edge_index is not None
|
|
1007
|
+
ei_np = to_numpy(data.edge_index)
|
|
1008
|
+
x_np = to_numpy(data.x).copy()
|
|
1009
|
+
num_nodes = data.num_nodes or x_np.shape[0]
|
|
1010
|
+
|
|
1011
|
+
adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
|
|
1012
|
+
deg = np.array(adj.sum(axis=1)).flatten()
|
|
1013
|
+
deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
|
|
1014
|
+
D_inv = sp.diags(deg_inv_sqrt)
|
|
1015
|
+
P = (D_inv @ adj @ D_inv).tocsr()
|
|
1016
|
+
|
|
1017
|
+
if self.missing_mask is not None:
|
|
1018
|
+
if isinstance(self.missing_mask, str):
|
|
1019
|
+
mask = to_numpy(data[self.missing_mask])
|
|
1020
|
+
else:
|
|
1021
|
+
mask = to_numpy(self.missing_mask)
|
|
1022
|
+
else:
|
|
1023
|
+
mask = np.isnan(x_np)
|
|
1024
|
+
|
|
1025
|
+
if mask.ndim == 1:
|
|
1026
|
+
mask = mask[:, None]
|
|
1027
|
+
|
|
1028
|
+
x_orig = x_np.copy()
|
|
1029
|
+
x_np[mask.squeeze(-1) if mask.shape[-1] == 1 else mask] = 0.0
|
|
1030
|
+
|
|
1031
|
+
for _ in range(self.num_iterations):
|
|
1032
|
+
x_np = P @ x_np
|
|
1033
|
+
x_np[~mask.squeeze(-1) if mask.shape[-1] == 1 else ~mask] = x_orig[~mask.squeeze(-1) if mask.shape[-1] == 1 else ~mask]
|
|
1034
|
+
|
|
1035
|
+
data.x = match_tensor(x_np, data.x)
|
|
1036
|
+
return data
|
|
1037
|
+
|
|
1038
|
+
def __repr__(self) -> str:
|
|
1039
|
+
return f"{self.__class__.__name__}(num_iterations={self.num_iterations})"
|
|
1040
|
+
|
|
1041
|
+
|
|
1042
|
+
@functional_transform("half_hop")
|
|
1043
|
+
class HalfHop(BaseTransform):
|
|
1044
|
+
r"""Adds half-hop intermediate nodes."""
|
|
1045
|
+
|
|
1046
|
+
def __init__(self, alpha: float = 0.5):
|
|
1047
|
+
self.alpha = alpha
|
|
1048
|
+
|
|
1049
|
+
def forward(self, data: Data) -> Data:
|
|
1050
|
+
assert data.edge_index is not None
|
|
1051
|
+
ei_np = to_numpy(data.edge_index)
|
|
1052
|
+
num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
|
|
1053
|
+
num_edges = ei_np.shape[1]
|
|
1054
|
+
|
|
1055
|
+
# Each edge (u, v) gets an intermediate node num_nodes + idx
|
|
1056
|
+
inter = np.arange(num_nodes, num_nodes + num_edges, dtype=ei_np.dtype)
|
|
1057
|
+
e1 = np.stack([ei_np[0], inter], axis=0)
|
|
1058
|
+
e2 = np.stack([inter, ei_np[1]], axis=0)
|
|
1059
|
+
data.edge_index = match_tensor(np.concatenate([ei_np, e1, e2], axis=1), data.edge_index)
|
|
1060
|
+
|
|
1061
|
+
if hasattr(data, "x") and data.x is not None:
|
|
1062
|
+
x_np = to_numpy(data.x)
|
|
1063
|
+
inter_x = self.alpha * x_np[ei_np[0]] + (1 - self.alpha) * x_np[ei_np[1]]
|
|
1064
|
+
data.x = match_tensor(np.concatenate([x_np, inter_x], axis=0), data.x)
|
|
1065
|
+
|
|
1066
|
+
data.num_nodes = num_nodes + num_edges
|
|
1067
|
+
return data
|
|
1068
|
+
|
|
1069
|
+
def __repr__(self) -> str:
|
|
1070
|
+
return f"{self.__class__.__name__}(alpha={self.alpha})"
|