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,89 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
from .utils import PoolingByMultiheadAttention, SetAttentionBlock
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class GraphMultisetTransformer(Aggregation):
|
|
9
|
+
r"""The Graph Multiset Transformer pooling operator from the
|
|
10
|
+
`"Accurate Learning of Graph Representations
|
|
11
|
+
with Graph Multiset Pooling" <https://arxiv.org/abs/2102.11533>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import GraphMultisetTransformer
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
20
|
+
|
|
21
|
+
aggr = GraphMultisetTransformer(channels=8, k=2)
|
|
22
|
+
out = aggr(x, index=index, dim_size=2)
|
|
23
|
+
print(tuple(out.shape)) # (2, 8)
|
|
24
|
+
```
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
channels: int,
|
|
30
|
+
k: int,
|
|
31
|
+
num_encoder_blocks: int = 1,
|
|
32
|
+
heads: int = 1,
|
|
33
|
+
layer_norm: bool = False,
|
|
34
|
+
dropout: float = 0.0,
|
|
35
|
+
**kwargs,
|
|
36
|
+
):
|
|
37
|
+
super().__init__(**kwargs)
|
|
38
|
+
self.channels = channels
|
|
39
|
+
self.k = k
|
|
40
|
+
self.heads = heads
|
|
41
|
+
self.layer_norm = layer_norm
|
|
42
|
+
self.dropout = dropout
|
|
43
|
+
|
|
44
|
+
self.pma1 = PoolingByMultiheadAttention(channels, k, heads, layer_norm, dropout)
|
|
45
|
+
self.encoders = [
|
|
46
|
+
SetAttentionBlock(channels, heads, layer_norm, dropout)
|
|
47
|
+
for _ in range(num_encoder_blocks)
|
|
48
|
+
]
|
|
49
|
+
self.pma2 = PoolingByMultiheadAttention(channels, 1, heads, layer_norm, dropout)
|
|
50
|
+
|
|
51
|
+
def reset_parameters(self):
|
|
52
|
+
self.pma1.reset_parameters()
|
|
53
|
+
for encoder in self.encoders:
|
|
54
|
+
encoder.reset_parameters()
|
|
55
|
+
self.pma2.reset_parameters()
|
|
56
|
+
|
|
57
|
+
def call(
|
|
58
|
+
self,
|
|
59
|
+
x,
|
|
60
|
+
index: Optional[any] = None,
|
|
61
|
+
ptr: Optional[any] = None,
|
|
62
|
+
dim_size: Optional[int] = None,
|
|
63
|
+
dim: int = -2,
|
|
64
|
+
max_num_elements: Optional[int] = None,
|
|
65
|
+
training: bool = False,
|
|
66
|
+
**kwargs,
|
|
67
|
+
):
|
|
68
|
+
x_dense, mask = self.to_dense_batch(
|
|
69
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
70
|
+
max_num_elements=max_num_elements,
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
x_dense = self.pma1(x_dense, mask=mask, training=training)
|
|
74
|
+
|
|
75
|
+
for encoder in self.encoders:
|
|
76
|
+
x_dense = encoder(x_dense, training=training)
|
|
77
|
+
|
|
78
|
+
x_dense = self.pma2(x_dense, training=training)
|
|
79
|
+
|
|
80
|
+
# Output shape: [B, channels]
|
|
81
|
+
return ops.squeeze(x_dense, axis=1)
|
|
82
|
+
|
|
83
|
+
def __repr__(self) -> str:
|
|
84
|
+
return (
|
|
85
|
+
f"{self.__class__.__name__}({self.channels}, k={self.k}, "
|
|
86
|
+
f"heads={self.heads}, layer_norm={self.layer_norm}, "
|
|
87
|
+
f"dropout={self.dropout})"
|
|
88
|
+
)
|
|
89
|
+
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class GRUAggregation(Aggregation):
|
|
8
|
+
r"""Performs GRU aggregation in which the elements to aggregate are
|
|
9
|
+
interpreted as a sequence, as described in the `"Graph Neural Networks
|
|
10
|
+
with Adaptive Readouts" <https://arxiv.org/abs/2211.04952>`_ paper.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import GRUAggregation
|
|
16
|
+
|
|
17
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
18
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
19
|
+
|
|
20
|
+
aggr = GRUAggregation(in_channels=8, out_channels=16)
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 16)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(self, in_channels: int, out_channels: int, **kwargs):
|
|
27
|
+
super().__init__(**kwargs)
|
|
28
|
+
self.in_channels = in_channels
|
|
29
|
+
self.out_channels = out_channels
|
|
30
|
+
self.gru = layers.GRU(out_channels, return_sequences=False)
|
|
31
|
+
|
|
32
|
+
def build(self, input_shape=None):
|
|
33
|
+
self.gru.build((None, None, self.in_channels))
|
|
34
|
+
super().build(input_shape)
|
|
35
|
+
|
|
36
|
+
def reset_parameters(self):
|
|
37
|
+
self.gru.reset_parameters()
|
|
38
|
+
|
|
39
|
+
def call(
|
|
40
|
+
self,
|
|
41
|
+
x,
|
|
42
|
+
index: Optional[any] = None,
|
|
43
|
+
ptr: Optional[any] = None,
|
|
44
|
+
dim_size: Optional[int] = None,
|
|
45
|
+
dim: int = -2,
|
|
46
|
+
max_num_elements: Optional[int] = None,
|
|
47
|
+
training: bool = False,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
x_dense, _ = self.to_dense_batch(
|
|
51
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
52
|
+
max_num_elements=max_num_elements,
|
|
53
|
+
)
|
|
54
|
+
return self.gru(x_dense, training=training)
|
|
55
|
+
|
|
56
|
+
def __repr__(self) -> str:
|
|
57
|
+
return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
|
|
58
|
+
|
|
@@ -0,0 +1,143 @@
|
|
|
1
|
+
from math import ceil, log2
|
|
2
|
+
from typing import Optional
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from .base import Aggregation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class LCMAggregation(Aggregation):
|
|
9
|
+
r"""The Learnable Commutative Monoid aggregation from the
|
|
10
|
+
`"Learnable Commutative Monoids for Graph Neural Networks"
|
|
11
|
+
<https://arxiv.org/abs/2212.08541>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import LCMAggregation
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
20
|
+
|
|
21
|
+
aggr = LCMAggregation(in_channels=8, out_channels=16)
|
|
22
|
+
out = aggr(x, index=index, dim_size=2)
|
|
23
|
+
print(tuple(out.shape)) # (2, 16)
|
|
24
|
+
```
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: int,
|
|
30
|
+
out_channels: int,
|
|
31
|
+
project: bool = True,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
if in_channels != out_channels and not project:
|
|
36
|
+
raise ValueError(
|
|
37
|
+
f"Inputs of '{self.__class__.__name__}' must be projected if `in_channels != out_channels`"
|
|
38
|
+
)
|
|
39
|
+
|
|
40
|
+
self.in_channels = in_channels
|
|
41
|
+
self.out_channels = out_channels
|
|
42
|
+
self.project = project
|
|
43
|
+
|
|
44
|
+
self.lin = layers.Dense(out_channels) if project else None
|
|
45
|
+
self.gru_cell = layers.GRUCell(out_channels)
|
|
46
|
+
|
|
47
|
+
def reset_parameters(self):
|
|
48
|
+
if self.lin is not None:
|
|
49
|
+
self.lin.reset_parameters()
|
|
50
|
+
self.gru_cell.reset_parameters()
|
|
51
|
+
|
|
52
|
+
def call(
|
|
53
|
+
self,
|
|
54
|
+
x,
|
|
55
|
+
index: Optional[any] = None,
|
|
56
|
+
ptr: Optional[any] = None,
|
|
57
|
+
dim_size: Optional[int] = None,
|
|
58
|
+
dim: int = -2,
|
|
59
|
+
max_num_elements: Optional[int] = None,
|
|
60
|
+
**kwargs,
|
|
61
|
+
):
|
|
62
|
+
if self.lin is not None:
|
|
63
|
+
x = ops.relu(self.lin(x))
|
|
64
|
+
|
|
65
|
+
x_dense, _ = self.to_dense_batch(
|
|
66
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
67
|
+
max_num_elements=max_num_elements,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
# Transpose to [num_neighbors, num_nodes, num_features]
|
|
71
|
+
x_dense = ops.transpose(x_dense, (1, 0, 2))
|
|
72
|
+
num_neighbors = ops.shape(x_dense)[0]
|
|
73
|
+
num_nodes = ops.shape(x_dense)[1]
|
|
74
|
+
num_features = ops.shape(x_dense)[2]
|
|
75
|
+
|
|
76
|
+
if not isinstance(num_neighbors, int): # a tensor while tracing (e.g. TensorFlow's fit)
|
|
77
|
+
return self._reduce_traced(x_dense)
|
|
78
|
+
if num_neighbors == 0:
|
|
79
|
+
return ops.zeros((num_nodes, self.out_channels), dtype=x.dtype)
|
|
80
|
+
|
|
81
|
+
depth = ceil(log2(max(num_neighbors, 1)))
|
|
82
|
+
for _ in range(depth):
|
|
83
|
+
curr_len = ops.shape(x_dense)[0]
|
|
84
|
+
if curr_len <= 1:
|
|
85
|
+
break
|
|
86
|
+
half_size = ceil(curr_len / 2)
|
|
87
|
+
|
|
88
|
+
if curr_len % 2 == 1:
|
|
89
|
+
x_pair = x_dense[:-1]
|
|
90
|
+
remainder = x_dense[-1:]
|
|
91
|
+
else:
|
|
92
|
+
x_pair = x_dense
|
|
93
|
+
remainder = None
|
|
94
|
+
|
|
95
|
+
# x_pair: [2 * half, num_nodes, num_features]
|
|
96
|
+
pair_count = ops.shape(x_pair)[0] // 2
|
|
97
|
+
x_pair = ops.reshape(x_pair, (pair_count, 2, num_nodes, num_features))
|
|
98
|
+
left = x_pair[:, 0] # [pair_count, num_nodes, num_features]
|
|
99
|
+
right = x_pair[:, 1] # [pair_count, num_nodes, num_features]
|
|
100
|
+
|
|
101
|
+
left_flat = ops.reshape(left, (-1, num_features))
|
|
102
|
+
right_flat = ops.reshape(right, (-1, num_features))
|
|
103
|
+
|
|
104
|
+
# GRUCell: inputs=left, state=[right]
|
|
105
|
+
out1, _ = self.gru_cell(left_flat, [right_flat])
|
|
106
|
+
out2, _ = self.gru_cell(right_flat, [left_flat])
|
|
107
|
+
out = 0.5 * (out1 + out2)
|
|
108
|
+
out = ops.reshape(out, (pair_count, num_nodes, num_features))
|
|
109
|
+
|
|
110
|
+
if remainder is not None:
|
|
111
|
+
out = ops.concatenate([out, remainder], axis=0)
|
|
112
|
+
|
|
113
|
+
x_dense = out
|
|
114
|
+
|
|
115
|
+
return ops.squeeze(x_dense, axis=0)
|
|
116
|
+
|
|
117
|
+
def _combine(self, left, right):
|
|
118
|
+
num_features = ops.shape(left)[-1]
|
|
119
|
+
left_flat, right_flat = ops.reshape(left, (-1, num_features)), ops.reshape(right, (-1, num_features))
|
|
120
|
+
out1, _ = self.gru_cell(left_flat, [right_flat])
|
|
121
|
+
out2, _ = self.gru_cell(right_flat, [left_flat])
|
|
122
|
+
return ops.reshape(0.5 * (out1 + out2), ops.shape(left))
|
|
123
|
+
|
|
124
|
+
def _reduce_traced(self, x_dense):
|
|
125
|
+
"""The pairwise reduction of ``call`` for a neighbor count only known at run time: position
|
|
126
|
+
``i`` absorbs position ``i + stride`` when ``i`` is a multiple of ``2 * stride``, with the
|
|
127
|
+
stride doubling every level. This pairs the same elements as ``call``, in a fixed shape."""
|
|
128
|
+
length = ops.shape(x_dense)[0]
|
|
129
|
+
if not self.gru_cell.built: # no weights may be created inside the loop
|
|
130
|
+
self.gru_cell.build((None, x_dense.shape[-1]))
|
|
131
|
+
position = ops.arange(length, dtype="int32")
|
|
132
|
+
|
|
133
|
+
def body(stride, h):
|
|
134
|
+
combined = self._combine(h, ops.roll(h, -stride, axis=0))
|
|
135
|
+
active = ops.logical_and(position % (2 * stride) == 0, position + stride < length)
|
|
136
|
+
return stride * 2, ops.where(active[:, None, None], combined, h)
|
|
137
|
+
|
|
138
|
+
_, h = ops.while_loop(lambda stride, h: stride < length, body, (ops.convert_to_tensor(1, "int32"), x_dense))
|
|
139
|
+
return h[0]
|
|
140
|
+
|
|
141
|
+
def __repr__(self) -> str:
|
|
142
|
+
return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels}, project={self.project})"
|
|
143
|
+
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class LSTMAggregation(Aggregation):
|
|
8
|
+
r"""Performs LSTM-style aggregation in which the elements to aggregate are
|
|
9
|
+
interpreted as a sequence, as described in the `"Inductive Representation
|
|
10
|
+
Learning on Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import LSTMAggregation
|
|
16
|
+
|
|
17
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
18
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
19
|
+
|
|
20
|
+
aggr = LSTMAggregation(in_channels=8, out_channels=16)
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 16)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(self, in_channels: int, out_channels: int, **kwargs):
|
|
27
|
+
super().__init__(**kwargs)
|
|
28
|
+
self.in_channels = in_channels
|
|
29
|
+
self.out_channels = out_channels
|
|
30
|
+
self.lstm = layers.LSTM(out_channels, return_sequences=False)
|
|
31
|
+
|
|
32
|
+
def build(self, input_shape=None):
|
|
33
|
+
self.lstm.build((None, None, self.in_channels))
|
|
34
|
+
super().build(input_shape)
|
|
35
|
+
|
|
36
|
+
def reset_parameters(self):
|
|
37
|
+
self.lstm.reset_parameters()
|
|
38
|
+
|
|
39
|
+
def call(
|
|
40
|
+
self,
|
|
41
|
+
x,
|
|
42
|
+
index: Optional[any] = None,
|
|
43
|
+
ptr: Optional[any] = None,
|
|
44
|
+
dim_size: Optional[int] = None,
|
|
45
|
+
dim: int = -2,
|
|
46
|
+
max_num_elements: Optional[int] = None,
|
|
47
|
+
training: bool = False,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
x_dense, _ = self.to_dense_batch(
|
|
51
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
52
|
+
max_num_elements=max_num_elements,
|
|
53
|
+
)
|
|
54
|
+
return self.lstm(x_dense, training=training)
|
|
55
|
+
|
|
56
|
+
def __repr__(self) -> str:
|
|
57
|
+
return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
|
|
58
|
+
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class MLPAggregation(Aggregation):
|
|
8
|
+
r"""Performs MLP aggregation in which the elements to aggregate are
|
|
9
|
+
flattened into a single vectorial representation, and are then processed by
|
|
10
|
+
a Multi-Layer Perceptron (MLP).
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import MLPAggregation
|
|
16
|
+
|
|
17
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
18
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
19
|
+
|
|
20
|
+
aggr = MLPAggregation(in_channels=8, out_channels=16, max_num_elements=5)
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 16)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
in_channels: int,
|
|
29
|
+
out_channels: int,
|
|
30
|
+
max_num_elements: int,
|
|
31
|
+
mlp: Optional[any] = None,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
self.in_channels = in_channels
|
|
36
|
+
self.out_channels = out_channels
|
|
37
|
+
self.max_num_elements = max_num_elements
|
|
38
|
+
|
|
39
|
+
if mlp is None:
|
|
40
|
+
self.mlp = layers.Dense(out_channels)
|
|
41
|
+
else:
|
|
42
|
+
self.mlp = mlp
|
|
43
|
+
|
|
44
|
+
def build(self, input_shape=None):
|
|
45
|
+
if hasattr(self.mlp, "build"):
|
|
46
|
+
self.mlp.build((None, self.in_channels * self.max_num_elements))
|
|
47
|
+
super().build(input_shape)
|
|
48
|
+
|
|
49
|
+
def reset_parameters(self):
|
|
50
|
+
if hasattr(self.mlp, "reset_parameters"):
|
|
51
|
+
self.mlp.reset_parameters()
|
|
52
|
+
|
|
53
|
+
def call(
|
|
54
|
+
self,
|
|
55
|
+
x,
|
|
56
|
+
index: Optional[any] = None,
|
|
57
|
+
ptr: Optional[any] = None,
|
|
58
|
+
dim_size: Optional[int] = None,
|
|
59
|
+
dim: int = -2,
|
|
60
|
+
**kwargs,
|
|
61
|
+
):
|
|
62
|
+
x_dense, _ = self.to_dense_batch(
|
|
63
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
64
|
+
max_num_elements=self.max_num_elements,
|
|
65
|
+
)
|
|
66
|
+
B = ops.shape(x_dense)[0]
|
|
67
|
+
flattened = ops.reshape(x_dense, (B, self.max_num_elements * self.in_channels))
|
|
68
|
+
return self.mlp(flattened)
|
|
69
|
+
|
|
70
|
+
def __repr__(self) -> str:
|
|
71
|
+
return (
|
|
72
|
+
f"{self.__class__.__name__}({self.in_channels}, {self.out_channels}, "
|
|
73
|
+
f"max_num_elements={self.max_num_elements})"
|
|
74
|
+
)
|
|
75
|
+
|
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Any, Dict, List, Optional, Union
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from .base import Aggregation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class MultiAggregation(Aggregation):
|
|
9
|
+
r"""Performs aggregations with one or more aggregators and combines
|
|
10
|
+
aggregated results, as described in the `"Principal Neighbourhood
|
|
11
|
+
Aggregation for Graph Nets" <https://arxiv.org/abs/2004.05718>`_ and
|
|
12
|
+
`"Adaptive Filters and Aggregator Fusion for Efficient Graph Convolutions"
|
|
13
|
+
<https://arxiv.org/abs/2104.01481>`_ papers.
|
|
14
|
+
|
|
15
|
+
Example:
|
|
16
|
+
```python
|
|
17
|
+
import numpy as np
|
|
18
|
+
from k3_node.layers import MultiAggregation
|
|
19
|
+
|
|
20
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
21
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
22
|
+
|
|
23
|
+
aggr = MultiAggregation(aggrs=["sum", "mean", "max"])
|
|
24
|
+
out = aggr(x, index=index, dim_size=2)
|
|
25
|
+
print(tuple(out.shape)) # (2, 24)
|
|
26
|
+
```
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self,
|
|
31
|
+
aggrs: List[Union[Aggregation, str]],
|
|
32
|
+
aggrs_kwargs: Optional[List[Dict[str, Any]]] = None,
|
|
33
|
+
mode: Optional[str] = "cat",
|
|
34
|
+
mode_kwargs: Optional[Dict[str, Any]] = None,
|
|
35
|
+
**kwargs,
|
|
36
|
+
):
|
|
37
|
+
super().__init__(**kwargs)
|
|
38
|
+
|
|
39
|
+
if not isinstance(aggrs, (list, tuple)):
|
|
40
|
+
raise ValueError(f"'aggrs' of '{self.__class__.__name__}' should be a list or tuple.")
|
|
41
|
+
|
|
42
|
+
if len(aggrs) == 0:
|
|
43
|
+
raise ValueError(f"'aggrs' of '{self.__class__.__name__}' should not be empty.")
|
|
44
|
+
|
|
45
|
+
if aggrs_kwargs is None:
|
|
46
|
+
aggrs_kwargs = [{}] * len(aggrs)
|
|
47
|
+
elif len(aggrs) != len(aggrs_kwargs):
|
|
48
|
+
raise ValueError(
|
|
49
|
+
f"'aggrs_kwargs' with invalid length passed to '{self.__class__.__name__}' "
|
|
50
|
+
f"(got '{len(aggrs_kwargs)}', expected '{len(aggrs)}')."
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
from .resolver import aggregation_resolver
|
|
54
|
+
self.aggrs = [
|
|
55
|
+
aggregation_resolver(aggr, **aggr_kw)
|
|
56
|
+
for aggr, aggr_kw in zip(aggrs, aggrs_kwargs)
|
|
57
|
+
]
|
|
58
|
+
|
|
59
|
+
self.mode = mode
|
|
60
|
+
mode_kwargs = copy.copy(mode_kwargs) or {}
|
|
61
|
+
self.in_channels = mode_kwargs.pop("in_channels", None)
|
|
62
|
+
self.out_channels = mode_kwargs.pop("out_channels", None)
|
|
63
|
+
|
|
64
|
+
if mode in ["proj", "attn"]:
|
|
65
|
+
if len(aggrs) == 1:
|
|
66
|
+
raise ValueError("Multiple aggregations are required for 'proj' or 'attn' combine mode.")
|
|
67
|
+
if self.in_channels is None or self.out_channels is None:
|
|
68
|
+
raise ValueError(f"Combine mode '{mode}' must have `in_channels` and `out_channels` specified.")
|
|
69
|
+
|
|
70
|
+
if isinstance(self.in_channels, int):
|
|
71
|
+
self.in_channels = [self.in_channels] * len(aggrs)
|
|
72
|
+
|
|
73
|
+
if mode == "proj":
|
|
74
|
+
self.lin = layers.Dense(self.out_channels, **mode_kwargs)
|
|
75
|
+
elif mode == "attn":
|
|
76
|
+
from ..dense import HeteroDictLinear
|
|
77
|
+
channels = {str(k): v for k, v in enumerate(self.in_channels)}
|
|
78
|
+
self.lin_heads = HeteroDictLinear(channels, self.out_channels)
|
|
79
|
+
num_heads = mode_kwargs.pop("num_heads", 1)
|
|
80
|
+
self.multihead_attn = layers.MultiHeadAttention(
|
|
81
|
+
num_heads=num_heads,
|
|
82
|
+
key_dim=max(self.out_channels // num_heads, 1),
|
|
83
|
+
**mode_kwargs,
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
def reset_parameters(self):
|
|
87
|
+
for aggr in self.aggrs:
|
|
88
|
+
if hasattr(aggr, "reset_parameters"):
|
|
89
|
+
aggr.reset_parameters()
|
|
90
|
+
if hasattr(self, "lin") and hasattr(self.lin, "reset_parameters"):
|
|
91
|
+
self.lin.reset_parameters()
|
|
92
|
+
if hasattr(self, "lin_heads") and hasattr(self.lin_heads, "reset_parameters"):
|
|
93
|
+
self.lin_heads.reset_parameters()
|
|
94
|
+
|
|
95
|
+
def get_out_channels(self, in_channels: int) -> int:
|
|
96
|
+
if self.out_channels is not None:
|
|
97
|
+
return self.out_channels
|
|
98
|
+
if self.mode == "cat":
|
|
99
|
+
return in_channels * len(self.aggrs)
|
|
100
|
+
return in_channels
|
|
101
|
+
|
|
102
|
+
def call(
|
|
103
|
+
self,
|
|
104
|
+
x,
|
|
105
|
+
index: Optional[any] = None,
|
|
106
|
+
ptr: Optional[any] = None,
|
|
107
|
+
dim_size: Optional[int] = None,
|
|
108
|
+
dim: int = -2,
|
|
109
|
+
**kwargs,
|
|
110
|
+
):
|
|
111
|
+
outs = [aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, **kwargs) for aggr in self.aggrs]
|
|
112
|
+
return self.combine(outs)
|
|
113
|
+
|
|
114
|
+
def combine(self, inputs: List[any]):
|
|
115
|
+
if len(inputs) == 1:
|
|
116
|
+
return inputs[0]
|
|
117
|
+
|
|
118
|
+
if self.mode == "cat":
|
|
119
|
+
return ops.concatenate(inputs, axis=-1)
|
|
120
|
+
|
|
121
|
+
if hasattr(self, "lin"):
|
|
122
|
+
return self.lin(ops.concatenate(inputs, axis=-1))
|
|
123
|
+
|
|
124
|
+
if hasattr(self, "multihead_attn"):
|
|
125
|
+
x_dict = {str(k): v for k, v in enumerate(inputs)}
|
|
126
|
+
x_dict = self.lin_heads(x_dict)
|
|
127
|
+
xs = [x_dict[str(key)] for key in range(len(inputs))]
|
|
128
|
+
# xs: [num_aggrs, B, D] -> transpose to [B, num_aggrs, D]
|
|
129
|
+
x_stack = ops.transpose(ops.stack(xs, axis=0), (1, 0, 2))
|
|
130
|
+
attn_out = self.multihead_attn(x_stack, x_stack, x_stack)
|
|
131
|
+
return ops.mean(attn_out, axis=1)
|
|
132
|
+
|
|
133
|
+
stacked = ops.stack(inputs, axis=0) # [num_aggrs, B, D]
|
|
134
|
+
if self.mode == "sum":
|
|
135
|
+
return ops.sum(stacked, axis=0)
|
|
136
|
+
elif self.mode == "mean":
|
|
137
|
+
return ops.mean(stacked, axis=0)
|
|
138
|
+
elif self.mode == "max":
|
|
139
|
+
return ops.max(stacked, axis=0)
|
|
140
|
+
elif self.mode == "min":
|
|
141
|
+
return ops.min(stacked, axis=0)
|
|
142
|
+
elif self.mode == "logsumexp":
|
|
143
|
+
return ops.logsumexp(stacked, axis=0)
|
|
144
|
+
elif self.mode == "std":
|
|
145
|
+
return ops.std(stacked, axis=0)
|
|
146
|
+
elif self.mode == "var":
|
|
147
|
+
return ops.var(stacked, axis=0)
|
|
148
|
+
|
|
149
|
+
raise ValueError(f"Combine mode '{self.mode}' is not supported.")
|
|
150
|
+
|
|
151
|
+
def __repr__(self) -> str:
|
|
152
|
+
aggrs = ",\n".join([f" {aggr}" for aggr in self.aggrs])
|
|
153
|
+
return f"{self.__class__.__name__}([\n{aggrs}\n], mode={self.mode})"
|
|
154
|
+
|