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,244 @@
|
|
|
1
|
+
from typing import Optional, Union, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import (
|
|
6
|
+
add_self_loops,
|
|
7
|
+
extend_mask_for_self_loops,
|
|
8
|
+
mask_edge_logits,
|
|
9
|
+
remove_self_loops_masked,
|
|
10
|
+
softmax,
|
|
11
|
+
)
|
|
12
|
+
from k3_node.ops.segment import segment_sum
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class GATConv(MessagePassing):
|
|
16
|
+
r"""The graph attentional operator from the `"Graph Attention Networks"
|
|
17
|
+
<https://arxiv.org/abs/1710.10903>`_ paper.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
21
|
+
out_channels: Size of each output sample.
|
|
22
|
+
heads: Number of multi-head-attentions. (default: ``1``)
|
|
23
|
+
concat: If set to :obj:`False`, the multi-head-attentions are averaged
|
|
24
|
+
instead of concatenated. (default: ``True``)
|
|
25
|
+
negative_slope: LeakyReLU angle of the negative slope. (default: ``0.2``)
|
|
26
|
+
dropout: Dropout probability of the normalized attention coefficients.
|
|
27
|
+
(default: ``0.0``)
|
|
28
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops to
|
|
29
|
+
the input graph. (default: ``True``)
|
|
30
|
+
edge_dim: Edge feature dimensionality (in case there are any).
|
|
31
|
+
(default: :obj:`None`)
|
|
32
|
+
fill_value: The way to generate edge features of self-loops
|
|
33
|
+
(default: ``"mean"``)
|
|
34
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
35
|
+
(default: ``True``)
|
|
36
|
+
share_weights: If set to :obj:`True`, the same matrix will be applied
|
|
37
|
+
to the source and target node features. (default: ``False``)
|
|
38
|
+
residual: If set to :obj:`True`, will compute residual connections.
|
|
39
|
+
(default: ``False``)
|
|
40
|
+
|
|
41
|
+
Example:
|
|
42
|
+
```python
|
|
43
|
+
import numpy as np
|
|
44
|
+
from k3_node.layers import GATConv
|
|
45
|
+
|
|
46
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
47
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
48
|
+
|
|
49
|
+
layer = GATConv(in_channels=8, out_channels=16, heads=2)
|
|
50
|
+
out = layer(x, edge_index)
|
|
51
|
+
print(tuple(out.shape)) # (10, 32)
|
|
52
|
+
```
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
def __init__(
|
|
56
|
+
self,
|
|
57
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
58
|
+
out_channels: int,
|
|
59
|
+
heads: int = 1,
|
|
60
|
+
concat: bool = True,
|
|
61
|
+
negative_slope: float = 0.2,
|
|
62
|
+
dropout: float = 0.0,
|
|
63
|
+
add_self_loops: bool = True,
|
|
64
|
+
edge_dim: Optional[int] = None,
|
|
65
|
+
fill_value: Union[float, str] = "mean",
|
|
66
|
+
bias: bool = True,
|
|
67
|
+
share_weights: bool = False,
|
|
68
|
+
residual: bool = False,
|
|
69
|
+
**kwargs,
|
|
70
|
+
):
|
|
71
|
+
super().__init__(node_dim=0, **kwargs)
|
|
72
|
+
self.in_channels = in_channels
|
|
73
|
+
self.out_channels = out_channels
|
|
74
|
+
self.heads = heads
|
|
75
|
+
self.concat = concat
|
|
76
|
+
self.negative_slope = negative_slope
|
|
77
|
+
self.dropout_rate = dropout
|
|
78
|
+
self.add_self_loops = add_self_loops
|
|
79
|
+
self.edge_dim = edge_dim
|
|
80
|
+
self.fill_value = fill_value
|
|
81
|
+
self.use_bias = bias
|
|
82
|
+
self.share_weights = share_weights
|
|
83
|
+
self.residual = residual
|
|
84
|
+
|
|
85
|
+
total_out_channels = out_channels * (heads if concat else 1)
|
|
86
|
+
|
|
87
|
+
if isinstance(in_channels, int):
|
|
88
|
+
self.lin = layers.Dense(heads * out_channels, use_bias=False)
|
|
89
|
+
self.lin_src = self.lin
|
|
90
|
+
self.lin_dst = self.lin
|
|
91
|
+
else:
|
|
92
|
+
self.lin = None
|
|
93
|
+
self.lin_src = layers.Dense(heads * out_channels, use_bias=False)
|
|
94
|
+
if share_weights:
|
|
95
|
+
self.lin_dst = self.lin_src
|
|
96
|
+
else:
|
|
97
|
+
self.lin_dst = layers.Dense(heads * out_channels, use_bias=False)
|
|
98
|
+
|
|
99
|
+
if edge_dim is not None:
|
|
100
|
+
self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
|
|
101
|
+
else:
|
|
102
|
+
self.lin_edge = None
|
|
103
|
+
|
|
104
|
+
if residual:
|
|
105
|
+
self.res = layers.Dense(total_out_channels, use_bias=False)
|
|
106
|
+
else:
|
|
107
|
+
self.res = None
|
|
108
|
+
|
|
109
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
110
|
+
|
|
111
|
+
def build(self, input_shape):
|
|
112
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
113
|
+
in_channels_src = input_shape[0][-1]
|
|
114
|
+
in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
|
|
115
|
+
else:
|
|
116
|
+
in_channels_src = input_shape[-1]
|
|
117
|
+
in_channels_dst = input_shape[-1]
|
|
118
|
+
|
|
119
|
+
self.lin_src.build((None, in_channels_src))
|
|
120
|
+
if self.lin_dst is not self.lin_src:
|
|
121
|
+
self.lin_dst.build((None, in_channels_dst))
|
|
122
|
+
if self.lin_edge is not None:
|
|
123
|
+
self.lin_edge.build((None, self.edge_dim))
|
|
124
|
+
if self.res is not None:
|
|
125
|
+
self.res.build((None, in_channels_dst))
|
|
126
|
+
|
|
127
|
+
self.att_src = self.add_weight(
|
|
128
|
+
shape=(1, self.heads, self.out_channels),
|
|
129
|
+
initializer="glorot_uniform",
|
|
130
|
+
name="att_src",
|
|
131
|
+
)
|
|
132
|
+
self.att_dst = self.add_weight(
|
|
133
|
+
shape=(1, self.heads, self.out_channels),
|
|
134
|
+
initializer="glorot_uniform",
|
|
135
|
+
name="att_dst",
|
|
136
|
+
)
|
|
137
|
+
if self.edge_dim is not None:
|
|
138
|
+
self.att_edge = self.add_weight(
|
|
139
|
+
shape=(1, self.heads, self.out_channels),
|
|
140
|
+
initializer="glorot_uniform",
|
|
141
|
+
name="att_edge",
|
|
142
|
+
)
|
|
143
|
+
else:
|
|
144
|
+
self.att_edge = None
|
|
145
|
+
|
|
146
|
+
total_out_channels = self.out_channels * (self.heads if self.concat else 1)
|
|
147
|
+
if self.use_bias:
|
|
148
|
+
self.bias = self.add_weight(
|
|
149
|
+
shape=(total_out_channels,),
|
|
150
|
+
initializer="zeros",
|
|
151
|
+
name="bias",
|
|
152
|
+
)
|
|
153
|
+
else:
|
|
154
|
+
self.bias = None
|
|
155
|
+
self.built = True
|
|
156
|
+
|
|
157
|
+
def call(self, x, edge_index=None, edge_attr=None, size=None, return_attention_weights=None, training=None, **kwargs):
|
|
158
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
159
|
+
x, edge_index = x[0], x[1]
|
|
160
|
+
|
|
161
|
+
H, C = self.heads, self.out_channels
|
|
162
|
+
if isinstance(x, (tuple, list)):
|
|
163
|
+
x_src, x_dst = x[0], x[1]
|
|
164
|
+
else:
|
|
165
|
+
x_src, x_dst = x, x
|
|
166
|
+
|
|
167
|
+
x_src_proj = ops.reshape(self.lin_src(x_src), (-1, H, C))
|
|
168
|
+
x_dst_proj = ops.reshape(self.lin_dst(x_dst), (-1, H, C)) if x_dst is not None else None
|
|
169
|
+
|
|
170
|
+
alpha_src = ops.sum(x_src_proj * self.att_src, axis=-1)
|
|
171
|
+
alpha_dst = ops.sum(x_dst_proj * self.att_dst, axis=-1) if x_dst_proj is not None else None
|
|
172
|
+
|
|
173
|
+
if self.add_self_loops:
|
|
174
|
+
if not isinstance(x, (tuple, list)):
|
|
175
|
+
num_nodes = ops.shape(x)[0]
|
|
176
|
+
else:
|
|
177
|
+
num_nodes = ops.shape(x_src)[0]
|
|
178
|
+
if x_dst is not None:
|
|
179
|
+
num_nodes = ops.minimum(num_nodes, ops.shape(x_dst)[0])
|
|
180
|
+
edge_index, edge_attr, keep_mask = remove_self_loops_masked(edge_index, edge_attr)
|
|
181
|
+
edge_index, edge_attr = add_self_loops(
|
|
182
|
+
edge_index, edge_attr, fill_value=self.fill_value, num_nodes=num_nodes
|
|
183
|
+
)
|
|
184
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
185
|
+
else:
|
|
186
|
+
keep_mask = None
|
|
187
|
+
|
|
188
|
+
row, col = edge_index[0], edge_index[1]
|
|
189
|
+
row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
|
|
190
|
+
|
|
191
|
+
alpha_j = ops.take(alpha_src, row, axis=0)
|
|
192
|
+
alpha_i = ops.take(alpha_dst, col, axis=0) if alpha_dst is not None else 0.0
|
|
193
|
+
alpha = alpha_j + alpha_i
|
|
194
|
+
|
|
195
|
+
if edge_attr is not None and self.lin_edge is not None and self.att_edge is not None:
|
|
196
|
+
edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
|
|
197
|
+
alpha = alpha + ops.sum(edge_attr_proj * self.att_edge, axis=-1)
|
|
198
|
+
|
|
199
|
+
alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
|
|
200
|
+
alpha = mask_edge_logits(alpha, keep_mask)
|
|
201
|
+
num_nodes_dst = ops.shape(x_dst)[0] if x_dst is not None else ops.shape(x_src)[0]
|
|
202
|
+
alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
|
|
203
|
+
|
|
204
|
+
if self.dropout is not None:
|
|
205
|
+
alpha = self.dropout(alpha, training=training)
|
|
206
|
+
|
|
207
|
+
# Message & aggregate
|
|
208
|
+
x_src_j = ops.take(x_src_proj, row, axis=0)
|
|
209
|
+
out = ops.expand_dims(alpha, -1) * x_src_j
|
|
210
|
+
out = segment_sum(out, col, num_segments=num_nodes_dst)
|
|
211
|
+
|
|
212
|
+
if self.concat:
|
|
213
|
+
out = ops.reshape(out, (-1, H * C))
|
|
214
|
+
else:
|
|
215
|
+
out = ops.mean(out, axis=1)
|
|
216
|
+
|
|
217
|
+
if self.res is not None and x_dst is not None:
|
|
218
|
+
out = out + self.res(x_dst)
|
|
219
|
+
|
|
220
|
+
if self.bias is not None:
|
|
221
|
+
out = out + self.bias
|
|
222
|
+
|
|
223
|
+
if return_attention_weights:
|
|
224
|
+
return out, (edge_index, alpha)
|
|
225
|
+
return out
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
class FusedGATConv(GATConv):
|
|
229
|
+
r"""The fused graph attentional operator.
|
|
230
|
+
|
|
231
|
+
Example:
|
|
232
|
+
```python
|
|
233
|
+
import numpy as np
|
|
234
|
+
from k3_node.layers import FusedGATConv
|
|
235
|
+
|
|
236
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
237
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
238
|
+
|
|
239
|
+
layer = FusedGATConv(in_channels=8, out_channels=16)
|
|
240
|
+
out = layer(x, edge_index)
|
|
241
|
+
print(tuple(out.shape)) # (10, 16)
|
|
242
|
+
```
|
|
243
|
+
"""
|
|
244
|
+
pass
|
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
# ported from spektral
|
|
2
|
+
|
|
3
|
+
from keras import ops
|
|
4
|
+
from keras.layers import GRUCell
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class GatedGraphConv(MessagePassing):
|
|
10
|
+
"""
|
|
11
|
+
`k3_node.layers.GatedGraphConv`
|
|
12
|
+
|
|
13
|
+
Implementation of Gated Graph Convolution (GGC) layer
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
channels: The number of output channels.
|
|
17
|
+
n_layers: The number of GGC layers to stack.
|
|
18
|
+
activation: Activation function to use.
|
|
19
|
+
use_bias: Whether to add a bias to the linear transformation.
|
|
20
|
+
kernel_initializer: Initializer for the `kernel` weights matrix.
|
|
21
|
+
bias_initializer: Initializer for the bias vector.
|
|
22
|
+
kernel_regularizer: Regularizer for the `kernel` weights matrix.
|
|
23
|
+
bias_regularizer: Regularizer for the bias vector.
|
|
24
|
+
activity_regularizer: Regularizer for the output.
|
|
25
|
+
kernel_constraint: Constraint for the `kernel` weights matrix.
|
|
26
|
+
bias_constraint: Constraint for the bias vector.
|
|
27
|
+
**kwargs: Additional arguments to pass to the `MessagePassing` superclass.
|
|
28
|
+
|
|
29
|
+
Example:
|
|
30
|
+
```python
|
|
31
|
+
import numpy as np
|
|
32
|
+
from k3_node.layers import GatedGraphConv
|
|
33
|
+
|
|
34
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
35
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
36
|
+
|
|
37
|
+
layer = GatedGraphConv(out_channels=16, num_layers=2)
|
|
38
|
+
out = layer(x, edge_index)
|
|
39
|
+
print(tuple(out.shape)) # (10, 16)
|
|
40
|
+
```
|
|
41
|
+
"""
|
|
42
|
+
def __init__(
|
|
43
|
+
self,
|
|
44
|
+
channels=None,
|
|
45
|
+
n_layers=None,
|
|
46
|
+
out_channels=None,
|
|
47
|
+
num_layers=None,
|
|
48
|
+
activation=None,
|
|
49
|
+
use_bias=True,
|
|
50
|
+
kernel_initializer="glorot_uniform",
|
|
51
|
+
bias_initializer="zeros",
|
|
52
|
+
kernel_regularizer=None,
|
|
53
|
+
bias_regularizer=None,
|
|
54
|
+
activity_regularizer=None,
|
|
55
|
+
kernel_constraint=None,
|
|
56
|
+
bias_constraint=None,
|
|
57
|
+
**kwargs,
|
|
58
|
+
):
|
|
59
|
+
channels = out_channels if out_channels is not None else channels
|
|
60
|
+
n_layers = num_layers if num_layers is not None else n_layers
|
|
61
|
+
super().__init__(
|
|
62
|
+
activation=activation,
|
|
63
|
+
use_bias=use_bias,
|
|
64
|
+
kernel_initializer=kernel_initializer,
|
|
65
|
+
bias_initializer=bias_initializer,
|
|
66
|
+
kernel_regularizer=kernel_regularizer,
|
|
67
|
+
bias_regularizer=bias_regularizer,
|
|
68
|
+
activity_regularizer=activity_regularizer,
|
|
69
|
+
kernel_constraint=kernel_constraint,
|
|
70
|
+
bias_constraint=bias_constraint,
|
|
71
|
+
**kwargs,
|
|
72
|
+
)
|
|
73
|
+
self.channels = channels
|
|
74
|
+
self.out_channels = channels
|
|
75
|
+
self.n_layers = n_layers
|
|
76
|
+
self.num_layers = n_layers
|
|
77
|
+
|
|
78
|
+
def build(self, input_shape=None):
|
|
79
|
+
self.kernel = self.add_weight(
|
|
80
|
+
name="kernel",
|
|
81
|
+
shape=(self.n_layers, self.channels, self.channels),
|
|
82
|
+
initializer=self.kernel_initializer,
|
|
83
|
+
regularizer=self.kernel_regularizer,
|
|
84
|
+
constraint=self.kernel_constraint,
|
|
85
|
+
)
|
|
86
|
+
self.rnn = GRUCell(
|
|
87
|
+
self.channels,
|
|
88
|
+
kernel_initializer=self.kernel_initializer,
|
|
89
|
+
bias_initializer=self.bias_initializer,
|
|
90
|
+
kernel_regularizer=self.kernel_regularizer,
|
|
91
|
+
bias_regularizer=self.bias_regularizer,
|
|
92
|
+
activity_regularizer=self.activity_regularizer,
|
|
93
|
+
kernel_constraint=self.kernel_constraint,
|
|
94
|
+
bias_constraint=self.bias_constraint,
|
|
95
|
+
use_bias=self.use_bias,
|
|
96
|
+
dtype=self.dtype,
|
|
97
|
+
)
|
|
98
|
+
self.rnn.build((self.channels,))
|
|
99
|
+
super().build(input_shape)
|
|
100
|
+
self.built = True
|
|
101
|
+
|
|
102
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
103
|
+
is_legacy = False
|
|
104
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
105
|
+
x, a, _ = self.get_inputs(x)
|
|
106
|
+
edge_index = a
|
|
107
|
+
is_legacy = True
|
|
108
|
+
|
|
109
|
+
F = ops.shape(x)[-1]
|
|
110
|
+
if F < self.channels:
|
|
111
|
+
to_pad = self.channels - F
|
|
112
|
+
ndims = len(ops.shape(x)) - 1
|
|
113
|
+
output = ops.pad(x, [[0, 0]] * ndims + [[0, to_pad]])
|
|
114
|
+
elif F > self.channels:
|
|
115
|
+
output = x[..., :self.channels]
|
|
116
|
+
else:
|
|
117
|
+
output = x
|
|
118
|
+
|
|
119
|
+
for i in range(self.n_layers):
|
|
120
|
+
m = ops.matmul(output, self.kernel[i])
|
|
121
|
+
if is_legacy:
|
|
122
|
+
m = self.propagate(m, edge_index)
|
|
123
|
+
else:
|
|
124
|
+
m = self.propagate(edge_index, x=m, edge_weight=edge_weight)
|
|
125
|
+
output = self.rnn(m, [output])[0]
|
|
126
|
+
|
|
127
|
+
if hasattr(self, "activation") and self.activation is not None and callable(self.activation):
|
|
128
|
+
output = self.activation(output)
|
|
129
|
+
return output
|
|
130
|
+
|
|
131
|
+
@property
|
|
132
|
+
def config(self):
|
|
133
|
+
return {
|
|
134
|
+
"channels": self.channels,
|
|
135
|
+
"n_layers": self.n_layers,
|
|
136
|
+
}
|
|
@@ -0,0 +1,205 @@
|
|
|
1
|
+
from typing import Optional, Union, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import (
|
|
6
|
+
add_self_loops,
|
|
7
|
+
extend_mask_for_self_loops,
|
|
8
|
+
mask_edge_logits,
|
|
9
|
+
remove_self_loops_masked,
|
|
10
|
+
softmax,
|
|
11
|
+
)
|
|
12
|
+
from k3_node.ops.segment import segment_sum
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class GATv2Conv(MessagePassing):
|
|
16
|
+
r"""The GATv2 operator from the `"How Attentive are Graph Attention Networks?"
|
|
17
|
+
<https://arxiv.org/abs/2105.14491>`_ paper, which fixes the static
|
|
18
|
+
attention problem of standard :class:`~k3_node.layers.conv.GATConv`.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
22
|
+
out_channels: Size of each output sample.
|
|
23
|
+
heads: Number of multi-head-attentions. (default: ``1``)
|
|
24
|
+
concat: If set to :obj:`False`, the multi-head-attentions are averaged
|
|
25
|
+
instead of concatenated. (default: ``True``)
|
|
26
|
+
negative_slope: LeakyReLU angle of the negative slope. (default: ``0.2``)
|
|
27
|
+
dropout: Dropout probability of the normalized attention coefficients.
|
|
28
|
+
(default: ``0.0``)
|
|
29
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops to
|
|
30
|
+
the input graph. (default: ``True``)
|
|
31
|
+
edge_dim: Edge feature dimensionality (in case there are any).
|
|
32
|
+
(default: :obj:`None`)
|
|
33
|
+
fill_value: The way to generate edge features of self-loops
|
|
34
|
+
(default: ``"mean"``)
|
|
35
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
36
|
+
(default: ``True``)
|
|
37
|
+
share_weights: If set to :obj:`True`, the same matrix will be applied
|
|
38
|
+
to the source and target node features. (default: ``False``)
|
|
39
|
+
residual: If set to :obj:`True`, will compute residual connections.
|
|
40
|
+
(default: ``False``)
|
|
41
|
+
|
|
42
|
+
Example:
|
|
43
|
+
```python
|
|
44
|
+
import numpy as np
|
|
45
|
+
from k3_node.layers import GATv2Conv
|
|
46
|
+
|
|
47
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
48
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
49
|
+
|
|
50
|
+
layer = GATv2Conv(in_channels=8, out_channels=16, heads=2)
|
|
51
|
+
out = layer(x, edge_index)
|
|
52
|
+
print(tuple(out.shape)) # (10, 32)
|
|
53
|
+
```
|
|
54
|
+
"""
|
|
55
|
+
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
59
|
+
out_channels: int,
|
|
60
|
+
heads: int = 1,
|
|
61
|
+
concat: bool = True,
|
|
62
|
+
negative_slope: float = 0.2,
|
|
63
|
+
dropout: float = 0.0,
|
|
64
|
+
add_self_loops: bool = True,
|
|
65
|
+
edge_dim: Optional[int] = None,
|
|
66
|
+
fill_value: Union[float, str] = "mean",
|
|
67
|
+
bias: bool = True,
|
|
68
|
+
share_weights: bool = False,
|
|
69
|
+
residual: bool = False,
|
|
70
|
+
**kwargs,
|
|
71
|
+
):
|
|
72
|
+
super().__init__(node_dim=0, **kwargs)
|
|
73
|
+
self.in_channels = in_channels
|
|
74
|
+
self.out_channels = out_channels
|
|
75
|
+
self.heads = heads
|
|
76
|
+
self.concat = concat
|
|
77
|
+
self.negative_slope = negative_slope
|
|
78
|
+
self.dropout_rate = dropout
|
|
79
|
+
self.add_self_loops = add_self_loops
|
|
80
|
+
self.edge_dim = edge_dim
|
|
81
|
+
self.fill_value = fill_value
|
|
82
|
+
self.use_bias = bias
|
|
83
|
+
self.share_weights = share_weights
|
|
84
|
+
self.residual = residual
|
|
85
|
+
|
|
86
|
+
total_out_channels = out_channels * (heads if concat else 1)
|
|
87
|
+
|
|
88
|
+
self.lin_l = layers.Dense(heads * out_channels, use_bias=bias)
|
|
89
|
+
if share_weights:
|
|
90
|
+
self.lin_r = self.lin_l
|
|
91
|
+
else:
|
|
92
|
+
self.lin_r = layers.Dense(heads * out_channels, use_bias=bias)
|
|
93
|
+
|
|
94
|
+
if edge_dim is not None:
|
|
95
|
+
self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
|
|
96
|
+
else:
|
|
97
|
+
self.lin_edge = None
|
|
98
|
+
|
|
99
|
+
if residual:
|
|
100
|
+
self.res = layers.Dense(total_out_channels, use_bias=False)
|
|
101
|
+
else:
|
|
102
|
+
self.res = None
|
|
103
|
+
|
|
104
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
105
|
+
|
|
106
|
+
def build(self, input_shape):
|
|
107
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
108
|
+
in_channels_src = input_shape[0][-1]
|
|
109
|
+
in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
|
|
110
|
+
else:
|
|
111
|
+
in_channels_src = input_shape[-1]
|
|
112
|
+
in_channels_dst = input_shape[-1]
|
|
113
|
+
|
|
114
|
+
self.lin_l.build((None, in_channels_src))
|
|
115
|
+
if self.lin_r is not self.lin_l:
|
|
116
|
+
self.lin_r.build((None, in_channels_dst))
|
|
117
|
+
if self.lin_edge is not None:
|
|
118
|
+
self.lin_edge.build((None, self.edge_dim))
|
|
119
|
+
if self.res is not None:
|
|
120
|
+
self.res.build((None, in_channels_dst))
|
|
121
|
+
|
|
122
|
+
self.att = self.add_weight(
|
|
123
|
+
shape=(1, self.heads, self.out_channels),
|
|
124
|
+
initializer="glorot_uniform",
|
|
125
|
+
name="att",
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
total_out_channels = self.out_channels * (self.heads if self.concat else 1)
|
|
129
|
+
if self.use_bias:
|
|
130
|
+
self.bias = self.add_weight(
|
|
131
|
+
shape=(total_out_channels,),
|
|
132
|
+
initializer="zeros",
|
|
133
|
+
name="bias",
|
|
134
|
+
)
|
|
135
|
+
else:
|
|
136
|
+
self.bias = None
|
|
137
|
+
self.built = True
|
|
138
|
+
|
|
139
|
+
def call(self, x, edge_index=None, edge_attr=None, size=None, return_attention_weights=None, training=None, **kwargs):
|
|
140
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
141
|
+
x, edge_index = x[0], x[1]
|
|
142
|
+
|
|
143
|
+
H, C = self.heads, self.out_channels
|
|
144
|
+
if isinstance(x, (tuple, list)):
|
|
145
|
+
x_src, x_dst = x[0], x[1]
|
|
146
|
+
else:
|
|
147
|
+
x_src, x_dst = x, x
|
|
148
|
+
|
|
149
|
+
x_l = ops.reshape(self.lin_l(x_src), (-1, H, C))
|
|
150
|
+
x_r = ops.reshape(self.lin_r(x_dst), (-1, H, C)) if x_dst is not None else x_l
|
|
151
|
+
|
|
152
|
+
if self.add_self_loops:
|
|
153
|
+
if not isinstance(x, (tuple, list)):
|
|
154
|
+
num_nodes = ops.shape(x)[0]
|
|
155
|
+
else:
|
|
156
|
+
num_nodes = ops.shape(x_l)[0]
|
|
157
|
+
if x_r is not None:
|
|
158
|
+
num_nodes = ops.minimum(num_nodes, ops.shape(x_r)[0])
|
|
159
|
+
edge_index, edge_attr, keep_mask = remove_self_loops_masked(edge_index, edge_attr)
|
|
160
|
+
edge_index, edge_attr = add_self_loops(
|
|
161
|
+
edge_index, edge_attr, fill_value=self.fill_value, num_nodes=num_nodes
|
|
162
|
+
)
|
|
163
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
164
|
+
else:
|
|
165
|
+
keep_mask = None
|
|
166
|
+
|
|
167
|
+
row, col = edge_index[0], edge_index[1]
|
|
168
|
+
row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
|
|
169
|
+
|
|
170
|
+
x_l_j = ops.take(x_l, row, axis=0)
|
|
171
|
+
x_r_i = ops.take(x_r, col, axis=0)
|
|
172
|
+
alpha = x_l_j + x_r_i
|
|
173
|
+
|
|
174
|
+
if edge_attr is not None and self.lin_edge is not None:
|
|
175
|
+
edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
|
|
176
|
+
alpha = alpha + edge_attr_proj
|
|
177
|
+
|
|
178
|
+
alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
|
|
179
|
+
alpha = ops.sum(alpha * self.att, axis=-1)
|
|
180
|
+
alpha = mask_edge_logits(alpha, keep_mask)
|
|
181
|
+
|
|
182
|
+
num_nodes_dst = ops.shape(x_r)[0]
|
|
183
|
+
alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
|
|
184
|
+
|
|
185
|
+
if self.dropout is not None:
|
|
186
|
+
alpha = self.dropout(alpha, training=training)
|
|
187
|
+
|
|
188
|
+
out = ops.expand_dims(alpha, -1) * x_l_j
|
|
189
|
+
out = segment_sum(out, col, num_segments=num_nodes_dst)
|
|
190
|
+
|
|
191
|
+
if self.concat:
|
|
192
|
+
out = ops.reshape(out, (-1, H * C))
|
|
193
|
+
else:
|
|
194
|
+
out = ops.mean(out, axis=1)
|
|
195
|
+
|
|
196
|
+
if self.res is not None and x_dst is not None:
|
|
197
|
+
out = out + self.res(x_dst)
|
|
198
|
+
|
|
199
|
+
if self.bias is not None:
|
|
200
|
+
out = out + self.bias
|
|
201
|
+
|
|
202
|
+
if return_attention_weights:
|
|
203
|
+
return out, (edge_index, alpha)
|
|
204
|
+
return out
|
|
205
|
+
|