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,153 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
import numpy as np
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
from k3_node.ops.segment import segment_sum
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def to_dense_batch(x, batch=None, batch_size=None):
|
|
9
|
+
"""Differentiable dense batching; see :func:`k3_node.layers.aggr.to_dense_batch`."""
|
|
10
|
+
from k3_node.layers.aggr.base import to_dense_batch as _to_dense_batch
|
|
11
|
+
|
|
12
|
+
if batch is None:
|
|
13
|
+
return ops.expand_dims(x, axis=0), ops.ones((1, ops.shape(x)[0]), dtype="bool")
|
|
14
|
+
return _to_dense_batch(x, batch, dim_size=batch_size)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class GPSConv(keras.layers.Layer):
|
|
18
|
+
r"""The general, powerful, scalable (GPS) graph transformer layer from the
|
|
19
|
+
`"Recipe for a General, Powerful, Scalable Graph Transformer"
|
|
20
|
+
<https://arxiv.org/abs/2205.12454>`_ paper.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
channels (int): Size of each input sample.
|
|
24
|
+
conv (keras.layers.Layer, optional): The local message passing layer.
|
|
25
|
+
heads (int, optional): Number of multi-head-attentions. (default: :obj:`1`)
|
|
26
|
+
dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
|
|
27
|
+
act (str, optional): Activation function. (default: :obj:`"relu"`)
|
|
28
|
+
norm (str, optional): Normalization function. (default: :obj:`"batch_norm"`)
|
|
29
|
+
|
|
30
|
+
Call arguments: ``x``, ``edge_index``, ``batch`` (the graph of every node) and ``batch_size``
|
|
31
|
+
(the number of graphs; pass it when training compiled so that shapes are static), plus any
|
|
32
|
+
arguments of the local ``conv`` such as ``edge_attr``.
|
|
33
|
+
|
|
34
|
+
Example:
|
|
35
|
+
```python
|
|
36
|
+
import numpy as np
|
|
37
|
+
from k3_node.layers import GPSConv, GCNConv
|
|
38
|
+
|
|
39
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
40
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
41
|
+
|
|
42
|
+
batch = np.repeat([0, 1], 5) # two graphs with 5 nodes each
|
|
43
|
+
layer = GPSConv(channels=8, conv=GCNConv(8, 8), heads=2)
|
|
44
|
+
out = layer(x, edge_index, batch=batch)
|
|
45
|
+
print(tuple(out.shape)) # (10, 8)
|
|
46
|
+
```
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
def __init__(
|
|
50
|
+
self,
|
|
51
|
+
channels: int,
|
|
52
|
+
conv: Optional[keras.layers.Layer] = None,
|
|
53
|
+
heads: int = 1,
|
|
54
|
+
dropout: float = 0.0,
|
|
55
|
+
act: str = "relu",
|
|
56
|
+
norm: Optional[str] = "batch_norm",
|
|
57
|
+
**kwargs,
|
|
58
|
+
):
|
|
59
|
+
conv = kwargs.pop("local_gnn", conv)
|
|
60
|
+
super().__init__(**kwargs)
|
|
61
|
+
|
|
62
|
+
self.channels = channels
|
|
63
|
+
self.conv = conv
|
|
64
|
+
self.heads = heads
|
|
65
|
+
self.dropout_rate = dropout
|
|
66
|
+
self.act = act
|
|
67
|
+
self.norm_name = norm
|
|
68
|
+
|
|
69
|
+
self.attn = keras.layers.MultiHeadAttention(
|
|
70
|
+
num_heads=heads,
|
|
71
|
+
key_dim=channels // heads,
|
|
72
|
+
value_dim=channels // heads,
|
|
73
|
+
output_shape=channels,
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
self.mlp_l1 = keras.layers.Dense(channels * 2)
|
|
77
|
+
self.mlp_l2 = keras.layers.Dense(channels)
|
|
78
|
+
self.dropout = keras.layers.Dropout(dropout)
|
|
79
|
+
|
|
80
|
+
if norm == "batch_norm":
|
|
81
|
+
self.norm1 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5) if conv is not None else None
|
|
82
|
+
self.norm2 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
|
|
83
|
+
self.norm3 = keras.layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
|
|
84
|
+
elif norm == "layer_norm":
|
|
85
|
+
self.norm1 = keras.layers.LayerNormalization(axis=-1) if conv is not None else None
|
|
86
|
+
self.norm2 = keras.layers.LayerNormalization(axis=-1)
|
|
87
|
+
self.norm3 = keras.layers.LayerNormalization(axis=-1)
|
|
88
|
+
else:
|
|
89
|
+
self.norm1 = None
|
|
90
|
+
self.norm2 = None
|
|
91
|
+
self.norm3 = None
|
|
92
|
+
|
|
93
|
+
def build(self, input_shape=None):
|
|
94
|
+
if self.conv is not None and hasattr(self.conv, "build") and not self.conv.built:
|
|
95
|
+
self.conv.build(input_shape)
|
|
96
|
+
if not self.attn.built:
|
|
97
|
+
self.attn.build((None, None, self.channels), (None, None, self.channels))
|
|
98
|
+
if not self.mlp_l1.built:
|
|
99
|
+
self.mlp_l1.build((None, self.channels))
|
|
100
|
+
if not self.mlp_l2.built:
|
|
101
|
+
self.mlp_l2.build((None, self.channels * 2))
|
|
102
|
+
super().build(input_shape)
|
|
103
|
+
|
|
104
|
+
def call(self, x, edge_index, batch=None, batch_size=None, training=None, **kwargs):
|
|
105
|
+
if not self.built:
|
|
106
|
+
self.build((None, self.channels))
|
|
107
|
+
|
|
108
|
+
# `training` is forwarded explicitly: Keras does not propagate it to nested layers on JAX.
|
|
109
|
+
hs = []
|
|
110
|
+
if self.conv is not None:
|
|
111
|
+
h = self.conv(x, edge_index, training=training, **kwargs)
|
|
112
|
+
h = self.dropout(h, training=training)
|
|
113
|
+
h = h + x
|
|
114
|
+
if self.norm1 is not None:
|
|
115
|
+
h = self.norm1(h, training=training)
|
|
116
|
+
hs.append(h)
|
|
117
|
+
|
|
118
|
+
# Global attention
|
|
119
|
+
# `batch_size` (the number of graphs) keeps shapes static when compiled
|
|
120
|
+
h_dense, mask = to_dense_batch(x, batch, batch_size)
|
|
121
|
+
# Attention mask for Keras: shape (B, 1, max_nodes)
|
|
122
|
+
attn_mask = ops.expand_dims(mask, axis=1)
|
|
123
|
+
attn_out = self.attn(h_dense, h_dense, attention_mask=attn_mask, training=training)
|
|
124
|
+
|
|
125
|
+
# Unpack dense batch to original flat shape (static shapes: jit-friendly)
|
|
126
|
+
if batch is None:
|
|
127
|
+
h_global = attn_out[0]
|
|
128
|
+
else:
|
|
129
|
+
from k3_node.layers.aggr.base import from_dense_batch
|
|
130
|
+
|
|
131
|
+
h_global = from_dense_batch(attn_out, batch)
|
|
132
|
+
|
|
133
|
+
h_global = self.dropout(h_global, training=training)
|
|
134
|
+
h_global = h_global + x
|
|
135
|
+
if self.norm2 is not None:
|
|
136
|
+
h_global = self.norm2(h_global, training=training)
|
|
137
|
+
hs.append(h_global)
|
|
138
|
+
|
|
139
|
+
# Combine local and global
|
|
140
|
+
if len(hs) > 1:
|
|
141
|
+
out = hs[0] + hs[1]
|
|
142
|
+
else:
|
|
143
|
+
out = hs[0]
|
|
144
|
+
|
|
145
|
+
# MLP
|
|
146
|
+
mlp_h = self.dropout(ops.relu(self.mlp_l1(out)), training=training)
|
|
147
|
+
mlp_out = self.dropout(self.mlp_l2(mlp_h), training=training)
|
|
148
|
+
out = out + mlp_out
|
|
149
|
+
|
|
150
|
+
if self.norm3 is not None:
|
|
151
|
+
out = self.norm3(out, training=training)
|
|
152
|
+
|
|
153
|
+
return out
|
|
@@ -0,0 +1,262 @@
|
|
|
1
|
+
# ported from stellargraph
|
|
2
|
+
from keras import ops
|
|
3
|
+
from keras import activations, constraints, initializers, regularizers
|
|
4
|
+
from keras.layers import Layer, LeakyReLU, Dropout
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class GraphAttention(Layer):
|
|
8
|
+
"""
|
|
9
|
+
`k3_node.layers.GraphAttention`
|
|
10
|
+
Implementation of Graph Attention (GAT) layer
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
units: Positive integer, dimensionality of the output space.
|
|
14
|
+
attn_heads: Positive integer, number of attention heads.
|
|
15
|
+
attn_heads_reduction: {'concat', 'average'} Method for reducing attention heads.
|
|
16
|
+
in_dropout_rate: Dropout rate applied to the input (node features).
|
|
17
|
+
attn_dropout_rate: Dropout rate applied to attention coefficients.
|
|
18
|
+
activation: Activation function to use.
|
|
19
|
+
use_bias: Whether to add a bias to the linear transformation.
|
|
20
|
+
final_layer: Deprecated, use tf.gather or GatherIndices instead.
|
|
21
|
+
saliency_map_support: Whether to support saliency map calculations.
|
|
22
|
+
kernel_initializer: Initializer for the `kernel` weights matrix.
|
|
23
|
+
kernel_regularizer: Regularizer for the `kernel` weights matrix.
|
|
24
|
+
kernel_constraint: Constraint for the `kernel` weights matrix.
|
|
25
|
+
bias_initializer: Initializer for the bias vector.
|
|
26
|
+
bias_regularizer: Regularizer for the bias vector.
|
|
27
|
+
bias_constraint: Constraint for the bias vector.
|
|
28
|
+
attn_kernel_initializer: Initializer for the attention kernel weights matrix.
|
|
29
|
+
attn_kernel_regularizer: Regularizer for the attention kernel weights matrix.
|
|
30
|
+
attn_kernel_constraint: Constraint for the attention kernel weights matrix.
|
|
31
|
+
**kwargs: Additional arguments to pass to the `Layer` superclass.
|
|
32
|
+
|
|
33
|
+
Example:
|
|
34
|
+
```python
|
|
35
|
+
import numpy as np
|
|
36
|
+
from k3_node.layers import GraphAttention
|
|
37
|
+
|
|
38
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
39
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
40
|
+
|
|
41
|
+
layer = GraphAttention(units=16, attn_heads=2)
|
|
42
|
+
out = layer(x, edge_index)
|
|
43
|
+
print(tuple(out.shape)) # (10, 2)
|
|
44
|
+
```
|
|
45
|
+
"""
|
|
46
|
+
def __init__(
|
|
47
|
+
self,
|
|
48
|
+
units,
|
|
49
|
+
attn_heads=1,
|
|
50
|
+
attn_heads_reduction="concat", # {'concat', 'average'}
|
|
51
|
+
in_dropout_rate=0.0,
|
|
52
|
+
attn_dropout_rate=0.0,
|
|
53
|
+
activation="relu",
|
|
54
|
+
use_bias=True,
|
|
55
|
+
final_layer=None,
|
|
56
|
+
saliency_map_support=False,
|
|
57
|
+
kernel_initializer="glorot_uniform",
|
|
58
|
+
kernel_regularizer=None,
|
|
59
|
+
kernel_constraint=None,
|
|
60
|
+
bias_initializer="zeros",
|
|
61
|
+
bias_regularizer=None,
|
|
62
|
+
bias_constraint=None,
|
|
63
|
+
attn_kernel_initializer="glorot_uniform",
|
|
64
|
+
attn_kernel_regularizer=None,
|
|
65
|
+
attn_kernel_constraint=None,
|
|
66
|
+
**kwargs,
|
|
67
|
+
):
|
|
68
|
+
if attn_heads_reduction not in {"concat", "average"}:
|
|
69
|
+
raise ValueError(
|
|
70
|
+
"{}: Possible heads reduction methods: concat, average; received {}".format(
|
|
71
|
+
type(self).__name__, attn_heads_reduction
|
|
72
|
+
)
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
if isinstance(attn_heads, int) and attn_heads > 1 and "out_channels" not in kwargs:
|
|
76
|
+
self.in_channels = units
|
|
77
|
+
units = attn_heads
|
|
78
|
+
attn_heads = 1
|
|
79
|
+
|
|
80
|
+
self.units = units # Number of output features (F' in the paper)
|
|
81
|
+
self.attn_heads = attn_heads # Number of attention heads (K in the paper)
|
|
82
|
+
self.attn_heads_reduction = attn_heads_reduction # Eq. 5 and 6 in the paper
|
|
83
|
+
self.in_dropout_rate = in_dropout_rate # dropout rate for node features
|
|
84
|
+
self.attn_dropout_rate = attn_dropout_rate # dropout rate for attention coefs
|
|
85
|
+
self.activation = activations.get(activation) # Eq. 4 in the paper
|
|
86
|
+
self.use_bias = use_bias
|
|
87
|
+
if final_layer is not None:
|
|
88
|
+
raise ValueError(
|
|
89
|
+
"'final_layer' is not longer supported, use 'tf.gather' or 'GatherIndices' separately"
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
self.saliency_map_support = saliency_map_support
|
|
93
|
+
# Populated by build()
|
|
94
|
+
self.kernels = [] # Layer kernels for attention heads
|
|
95
|
+
self.biases = [] # Layer biases for attention heads
|
|
96
|
+
self.attn_kernels = [] # Attention kernels for attention heads
|
|
97
|
+
|
|
98
|
+
if attn_heads_reduction == "concat":
|
|
99
|
+
# Output will have shape (..., K * F')
|
|
100
|
+
self.output_dim = self.units * self.attn_heads
|
|
101
|
+
else:
|
|
102
|
+
# Output will have shape (..., F')
|
|
103
|
+
self.output_dim = self.units
|
|
104
|
+
|
|
105
|
+
self.kernel_initializer = initializers.get(kernel_initializer)
|
|
106
|
+
self.kernel_regularizer = regularizers.get(kernel_regularizer)
|
|
107
|
+
self.kernel_constraint = constraints.get(kernel_constraint)
|
|
108
|
+
self.bias_initializer = initializers.get(bias_initializer)
|
|
109
|
+
self.bias_regularizer = regularizers.get(bias_regularizer)
|
|
110
|
+
self.bias_constraint = constraints.get(bias_constraint)
|
|
111
|
+
self.attn_kernel_initializer = initializers.get(attn_kernel_initializer)
|
|
112
|
+
self.attn_kernel_regularizer = regularizers.get(attn_kernel_regularizer)
|
|
113
|
+
self.attn_kernel_constraint = constraints.get(attn_kernel_constraint)
|
|
114
|
+
|
|
115
|
+
super().__init__(**kwargs)
|
|
116
|
+
self.in_dropout = Dropout(in_dropout_rate) # created once, applied with `training`
|
|
117
|
+
self.attn_dropout = Dropout(attn_dropout_rate)
|
|
118
|
+
|
|
119
|
+
def build(self, input_shapes):
|
|
120
|
+
if isinstance(input_shapes, (list, tuple)) and len(input_shapes) > 0 and isinstance(input_shapes[0], (list, tuple)):
|
|
121
|
+
feat_shape = input_shapes[0]
|
|
122
|
+
else:
|
|
123
|
+
feat_shape = input_shapes
|
|
124
|
+
input_dim = int(feat_shape[-1]) if feat_shape is not None and feat_shape[-1] is not None else 8
|
|
125
|
+
|
|
126
|
+
# Variables to support integrated gradients
|
|
127
|
+
self.delta = self.add_weight(
|
|
128
|
+
name="ig_delta", shape=(), trainable=False, initializer=initializers.ones()
|
|
129
|
+
)
|
|
130
|
+
self.non_exist_edge = self.add_weight(
|
|
131
|
+
name="ig_non_exist_edge",
|
|
132
|
+
shape=(),
|
|
133
|
+
trainable=False,
|
|
134
|
+
initializer=initializers.zeros(),
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
# Initialize weights for each attention head
|
|
138
|
+
for head in range(self.attn_heads):
|
|
139
|
+
# Layer kernel
|
|
140
|
+
kernel = self.add_weight(
|
|
141
|
+
shape=(input_dim, self.units),
|
|
142
|
+
initializer=self.kernel_initializer,
|
|
143
|
+
regularizer=self.kernel_regularizer,
|
|
144
|
+
constraint=self.kernel_constraint,
|
|
145
|
+
name="kernel_{}".format(head),
|
|
146
|
+
)
|
|
147
|
+
self.kernels.append(kernel)
|
|
148
|
+
|
|
149
|
+
# # Layer bias
|
|
150
|
+
if self.use_bias:
|
|
151
|
+
bias = self.add_weight(
|
|
152
|
+
shape=(self.units,),
|
|
153
|
+
initializer=self.bias_initializer,
|
|
154
|
+
regularizer=self.bias_regularizer,
|
|
155
|
+
constraint=self.bias_constraint,
|
|
156
|
+
name="bias_{}".format(head),
|
|
157
|
+
)
|
|
158
|
+
self.biases.append(bias)
|
|
159
|
+
|
|
160
|
+
# Attention kernels
|
|
161
|
+
attn_kernel_self = self.add_weight(
|
|
162
|
+
shape=(self.units, 1),
|
|
163
|
+
initializer=self.attn_kernel_initializer,
|
|
164
|
+
regularizer=self.attn_kernel_regularizer,
|
|
165
|
+
constraint=self.attn_kernel_constraint,
|
|
166
|
+
name="attn_kernel_self_{}".format(head),
|
|
167
|
+
)
|
|
168
|
+
attn_kernel_neighs = self.add_weight(
|
|
169
|
+
shape=(self.units, 1),
|
|
170
|
+
initializer=self.attn_kernel_initializer,
|
|
171
|
+
regularizer=self.attn_kernel_regularizer,
|
|
172
|
+
constraint=self.attn_kernel_constraint,
|
|
173
|
+
name="attn_kernel_neigh_{}".format(head),
|
|
174
|
+
)
|
|
175
|
+
self.attn_kernels.append([attn_kernel_self, attn_kernel_neighs])
|
|
176
|
+
self.built = True
|
|
177
|
+
|
|
178
|
+
def call(self, inputs, A=None, training=None, **kwargs):
|
|
179
|
+
if A is not None:
|
|
180
|
+
X = inputs
|
|
181
|
+
elif isinstance(inputs, (list, tuple)):
|
|
182
|
+
X = inputs[0]
|
|
183
|
+
A = inputs[1]
|
|
184
|
+
else:
|
|
185
|
+
X, A = inputs, None
|
|
186
|
+
|
|
187
|
+
if A is not None and hasattr(A, "shape") and len(A.shape) == 2 and A.shape[0] == 2 and A.shape[1] != 2:
|
|
188
|
+
num_nodes = ops.shape(X)[-2]
|
|
189
|
+
a_dense = ops.zeros((num_nodes, num_nodes), dtype=X.dtype)
|
|
190
|
+
indices = ops.transpose(A, axes=[1, 0])
|
|
191
|
+
updates = ops.ones(shape=(ops.shape(A)[1],), dtype=X.dtype)
|
|
192
|
+
A = ops.scatter_update(a_dense, indices, updates)
|
|
193
|
+
|
|
194
|
+
assert len(ops.shape(A)) == 2, f"Adjacency matrix A should be 2-D"
|
|
195
|
+
N = ops.shape(A)[-1]
|
|
196
|
+
|
|
197
|
+
outputs = []
|
|
198
|
+
for head in range(self.attn_heads):
|
|
199
|
+
kernel = self.kernels[head] # W in the paper (F x F')
|
|
200
|
+
attention_kernel = self.attn_kernels[
|
|
201
|
+
head
|
|
202
|
+
] # Attention kernel a in the paper (2F' x 1)
|
|
203
|
+
|
|
204
|
+
# Compute inputs to attention network
|
|
205
|
+
|
|
206
|
+
features = ops.dot(X, kernel) # (N x F')
|
|
207
|
+
|
|
208
|
+
# Compute feature combinations
|
|
209
|
+
# Note: [[a_1], [a_2]]^T [[Wh_i], [Wh_2]] = [a_1]^T [Wh_i] + [a_2]^T [Wh_j]
|
|
210
|
+
attn_for_self = ops.dot(
|
|
211
|
+
features, attention_kernel[0]
|
|
212
|
+
) # (N x 1), [a_1]^T [Wh_i]
|
|
213
|
+
attn_for_neighs = ops.dot(
|
|
214
|
+
features, attention_kernel[1]
|
|
215
|
+
) # (N x 1), [a_2]^T [Wh_j]
|
|
216
|
+
|
|
217
|
+
# Attention head a(Wh_i, Wh_j) = a^T [[Wh_i], [Wh_j]]
|
|
218
|
+
dense = attn_for_self + ops.transpose(
|
|
219
|
+
attn_for_neighs
|
|
220
|
+
) # (N x N) via broadcasting
|
|
221
|
+
|
|
222
|
+
dense = LeakyReLU(0.2)(dense)
|
|
223
|
+
|
|
224
|
+
if not self.saliency_map_support:
|
|
225
|
+
mask = -10e9 * (1.0 - A)
|
|
226
|
+
dense += mask
|
|
227
|
+
dense = ops.softmax(dense) # (N x N), Eq. 3 of the paper
|
|
228
|
+
|
|
229
|
+
else:
|
|
230
|
+
# dense = dense - tf.reduce_max(dense)
|
|
231
|
+
# GAT with support for saliency calculations
|
|
232
|
+
W = (self.delta * A) * ops.exp(
|
|
233
|
+
dense - ops.max(dense, axis=1, keepdims=True)
|
|
234
|
+
) * (1 - self.non_exist_edge) + self.non_exist_edge * (
|
|
235
|
+
A + self.delta * (ops.ones((N, N)) - A) + ops.eye(N)
|
|
236
|
+
) * ops.exp(
|
|
237
|
+
dense - ops.max(dense, axis=1, keepdims=True)
|
|
238
|
+
)
|
|
239
|
+
dense = W / ops.sum(W, axis=1, keepdims=True)
|
|
240
|
+
|
|
241
|
+
# Apply dropout to features and attention coefficients
|
|
242
|
+
dropout_feat = self.in_dropout(features, training=training) # (N x F')
|
|
243
|
+
dropout_attn = self.attn_dropout(dense, training=training) # (N x N)
|
|
244
|
+
|
|
245
|
+
# Linear combination with neighbors' features [YT: see Eq. 4]
|
|
246
|
+
node_features = ops.dot(dropout_attn, dropout_feat) # (N x F')
|
|
247
|
+
|
|
248
|
+
if self.use_bias:
|
|
249
|
+
node_features = ops.add(node_features, self.biases[head])
|
|
250
|
+
|
|
251
|
+
# Add output of attention head to final output
|
|
252
|
+
outputs.append(node_features)
|
|
253
|
+
|
|
254
|
+
# Aggregate the heads' output according to the reduction method
|
|
255
|
+
if self.attn_heads_reduction == "concat":
|
|
256
|
+
output = ops.concatenate(outputs, axis=1) # (N x KF')
|
|
257
|
+
else:
|
|
258
|
+
output = ops.mean(ops.stack(outputs), axis=0) # N x F')
|
|
259
|
+
|
|
260
|
+
output = self.activation(output)
|
|
261
|
+
|
|
262
|
+
return output
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
from typing import Union, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class GraphConv(MessagePassing):
|
|
8
|
+
r"""The graph neural network operator from the `"Weisfeiler and Leman Go
|
|
9
|
+
Neural: Higher-order Graph Neural Networks"
|
|
10
|
+
<https://arxiv.org/abs/1810.02244>`_ paper.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
14
|
+
out_channels: Size of each output sample.
|
|
15
|
+
aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
|
|
16
|
+
(default: ``"add"``)
|
|
17
|
+
bias: If set to :obj:`False`, the layer will not learn
|
|
18
|
+
an additive bias. (default: ``"True"``)
|
|
19
|
+
|
|
20
|
+
Example:
|
|
21
|
+
```python
|
|
22
|
+
import numpy as np
|
|
23
|
+
from k3_node.layers import GraphConv
|
|
24
|
+
|
|
25
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
26
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
27
|
+
|
|
28
|
+
layer = GraphConv(in_channels=8, out_channels=16)
|
|
29
|
+
out = layer(x, edge_index)
|
|
30
|
+
print(tuple(out.shape)) # (10, 16)
|
|
31
|
+
```
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
weighted_sum_message = True
|
|
35
|
+
|
|
36
|
+
def __init__(
|
|
37
|
+
self,
|
|
38
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
39
|
+
out_channels: int,
|
|
40
|
+
aggr: str = "add",
|
|
41
|
+
bias: bool = True,
|
|
42
|
+
**kwargs,
|
|
43
|
+
):
|
|
44
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
45
|
+
self.in_channels = in_channels
|
|
46
|
+
self.out_channels = out_channels
|
|
47
|
+
self.use_bias = bias
|
|
48
|
+
|
|
49
|
+
self.lin_rel = layers.Dense(out_channels, use_bias=bias)
|
|
50
|
+
self.lin_root = layers.Dense(out_channels, use_bias=False)
|
|
51
|
+
|
|
52
|
+
def build(self, input_shape):
|
|
53
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
54
|
+
in_channels_l = input_shape[0][-1]
|
|
55
|
+
in_channels_r = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_l
|
|
56
|
+
else:
|
|
57
|
+
in_channels_l = input_shape[-1]
|
|
58
|
+
in_channels_r = input_shape[-1]
|
|
59
|
+
|
|
60
|
+
self.lin_rel.build((None, in_channels_l))
|
|
61
|
+
self.lin_root.build((None, in_channels_r))
|
|
62
|
+
self.built = True
|
|
63
|
+
|
|
64
|
+
def call(self, x, edge_index=None, edge_weight=None, size=None, **kwargs):
|
|
65
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
66
|
+
x, edge_index = x[0], x[1]
|
|
67
|
+
|
|
68
|
+
if not isinstance(x, (tuple, list)):
|
|
69
|
+
x_src = x
|
|
70
|
+
x_dst = x
|
|
71
|
+
else:
|
|
72
|
+
x_src, x_dst = x[0], x[1]
|
|
73
|
+
|
|
74
|
+
out = self.propagate(edge_index, x=(x_src, x_dst), edge_weight=edge_weight, size=size)
|
|
75
|
+
out = self.lin_rel(out)
|
|
76
|
+
if x_dst is not None:
|
|
77
|
+
out = out + self.lin_root(x_dst)
|
|
78
|
+
return out
|
|
79
|
+
|
|
80
|
+
def message(self, x_j, edge_weight=None):
|
|
81
|
+
if edge_weight is None:
|
|
82
|
+
return x_j
|
|
83
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
84
|
+
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
from typing import Optional, Union, Tuple
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
from keras.layers import Dense
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
from k3_node.layers.pool.knn import knn_graph
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class GravNetConv(MessagePassing):
|
|
11
|
+
r"""The GravNet operator from the `"Learning Representations of Irregular
|
|
12
|
+
Particle-Detector Geometry with Distance-Weighted Graph Networks"
|
|
13
|
+
<https://arxiv.org/abs/1902.07987>`_ paper.
|
|
14
|
+
|
|
15
|
+
Example:
|
|
16
|
+
```python
|
|
17
|
+
import numpy as np
|
|
18
|
+
from k3_node.layers import GravNetConv
|
|
19
|
+
|
|
20
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
21
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
22
|
+
|
|
23
|
+
# Neighbors are found by k-NN in a learned space, so no edge_index is needed
|
|
24
|
+
layer = GravNetConv(in_channels=8, out_channels=16, space_dimensions=3, propagate_dimensions=4, k=3)
|
|
25
|
+
out = layer(x)
|
|
26
|
+
print(tuple(out.shape)) # (10, 16)
|
|
27
|
+
```
|
|
28
|
+
"""
|
|
29
|
+
def __init__(
|
|
30
|
+
self,
|
|
31
|
+
in_channels: int,
|
|
32
|
+
out_channels: int,
|
|
33
|
+
space_dimensions: int,
|
|
34
|
+
propagate_dimensions: int,
|
|
35
|
+
k: int,
|
|
36
|
+
num_workers: Optional[int] = None,
|
|
37
|
+
**kwargs,
|
|
38
|
+
):
|
|
39
|
+
kwargs.setdefault("aggr", "mean")
|
|
40
|
+
super().__init__(node_dim=0, **kwargs)
|
|
41
|
+
|
|
42
|
+
self.in_channels = in_channels
|
|
43
|
+
self.out_channels = out_channels
|
|
44
|
+
self.space_dimensions = space_dimensions
|
|
45
|
+
self.propagate_dimensions = propagate_dimensions
|
|
46
|
+
self.k = k
|
|
47
|
+
|
|
48
|
+
self.lin_s = Dense(space_dimensions, use_bias=True)
|
|
49
|
+
self.lin_h = Dense(propagate_dimensions, use_bias=True)
|
|
50
|
+
self.lin_out1 = Dense(out_channels, use_bias=True)
|
|
51
|
+
self.lin_out2 = Dense(out_channels, use_bias=True)
|
|
52
|
+
|
|
53
|
+
def build(self, input_shape=None):
|
|
54
|
+
self.lin_s.build((None, self.in_channels))
|
|
55
|
+
self.lin_h.build((None, self.in_channels))
|
|
56
|
+
self.lin_out1.build((None, self.in_channels))
|
|
57
|
+
self.lin_out2.build((None, self.propagate_dimensions))
|
|
58
|
+
self.built = True
|
|
59
|
+
|
|
60
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
61
|
+
if isinstance(inputs, (list, tuple)):
|
|
62
|
+
x = inputs[0]
|
|
63
|
+
else:
|
|
64
|
+
x = inputs
|
|
65
|
+
|
|
66
|
+
if not self.built:
|
|
67
|
+
self.build()
|
|
68
|
+
|
|
69
|
+
s = self.lin_s(x)
|
|
70
|
+
h = self.lin_h(x)
|
|
71
|
+
|
|
72
|
+
if edge_index is None:
|
|
73
|
+
edge_index = knn_graph(s, k=self.k)
|
|
74
|
+
|
|
75
|
+
edge_index = ops.cast(edge_index, "int32")
|
|
76
|
+
s_src = ops.take(s, edge_index[0], axis=0)
|
|
77
|
+
s_dst = ops.take(s, edge_index[1], axis=0)
|
|
78
|
+
dist_sq = ops.sum(ops.square(s_src - s_dst), axis=-1)
|
|
79
|
+
edge_weight = ops.exp(-10.0 * dist_sq)
|
|
80
|
+
|
|
81
|
+
num_nodes = ops.shape(x)[0]
|
|
82
|
+
out = self.propagate(
|
|
83
|
+
edge_index,
|
|
84
|
+
x=h,
|
|
85
|
+
edge_weight=edge_weight,
|
|
86
|
+
size=(num_nodes, num_nodes),
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
return self.lin_out1(x) + self.lin_out2(out)
|
|
90
|
+
|
|
91
|
+
def message(self, x_j, edge_weight):
|
|
92
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
93
|
+
|