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,144 @@
|
|
|
1
|
+
# ported from stellargraph
|
|
2
|
+
from keras import ops
|
|
3
|
+
from keras import activations, initializers, constraints, regularizers
|
|
4
|
+
from keras.layers import Layer, dot
|
|
5
|
+
from k3_node.ops.creation import repeat
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class GraphConvolution(Layer):
|
|
9
|
+
"""
|
|
10
|
+
`k3_node.layers.GraphConvolution`
|
|
11
|
+
Implementation of Graph Convolution (GCN) layer
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
units: Positive integer, dimensionality of the output space.
|
|
15
|
+
activation: Activation function to use.
|
|
16
|
+
use_bias: Whether to add a bias to the linear transformation.
|
|
17
|
+
final_layer: Deprecated, use tf.gather or GatherIndices instead.
|
|
18
|
+
input_dim: Deprecated, use `keras.layers.Input` with `input_shape` instead.
|
|
19
|
+
kernel_initializer: Initializer for the `kernel` weights matrix.
|
|
20
|
+
kernel_regularizer: Regularizer for the `kernel` weights matrix.
|
|
21
|
+
kernel_constraint: Constraint for the `kernel` weights matrix.
|
|
22
|
+
bias_initializer: Initializer for the bias vector.
|
|
23
|
+
bias_regularizer: Regularizer for the bias vector.
|
|
24
|
+
bias_constraint: Constraint for the bias vector.
|
|
25
|
+
**kwargs: Additional arguments to pass to the `Layer` superclass.
|
|
26
|
+
|
|
27
|
+
Example:
|
|
28
|
+
```python
|
|
29
|
+
import numpy as np
|
|
30
|
+
from k3_node.layers import GraphConvolution
|
|
31
|
+
|
|
32
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
33
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
34
|
+
|
|
35
|
+
layer = GraphConvolution(units=16, activation="relu")
|
|
36
|
+
out = layer(x, edge_index)
|
|
37
|
+
print(tuple(out.shape)) # (10, 16)
|
|
38
|
+
```
|
|
39
|
+
"""
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
units,
|
|
43
|
+
activation=None,
|
|
44
|
+
use_bias=True,
|
|
45
|
+
final_layer=None,
|
|
46
|
+
input_dim=None,
|
|
47
|
+
kernel_initializer="glorot_uniform",
|
|
48
|
+
kernel_regularizer=None,
|
|
49
|
+
kernel_constraint=None,
|
|
50
|
+
bias_initializer="zeros",
|
|
51
|
+
bias_regularizer=None,
|
|
52
|
+
bias_constraint=None,
|
|
53
|
+
**kwargs,
|
|
54
|
+
):
|
|
55
|
+
if isinstance(activation, int):
|
|
56
|
+
# Called as GraphConvolution(in_channels, out_channels)
|
|
57
|
+
self.in_channels = units
|
|
58
|
+
units = activation
|
|
59
|
+
activation = kwargs.pop("activation", None)
|
|
60
|
+
|
|
61
|
+
if "input_shape" not in kwargs and input_dim is not None:
|
|
62
|
+
kwargs["input_shape"] = (input_dim,)
|
|
63
|
+
|
|
64
|
+
self.units = units
|
|
65
|
+
self.activation = activations.get(activation)
|
|
66
|
+
self.use_bias = use_bias
|
|
67
|
+
if final_layer is not None:
|
|
68
|
+
raise ValueError(
|
|
69
|
+
"'final_layer' is not longer supported, use 'tf.gather' or 'GatherIndices' separately"
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
self.kernel_initializer = initializers.get(kernel_initializer)
|
|
73
|
+
self.kernel_regularizer = regularizers.get(kernel_regularizer)
|
|
74
|
+
self.kernel_constraint = constraints.get(kernel_constraint)
|
|
75
|
+
self.bias_initializer = initializers.get(bias_initializer)
|
|
76
|
+
self.bias_regularizer = regularizers.get(bias_regularizer)
|
|
77
|
+
self.bias_constraint = constraints.get(bias_constraint)
|
|
78
|
+
|
|
79
|
+
super().__init__(**kwargs)
|
|
80
|
+
|
|
81
|
+
def build(self, input_shapes):
|
|
82
|
+
if isinstance(input_shapes, (list, tuple)) and len(input_shapes) > 0 and isinstance(input_shapes[0], (list, tuple)):
|
|
83
|
+
feat_shape = input_shapes[0]
|
|
84
|
+
else:
|
|
85
|
+
feat_shape = input_shapes
|
|
86
|
+
input_dim = int(feat_shape[-1]) if feat_shape is not None and feat_shape[-1] is not None else 8
|
|
87
|
+
|
|
88
|
+
self.kernel = self.add_weight(
|
|
89
|
+
shape=(1, input_dim, self.units),
|
|
90
|
+
initializer=self.kernel_initializer,
|
|
91
|
+
name="kernel",
|
|
92
|
+
regularizer=self.kernel_regularizer,
|
|
93
|
+
constraint=self.kernel_constraint,
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
if self.use_bias:
|
|
97
|
+
self.bias = self.add_weight(
|
|
98
|
+
shape=(self.units,),
|
|
99
|
+
initializer=self.bias_initializer,
|
|
100
|
+
name="bias",
|
|
101
|
+
regularizer=self.bias_regularizer,
|
|
102
|
+
constraint=self.bias_constraint,
|
|
103
|
+
)
|
|
104
|
+
else:
|
|
105
|
+
self.bias = None
|
|
106
|
+
self.built = True
|
|
107
|
+
|
|
108
|
+
def call(self, inputs, A=None, **kwargs):
|
|
109
|
+
if A is not None:
|
|
110
|
+
features = inputs
|
|
111
|
+
elif isinstance(inputs, (list, tuple)):
|
|
112
|
+
features, A = inputs
|
|
113
|
+
else:
|
|
114
|
+
features, A = inputs, None
|
|
115
|
+
|
|
116
|
+
if A is not None and hasattr(A, "shape") and len(A.shape) == 2 and A.shape[0] == 2 and A.shape[1] != 2:
|
|
117
|
+
num_nodes = ops.shape(features)[-2]
|
|
118
|
+
a_dense = ops.zeros((num_nodes, num_nodes), dtype=features.dtype)
|
|
119
|
+
indices = ops.transpose(A, axes=[1, 0])
|
|
120
|
+
updates = ops.ones(shape=(ops.shape(A)[1],), dtype=features.dtype)
|
|
121
|
+
A = ops.scatter_update(a_dense, indices, updates)
|
|
122
|
+
|
|
123
|
+
was_2d = len(ops.shape(features)) == 2
|
|
124
|
+
if was_2d:
|
|
125
|
+
features = ops.expand_dims(features, 0)
|
|
126
|
+
if len(ops.shape(A)) == 2:
|
|
127
|
+
A = ops.expand_dims(A, 0)
|
|
128
|
+
|
|
129
|
+
# Calculate the layer operation of GCN
|
|
130
|
+
|
|
131
|
+
h_graph = dot((A, features), axes=1)
|
|
132
|
+
b = ops.shape(h_graph)[0]
|
|
133
|
+
kernel = repeat(self.kernel, b, axis=0)
|
|
134
|
+
output = dot((h_graph, kernel), axes=(-1, 1))
|
|
135
|
+
|
|
136
|
+
# Add optional bias & apply activation
|
|
137
|
+
if self.bias is not None:
|
|
138
|
+
output += self.bias
|
|
139
|
+
output = self.activation(output)
|
|
140
|
+
|
|
141
|
+
if was_2d:
|
|
142
|
+
output = ops.squeeze(output, 0)
|
|
143
|
+
|
|
144
|
+
return output
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
import math
|
|
2
|
+
from typing import Optional, Union, Tuple
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.utils import gcn_norm, is_tracing
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class GCN2Conv(MessagePassing):
|
|
10
|
+
r"""The graph convolutional operator from the `"Simple and Deep Graph
|
|
11
|
+
Convolutional Networks" <https://arxiv.org/abs/2007.02133>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
channels: Size of each input and output sample.
|
|
15
|
+
alpha: The strength of the initial residual connection :math:`\alpha`.
|
|
16
|
+
theta: The hyperparameter for the identity mapping :math:`\theta`.
|
|
17
|
+
(default: :obj:`None`)
|
|
18
|
+
layer: The layer index :math:`l`. (default: :obj:`None`)
|
|
19
|
+
shared_weights: If set to :obj:`True`, will use the same weights
|
|
20
|
+
for :math:`\mathbf{X}` and :math:`\mathbf{X}_0`. (default: :obj:`True`)
|
|
21
|
+
cached: If set to :obj:`True`, will cache the computation of normalization
|
|
22
|
+
coefficients. (default: ``False``)
|
|
23
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops.
|
|
24
|
+
(default: ``True``)
|
|
25
|
+
normalize: Whether to apply symmetric normalization. (default: ``True``)
|
|
26
|
+
|
|
27
|
+
Example:
|
|
28
|
+
```python
|
|
29
|
+
import numpy as np
|
|
30
|
+
from k3_node.layers import GCN2Conv
|
|
31
|
+
|
|
32
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
33
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
34
|
+
|
|
35
|
+
x_0 = x # initial node representations (from the first layer)
|
|
36
|
+
layer = GCN2Conv(channels=8, alpha=0.1, theta=0.5, layer=1)
|
|
37
|
+
out = layer(x, x_0, edge_index)
|
|
38
|
+
print(tuple(out.shape)) # (10, 8)
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
weighted_sum_message = True
|
|
43
|
+
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
channels: int,
|
|
47
|
+
alpha: float,
|
|
48
|
+
theta: Optional[float] = None,
|
|
49
|
+
layer: Optional[int] = None,
|
|
50
|
+
shared_weights: bool = True,
|
|
51
|
+
cached: bool = False,
|
|
52
|
+
add_self_loops: bool = True,
|
|
53
|
+
normalize: bool = True,
|
|
54
|
+
**kwargs,
|
|
55
|
+
):
|
|
56
|
+
super().__init__(aggr="add", **kwargs)
|
|
57
|
+
self.channels = channels
|
|
58
|
+
self.alpha = alpha
|
|
59
|
+
self.beta = 1.0
|
|
60
|
+
if theta is not None and layer is not None:
|
|
61
|
+
self.beta = math.log(theta / layer + 1.0)
|
|
62
|
+
self.cached = cached
|
|
63
|
+
self.normalize = normalize
|
|
64
|
+
self.add_self_loops = add_self_loops
|
|
65
|
+
self.shared_weights = shared_weights
|
|
66
|
+
|
|
67
|
+
self._cached_edge_index = None
|
|
68
|
+
self._cached_norm = None
|
|
69
|
+
|
|
70
|
+
def build(self, input_shape):
|
|
71
|
+
self.weight1 = self.add_weight(
|
|
72
|
+
shape=(self.channels, self.channels),
|
|
73
|
+
initializer="glorot_uniform",
|
|
74
|
+
name="weight1",
|
|
75
|
+
)
|
|
76
|
+
if not self.shared_weights:
|
|
77
|
+
self.weight2 = self.add_weight(
|
|
78
|
+
shape=(self.channels, self.channels),
|
|
79
|
+
initializer="glorot_uniform",
|
|
80
|
+
name="weight2",
|
|
81
|
+
)
|
|
82
|
+
else:
|
|
83
|
+
self.weight2 = None
|
|
84
|
+
self.built = True
|
|
85
|
+
|
|
86
|
+
def call(self, x, x_0, edge_index=None, edge_weight=None, **kwargs):
|
|
87
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
88
|
+
x, x_0, edge_index = x[0], x[1], x[2]
|
|
89
|
+
|
|
90
|
+
if self.normalize:
|
|
91
|
+
if self.cached and self._cached_edge_index is not None:
|
|
92
|
+
edge_index = self._cached_edge_index
|
|
93
|
+
edge_weight = self._cached_norm
|
|
94
|
+
else:
|
|
95
|
+
num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
|
|
96
|
+
edge_index, edge_weight = gcn_norm(
|
|
97
|
+
edge_index,
|
|
98
|
+
edge_weight,
|
|
99
|
+
num_nodes=num_nodes,
|
|
100
|
+
add_self_loops=self.add_self_loops,
|
|
101
|
+
flow=self.flow,
|
|
102
|
+
dtype=x.dtype,
|
|
103
|
+
)
|
|
104
|
+
if self.cached and not is_tracing(edge_index):
|
|
105
|
+
self._cached_edge_index = edge_index
|
|
106
|
+
self._cached_norm = edge_weight
|
|
107
|
+
|
|
108
|
+
h = self.propagate(edge_index, x=x, edge_weight=edge_weight)
|
|
109
|
+
h = (1.0 - self.alpha) * h
|
|
110
|
+
h_0 = self.alpha * x_0
|
|
111
|
+
|
|
112
|
+
if self.weight2 is None:
|
|
113
|
+
combined = h + h_0
|
|
114
|
+
out = (1.0 - self.beta) * combined + self.beta * ops.matmul(combined, self.weight1)
|
|
115
|
+
else:
|
|
116
|
+
term1 = (1.0 - self.beta) * h + self.beta * ops.matmul(h, self.weight1)
|
|
117
|
+
term2 = (1.0 - self.beta) * h_0 + self.beta * ops.matmul(h_0, self.weight2)
|
|
118
|
+
out = term1 + term2
|
|
119
|
+
|
|
120
|
+
return out
|
|
121
|
+
|
|
122
|
+
def message(self, x_j, edge_weight=None):
|
|
123
|
+
if edge_weight is None:
|
|
124
|
+
return x_j
|
|
125
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
126
|
+
|
|
@@ -0,0 +1,135 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
4
|
+
from k3_node.layers.conv.utils import gcn_norm, is_tracing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class GCNConv(MessagePassing):
|
|
8
|
+
r"""The graph convolutional operator from the `"Semi-supervised
|
|
9
|
+
Classification with Graph Convolutional Networks"
|
|
10
|
+
<https://arxiv.org/abs/1609.02907>`_ paper.
|
|
11
|
+
|
|
12
|
+
.. math::
|
|
13
|
+
\mathbf{X}^{\prime} = \mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}}
|
|
14
|
+
\mathbf{\hat{D}}^{-1/2} \mathbf{X} \mathbf{\Theta}
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
in_channels: Size of each input sample.
|
|
18
|
+
out_channels: Size of each output sample.
|
|
19
|
+
improved: If set to :obj:`True`, the layer computes
|
|
20
|
+
:math:`\mathbf{\hat{A}} = \mathbf{A} + 2 \mathbf{I}`.
|
|
21
|
+
(default: :obj:`False`)
|
|
22
|
+
cached: If set to :obj:`True`, the layer will cache the computation of
|
|
23
|
+
:math:`\mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}} \mathbf{\hat{D}}^{-1/2}`.
|
|
24
|
+
(default: :obj:`False`)
|
|
25
|
+
add_self_loops: If set to :obj:`False`, will not add
|
|
26
|
+
self-loops to the input graph. (default: :obj:`True`)
|
|
27
|
+
normalize: Whether to add self-loops and compute
|
|
28
|
+
symmetric normalization coefficients on the fly.
|
|
29
|
+
(default: :obj:`True`)
|
|
30
|
+
bias: If set to :obj:`False`, the layer will not learn
|
|
31
|
+
an additive bias. (default: :obj:`True`)
|
|
32
|
+
|
|
33
|
+
Example:
|
|
34
|
+
```python
|
|
35
|
+
import numpy as np
|
|
36
|
+
from k3_node.layers import GCNConv
|
|
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 = GCNConv(in_channels=8, out_channels=16)
|
|
42
|
+
out = layer(x, edge_index)
|
|
43
|
+
print(tuple(out.shape)) # (10, 16)
|
|
44
|
+
```
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
weighted_sum_message = True
|
|
48
|
+
|
|
49
|
+
def __init__(
|
|
50
|
+
self,
|
|
51
|
+
in_channels: int,
|
|
52
|
+
out_channels: int,
|
|
53
|
+
improved: bool = False,
|
|
54
|
+
cached: bool = False,
|
|
55
|
+
add_self_loops: bool = True,
|
|
56
|
+
normalize: bool = True,
|
|
57
|
+
bias: bool = True,
|
|
58
|
+
**kwargs,
|
|
59
|
+
):
|
|
60
|
+
super().__init__(aggr="add", **kwargs)
|
|
61
|
+
self.in_channels = in_channels
|
|
62
|
+
self.out_channels = out_channels
|
|
63
|
+
self.improved = improved
|
|
64
|
+
self.cached = cached
|
|
65
|
+
self.add_self_loops = add_self_loops
|
|
66
|
+
self.normalize = normalize
|
|
67
|
+
self.use_bias = bias
|
|
68
|
+
|
|
69
|
+
self.lin = layers.Dense(out_channels, use_bias=False)
|
|
70
|
+
self.bias = None
|
|
71
|
+
self._cached_edge_index = None
|
|
72
|
+
self._cached_norm = None
|
|
73
|
+
|
|
74
|
+
def build(self, input_shape):
|
|
75
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
76
|
+
feat_shape = input_shape[0]
|
|
77
|
+
else:
|
|
78
|
+
feat_shape = input_shape
|
|
79
|
+
self.lin.build(feat_shape)
|
|
80
|
+
if self.use_bias:
|
|
81
|
+
self.bias = self.add_weight(
|
|
82
|
+
shape=(self.out_channels,),
|
|
83
|
+
initializer="zeros",
|
|
84
|
+
name="bias",
|
|
85
|
+
)
|
|
86
|
+
else:
|
|
87
|
+
self.bias = None
|
|
88
|
+
self.built = True
|
|
89
|
+
|
|
90
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
91
|
+
# Handle legacy call conv((x, adj))
|
|
92
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
93
|
+
x, edge_index = x[0], x[1]
|
|
94
|
+
|
|
95
|
+
if not self.built:
|
|
96
|
+
feat_shape = x.shape if hasattr(x, "shape") and x.shape is not None else (None, self.in_channels)
|
|
97
|
+
self.build(feat_shape)
|
|
98
|
+
|
|
99
|
+
# Handle dense adjacency [N, N]
|
|
100
|
+
e_shape = getattr(edge_index, "shape", None)
|
|
101
|
+
if e_shape is not None and len(e_shape) == 2 and e_shape[0] is not None and e_shape[1] is not None and e_shape[0] > 2 and e_shape[0] == e_shape[1]:
|
|
102
|
+
where_adj = ops.where(edge_index != 0)
|
|
103
|
+
where_adj = where_adj if not isinstance(where_adj, list) else where_adj
|
|
104
|
+
edge_weight = ops.take(edge_index, where_adj[0] * e_shape[1] + where_adj[1]) if edge_weight is None else edge_weight
|
|
105
|
+
edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
|
|
106
|
+
|
|
107
|
+
if self.normalize:
|
|
108
|
+
if self.cached and self._cached_edge_index is not None:
|
|
109
|
+
edge_index = self._cached_edge_index
|
|
110
|
+
edge_weight = self._cached_norm
|
|
111
|
+
else:
|
|
112
|
+
num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
|
|
113
|
+
edge_index, edge_weight = gcn_norm(
|
|
114
|
+
edge_index,
|
|
115
|
+
edge_weight,
|
|
116
|
+
num_nodes=num_nodes,
|
|
117
|
+
improved=self.improved,
|
|
118
|
+
add_self_loops=self.add_self_loops,
|
|
119
|
+
flow=self.flow,
|
|
120
|
+
dtype=x.dtype,
|
|
121
|
+
)
|
|
122
|
+
if self.cached and not is_tracing(edge_index):
|
|
123
|
+
self._cached_edge_index = edge_index
|
|
124
|
+
self._cached_norm = edge_weight
|
|
125
|
+
|
|
126
|
+
x = self.lin(x)
|
|
127
|
+
out = self.propagate(edge_index, x=x, edge_weight=edge_weight)
|
|
128
|
+
if self.bias is not None:
|
|
129
|
+
out = out + self.bias
|
|
130
|
+
return out
|
|
131
|
+
|
|
132
|
+
def message(self, x_j, edge_weight=None):
|
|
133
|
+
if edge_weight is None:
|
|
134
|
+
return x_j
|
|
135
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
@@ -0,0 +1,163 @@
|
|
|
1
|
+
from typing import Optional, Union, Tuple
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
from keras.layers import Dense, BatchNormalization, LayerNormalization
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
from k3_node.layers.aggr import SoftmaxAggregation, PowerMeanAggregation
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class GENConv(MessagePassing):
|
|
11
|
+
r"""The generalized graph convolution operator from the `"DeeperGCN: All
|
|
12
|
+
You Need to Train Deeper GCNs" <https://arxiv.org/abs/2006.07739>`_ paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import GENConv
|
|
18
|
+
|
|
19
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
20
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
21
|
+
edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
|
|
22
|
+
|
|
23
|
+
layer = GENConv(in_channels=8, out_channels=16, edge_dim=3)
|
|
24
|
+
out = layer(x, edge_index, edge_attr)
|
|
25
|
+
print(tuple(out.shape)) # (10, 16)
|
|
26
|
+
```
|
|
27
|
+
"""
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
31
|
+
out_channels: int,
|
|
32
|
+
aggr: str = "softmax",
|
|
33
|
+
t: float = 1.0,
|
|
34
|
+
learn_t: bool = False,
|
|
35
|
+
p: float = 1.0,
|
|
36
|
+
learn_p: bool = False,
|
|
37
|
+
msg_norm: bool = False,
|
|
38
|
+
learn_msg_scale: bool = False,
|
|
39
|
+
norm: Optional[str] = "batch",
|
|
40
|
+
num_layers: int = 2,
|
|
41
|
+
expansion: int = 2,
|
|
42
|
+
eps: float = 1e-7,
|
|
43
|
+
bias: bool = False,
|
|
44
|
+
edge_dim: Optional[int] = None,
|
|
45
|
+
**kwargs,
|
|
46
|
+
):
|
|
47
|
+
if aggr in ("softmax", "softmax_sg"):
|
|
48
|
+
aggr_module = SoftmaxAggregation(t=t, learn=learn_t)
|
|
49
|
+
elif aggr in ("power", "powermean"):
|
|
50
|
+
aggr_module = PowerMeanAggregation(p=p, learn=learn_p)
|
|
51
|
+
else:
|
|
52
|
+
aggr_module = aggr
|
|
53
|
+
|
|
54
|
+
super().__init__(aggr=aggr_module, **kwargs)
|
|
55
|
+
|
|
56
|
+
self.in_channels = in_channels
|
|
57
|
+
self.out_channels = out_channels
|
|
58
|
+
self.eps = eps
|
|
59
|
+
self.edge_dim = edge_dim
|
|
60
|
+
self.use_bias = bias
|
|
61
|
+
|
|
62
|
+
if isinstance(in_channels, int):
|
|
63
|
+
self.in_channels_l = in_channels
|
|
64
|
+
self.in_channels_r = in_channels
|
|
65
|
+
else:
|
|
66
|
+
self.in_channels_l, self.in_channels_r = in_channels
|
|
67
|
+
|
|
68
|
+
if self.in_channels_l != out_channels:
|
|
69
|
+
self.lin_src = Dense(out_channels, use_bias=bias)
|
|
70
|
+
else:
|
|
71
|
+
self.lin_src = None
|
|
72
|
+
|
|
73
|
+
if edge_dim is not None and edge_dim != out_channels:
|
|
74
|
+
self.lin_edge = Dense(out_channels, use_bias=bias)
|
|
75
|
+
else:
|
|
76
|
+
self.lin_edge = None
|
|
77
|
+
|
|
78
|
+
if self.in_channels_r != out_channels:
|
|
79
|
+
self.lin_dst = Dense(out_channels, use_bias=bias)
|
|
80
|
+
else:
|
|
81
|
+
self.lin_dst = None
|
|
82
|
+
|
|
83
|
+
# MLP
|
|
84
|
+
self.mlp_layers = []
|
|
85
|
+
channels = [out_channels]
|
|
86
|
+
for _ in range(num_layers - 1):
|
|
87
|
+
channels.append(out_channels * expansion)
|
|
88
|
+
channels.append(out_channels)
|
|
89
|
+
|
|
90
|
+
for i in range(len(channels) - 1):
|
|
91
|
+
self.mlp_layers.append(Dense(channels[i + 1], use_bias=bias))
|
|
92
|
+
if i < len(channels) - 2:
|
|
93
|
+
if norm == "batch":
|
|
94
|
+
self.mlp_layers.append(BatchNormalization(momentum=0.9, epsilon=1e-5))
|
|
95
|
+
elif norm == "layer":
|
|
96
|
+
self.mlp_layers.append(LayerNormalization())
|
|
97
|
+
self.mlp_layers.append(keras.layers.ReLU())
|
|
98
|
+
|
|
99
|
+
def build(self, input_shape=None):
|
|
100
|
+
if self.lin_src is not None:
|
|
101
|
+
self.lin_src.build((None, self.in_channels_l))
|
|
102
|
+
if self.lin_dst is not None:
|
|
103
|
+
self.lin_dst.build((None, self.in_channels_r))
|
|
104
|
+
if self.lin_edge is not None:
|
|
105
|
+
self.lin_edge.build((None, self.edge_dim))
|
|
106
|
+
curr_dim = self.out_channels
|
|
107
|
+
for layer in self.mlp_layers:
|
|
108
|
+
if hasattr(layer, "build"):
|
|
109
|
+
layer.build((None, curr_dim))
|
|
110
|
+
if hasattr(layer, "units"):
|
|
111
|
+
curr_dim = layer.units
|
|
112
|
+
self.built = True
|
|
113
|
+
|
|
114
|
+
def call(self, inputs, edge_index=None, edge_attr=None, training=None, **kwargs):
|
|
115
|
+
if edge_index is None:
|
|
116
|
+
if isinstance(inputs, (list, tuple)):
|
|
117
|
+
if len(inputs) == 3:
|
|
118
|
+
x, edge_index, edge_attr = inputs
|
|
119
|
+
elif len(inputs) == 2:
|
|
120
|
+
x, edge_index = inputs
|
|
121
|
+
else:
|
|
122
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
123
|
+
else:
|
|
124
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
125
|
+
else:
|
|
126
|
+
x = inputs
|
|
127
|
+
|
|
128
|
+
if not self.built:
|
|
129
|
+
self.build()
|
|
130
|
+
|
|
131
|
+
if isinstance(x, (list, tuple)):
|
|
132
|
+
x_l, x_r = x
|
|
133
|
+
else:
|
|
134
|
+
x_l = x_r = x
|
|
135
|
+
|
|
136
|
+
if self.lin_src is not None:
|
|
137
|
+
x_l = self.lin_src(x_l)
|
|
138
|
+
|
|
139
|
+
num_nodes = ops.shape(x_r)[0]
|
|
140
|
+
out = self.propagate(
|
|
141
|
+
edge_index,
|
|
142
|
+
x=(x_l, x_r),
|
|
143
|
+
edge_attr=edge_attr,
|
|
144
|
+
size=(ops.shape(x_l)[0], num_nodes),
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
x_dst = x_r
|
|
148
|
+
if self.lin_dst is not None:
|
|
149
|
+
x_dst = self.lin_dst(x_dst)
|
|
150
|
+
out = out + x_dst
|
|
151
|
+
|
|
152
|
+
for layer in self.mlp_layers:
|
|
153
|
+
# Batch norm needs `training` explicitly (Keras does not propagate it on JAX).
|
|
154
|
+
out = layer(out, training=training) if isinstance(layer, BatchNormalization) else layer(out)
|
|
155
|
+
|
|
156
|
+
return out
|
|
157
|
+
|
|
158
|
+
def message(self, x_j, edge_attr=None):
|
|
159
|
+
if edge_attr is not None and self.lin_edge is not None:
|
|
160
|
+
edge_attr = self.lin_edge(edge_attr)
|
|
161
|
+
msg = x_j if edge_attr is None else x_j + edge_attr
|
|
162
|
+
return ops.relu(msg) + self.eps
|
|
163
|
+
|