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,451 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
import inspect
|
|
3
|
+
from typing import Any, List, Optional, Tuple, Union
|
|
4
|
+
from keras import layers, ops
|
|
5
|
+
|
|
6
|
+
from k3_node.utils import (
|
|
7
|
+
is_layer_kwarg,
|
|
8
|
+
deserialize_kwarg,
|
|
9
|
+
serialize_kwarg,
|
|
10
|
+
is_keras_kwarg,
|
|
11
|
+
deserialize_scatter,
|
|
12
|
+
serialize_scatter,
|
|
13
|
+
)
|
|
14
|
+
from k3_node.ops import get_source_target
|
|
15
|
+
from k3_node.layers.aggr.resolver import aggregation_resolver
|
|
16
|
+
from k3_node.layers.aggr.base import Aggregation
|
|
17
|
+
from k3_node.layers.aggr.basic import SumAggregation
|
|
18
|
+
from k3_node.ops.segment import segment_sum
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class MessagePassing(layers.Layer):
|
|
22
|
+
r"""Base class for creating Message Passing Neural Networks (MPNNs).
|
|
23
|
+
|
|
24
|
+
Args:
|
|
25
|
+
aggr: The aggregation scheme to use, such as ``"add"``, ``"sum"``,
|
|
26
|
+
``"mean"``, ``"min"``, ``"max"``, ``"mul"``, or an instance of
|
|
27
|
+
:class:`~k3_node.layers.aggr.Aggregation`. (default: ``"add"``)
|
|
28
|
+
flow: The direction of message passing (``"source_to_target"`` or
|
|
29
|
+
``"target_to_source"``). (default: ``"source_to_target"``)
|
|
30
|
+
node_dim: The axis along which to index node features. (default: ``-2``)
|
|
31
|
+
decomposed_layers: Number of decomposed layers for memory-efficient
|
|
32
|
+
aggregation. (default: ``1``)
|
|
33
|
+
|
|
34
|
+
Example:
|
|
35
|
+
```python
|
|
36
|
+
import numpy as np
|
|
37
|
+
import keras
|
|
38
|
+
from k3_node.layers import MessagePassing
|
|
39
|
+
|
|
40
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
41
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
42
|
+
|
|
43
|
+
# A minimal custom layer: average the neighbors' features, then transform them.
|
|
44
|
+
class MeanConv(MessagePassing):
|
|
45
|
+
def __init__(self, units, **kwargs):
|
|
46
|
+
super().__init__(aggr="mean", **kwargs)
|
|
47
|
+
self.dense = keras.layers.Dense(units)
|
|
48
|
+
|
|
49
|
+
def call(self, x, edge_index):
|
|
50
|
+
return self.dense(self.propagate(edge_index, x=x))
|
|
51
|
+
|
|
52
|
+
def message(self, x_j):
|
|
53
|
+
return x_j # features of the source node of every edge
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
layer = MeanConv(16)
|
|
57
|
+
out = layer(x, edge_index)
|
|
58
|
+
print(tuple(out.shape)) # (10, 16)
|
|
59
|
+
```
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
def __init__(
|
|
63
|
+
self,
|
|
64
|
+
aggr: Union[str, List[str], Aggregation, None] = "add",
|
|
65
|
+
flow: str = "source_to_target",
|
|
66
|
+
node_dim: int = -2,
|
|
67
|
+
decomposed_layers: int = 1,
|
|
68
|
+
**kwargs,
|
|
69
|
+
):
|
|
70
|
+
# Support legacy aggregate arg
|
|
71
|
+
if "aggregate" in kwargs:
|
|
72
|
+
aggr = kwargs.pop("aggregate")
|
|
73
|
+
|
|
74
|
+
# Extract and set layer kwargs for backwards compatibility
|
|
75
|
+
self.kwargs_keys = []
|
|
76
|
+
for key in list(kwargs.keys()):
|
|
77
|
+
if is_layer_kwarg(key):
|
|
78
|
+
attr = kwargs.pop(key)
|
|
79
|
+
attr = deserialize_kwarg(key, attr)
|
|
80
|
+
self.kwargs_keys.append(key)
|
|
81
|
+
setattr(self, key, attr)
|
|
82
|
+
|
|
83
|
+
unknown = sorted(k for k in kwargs if not is_keras_kwarg(k))
|
|
84
|
+
if unknown:
|
|
85
|
+
# Silently dropping these hides typos such as `num_heads=` for `heads=`.
|
|
86
|
+
raise TypeError(f"{type(self).__name__}() got unexpected keyword argument(s): {', '.join(unknown)}")
|
|
87
|
+
super().__init__(**kwargs)
|
|
88
|
+
self.aggr = aggr
|
|
89
|
+
self.flow = flow
|
|
90
|
+
self.node_dim = node_dim
|
|
91
|
+
self.decomposed_layers = decomposed_layers
|
|
92
|
+
|
|
93
|
+
if flow not in ["source_to_target", "target_to_source"]:
|
|
94
|
+
raise ValueError(f"Flow {flow} must be 'source_to_target' or 'target_to_source'")
|
|
95
|
+
|
|
96
|
+
# Legacy scatter operator support
|
|
97
|
+
try:
|
|
98
|
+
scatter_name = "sum" if aggr in ("add", "sum") else aggr
|
|
99
|
+
self.agg = deserialize_scatter(scatter_name)
|
|
100
|
+
except Exception:
|
|
101
|
+
self.agg = ops.segment_sum
|
|
102
|
+
|
|
103
|
+
if aggr is None:
|
|
104
|
+
self.aggr_module = None
|
|
105
|
+
elif isinstance(aggr, Aggregation):
|
|
106
|
+
self.aggr_module = aggr
|
|
107
|
+
elif isinstance(aggr, str):
|
|
108
|
+
aggr_name = "sum" if aggr == "add" else aggr
|
|
109
|
+
try:
|
|
110
|
+
self.aggr_module = aggregation_resolver(aggr_name)
|
|
111
|
+
except Exception:
|
|
112
|
+
self.aggr_module = SumAggregation()
|
|
113
|
+
elif isinstance(aggr, (list, tuple)):
|
|
114
|
+
from k3_node.layers.aggr.multi import MultiAggregation
|
|
115
|
+
self.aggr_module = MultiAggregation(aggr)
|
|
116
|
+
else:
|
|
117
|
+
self.aggr_module = SumAggregation()
|
|
118
|
+
|
|
119
|
+
self.msg_signature = inspect.signature(self.message).parameters
|
|
120
|
+
self.agg_signature = inspect.signature(self.aggregate).parameters
|
|
121
|
+
self.upd_signature = inspect.signature(self.update).parameters
|
|
122
|
+
|
|
123
|
+
def build(self, input_shape=None):
|
|
124
|
+
self.built = True
|
|
125
|
+
|
|
126
|
+
@staticmethod
|
|
127
|
+
def get_inputs(inputs):
|
|
128
|
+
if len(inputs) == 3:
|
|
129
|
+
x, a, e = inputs
|
|
130
|
+
if hasattr(e, "shape") and e.shape is not None:
|
|
131
|
+
assert len(e.shape) in (2, 3), "E must have rank 2 or 3"
|
|
132
|
+
elif len(inputs) == 2:
|
|
133
|
+
x, a = inputs
|
|
134
|
+
e = None
|
|
135
|
+
else:
|
|
136
|
+
raise ValueError(
|
|
137
|
+
"Expected 2 or 3 inputs tensors (X, A, E), got {}.".format(len(inputs))
|
|
138
|
+
)
|
|
139
|
+
if hasattr(a, "shape") and a.shape is not None:
|
|
140
|
+
assert len(a.shape) == 2, "A must have rank 2"
|
|
141
|
+
return x, a, e
|
|
142
|
+
|
|
143
|
+
def get_targets(self, x):
|
|
144
|
+
return ops.take(x, self.index_targets, axis=self.node_dim)
|
|
145
|
+
|
|
146
|
+
def get_sources(self, x):
|
|
147
|
+
return ops.take(x, self.index_sources, axis=self.node_dim)
|
|
148
|
+
|
|
149
|
+
def get_kwargs(self, x, a, e, signature, kwargs):
|
|
150
|
+
output = {}
|
|
151
|
+
for k in signature.keys():
|
|
152
|
+
if k == "kwargs":
|
|
153
|
+
pass
|
|
154
|
+
elif k == "x":
|
|
155
|
+
output[k] = x
|
|
156
|
+
elif k == "a":
|
|
157
|
+
output[k] = a
|
|
158
|
+
elif k == "e":
|
|
159
|
+
output[k] = e
|
|
160
|
+
elif k in kwargs:
|
|
161
|
+
output[k] = kwargs[k]
|
|
162
|
+
elif signature[k].default is inspect.Parameter.empty:
|
|
163
|
+
pass
|
|
164
|
+
else:
|
|
165
|
+
pass
|
|
166
|
+
return output
|
|
167
|
+
|
|
168
|
+
def _get_dim_size(self, kwargs, i, size=None):
|
|
169
|
+
if size is not None and size[1] is not None:
|
|
170
|
+
return size[1]
|
|
171
|
+
x = kwargs.get("x", None)
|
|
172
|
+
if x is not None:
|
|
173
|
+
if isinstance(x, (tuple, list)):
|
|
174
|
+
target_x = x[1] if x[1] is not None else x[0]
|
|
175
|
+
if target_x is not None:
|
|
176
|
+
if hasattr(target_x, "shape") and target_x.shape[self.node_dim] is not None:
|
|
177
|
+
return int(target_x.shape[self.node_dim])
|
|
178
|
+
return ops.shape(target_x)[self.node_dim]
|
|
179
|
+
else:
|
|
180
|
+
if hasattr(x, "shape") and x.shape[self.node_dim] is not None:
|
|
181
|
+
return int(x.shape[self.node_dim])
|
|
182
|
+
return ops.shape(x)[self.node_dim]
|
|
183
|
+
for k, val in kwargs.items():
|
|
184
|
+
if k in ("edge_index", "edge_attr", "edge_weight", "ptr"):
|
|
185
|
+
continue
|
|
186
|
+
if hasattr(val, "shape") and len(val.shape) >= 2:
|
|
187
|
+
if val.shape[self.node_dim] is not None:
|
|
188
|
+
return int(val.shape[self.node_dim])
|
|
189
|
+
return ops.shape(val)[self.node_dim]
|
|
190
|
+
if i is not None:
|
|
191
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
192
|
+
if is_tracing(i):
|
|
193
|
+
return 0
|
|
194
|
+
if hasattr(i, "shape") and len(i.shape) > 0 and i.shape[0] == 0:
|
|
195
|
+
return 0
|
|
196
|
+
try:
|
|
197
|
+
return int(ops.max(i)) + 1
|
|
198
|
+
except Exception:
|
|
199
|
+
return 0
|
|
200
|
+
return 0
|
|
201
|
+
|
|
202
|
+
def propagate(self, *args, **kwargs: Any):
|
|
203
|
+
r"""The initial call to start propagating messages."""
|
|
204
|
+
# Detect legacy Spektral-style call: propagate(x, a, e=None, **kwargs)
|
|
205
|
+
is_spektral = False
|
|
206
|
+
if len(args) >= 2:
|
|
207
|
+
arg0, arg1 = args[0], args[1]
|
|
208
|
+
shape0 = arg0.shape if hasattr(arg0, "shape") and arg0.shape is not None else ()
|
|
209
|
+
shape1 = arg1.shape if hasattr(arg1, "shape") and arg1.shape is not None else ()
|
|
210
|
+
if len(shape1) == 2 and shape1[0] is not None and shape1[0] == shape1[1]:
|
|
211
|
+
is_spektral = True
|
|
212
|
+
elif len(shape0) >= 2 and shape0[0] is not None and shape0[0] != 2 and not isinstance(arg1, (tuple, list, type(None))):
|
|
213
|
+
if len(shape1) == 2 and shape1[0] is not None and shape1[0] != 2:
|
|
214
|
+
is_spektral = True
|
|
215
|
+
|
|
216
|
+
if is_spektral:
|
|
217
|
+
x, a = args[0], args[1]
|
|
218
|
+
e = args[2] if len(args) >= 3 else kwargs.get("e", None)
|
|
219
|
+
self.n_nodes = x.shape[-2] if hasattr(x, "shape") and x.shape[-2] is not None else ops.shape(x)[-2]
|
|
220
|
+
self.index_sources, self.index_targets = get_source_target(a)
|
|
221
|
+
|
|
222
|
+
# Call legacy message
|
|
223
|
+
msg_kwargs = self.get_kwargs(x, a, e, self.msg_signature, kwargs)
|
|
224
|
+
messages = self.message(**msg_kwargs)
|
|
225
|
+
|
|
226
|
+
# Call legacy aggregate
|
|
227
|
+
agg_kwargs = self.get_kwargs(x, a, e, self.agg_signature, kwargs)
|
|
228
|
+
embeddings = self.aggregate(messages, **agg_kwargs)
|
|
229
|
+
|
|
230
|
+
# Call legacy update
|
|
231
|
+
upd_kwargs = self.get_kwargs(x, a, e, self.upd_signature, kwargs)
|
|
232
|
+
output = self.update(embeddings, **upd_kwargs)
|
|
233
|
+
return output
|
|
234
|
+
|
|
235
|
+
# Standard PyG MessagePassing propagate(edge_index, size=None, **kwargs)
|
|
236
|
+
if len(args) >= 1:
|
|
237
|
+
edge_index = args[0]
|
|
238
|
+
size = args[1] if len(args) >= 2 else kwargs.pop("size", None)
|
|
239
|
+
else:
|
|
240
|
+
edge_index = kwargs.pop("edge_index")
|
|
241
|
+
size = kwargs.pop("size", None)
|
|
242
|
+
|
|
243
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
244
|
+
|
|
245
|
+
# Handle dense adjacency [N, N] passed as edge_index
|
|
246
|
+
e_shape = getattr(edge_index, "shape", None)
|
|
247
|
+
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]:
|
|
248
|
+
where_adj = ops.where(edge_index != 0)
|
|
249
|
+
where_adj = where_adj if not isinstance(where_adj, list) else where_adj
|
|
250
|
+
edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
|
|
251
|
+
|
|
252
|
+
if self.flow == "source_to_target":
|
|
253
|
+
i = edge_index[1] # Target
|
|
254
|
+
j = edge_index[0] # Source
|
|
255
|
+
else:
|
|
256
|
+
i = edge_index[0]
|
|
257
|
+
j = edge_index[1]
|
|
258
|
+
|
|
259
|
+
i = ops.cast(i, "int32")
|
|
260
|
+
j = ops.cast(j, "int32")
|
|
261
|
+
self.index_targets = i
|
|
262
|
+
self.index_sources = j
|
|
263
|
+
|
|
264
|
+
dim_size = size[1] if size is not None and size[1] is not None else self._get_dim_size(kwargs, i, size)
|
|
265
|
+
self.n_nodes = dim_size
|
|
266
|
+
|
|
267
|
+
# Layers whose message is `edge_weight * x_j`, summed, can skip per-edge messages
|
|
268
|
+
fused = self._fused_weighted_sum(i, j, dim_size, size, kwargs)
|
|
269
|
+
if fused is not None:
|
|
270
|
+
return self.update(fused, **self._update_kwargs(kwargs))
|
|
271
|
+
|
|
272
|
+
# Construct message arguments
|
|
273
|
+
msg_kwargs = {}
|
|
274
|
+
for param_name in self.msg_signature.keys():
|
|
275
|
+
if param_name in kwargs:
|
|
276
|
+
msg_kwargs[param_name] = kwargs[param_name]
|
|
277
|
+
elif param_name in ("dim_size", "size_i"):
|
|
278
|
+
msg_kwargs[param_name] = dim_size
|
|
279
|
+
elif param_name == "size_j":
|
|
280
|
+
msg_kwargs["size_j"] = size[0] if size is not None and size[0] is not None else dim_size
|
|
281
|
+
elif param_name == "index":
|
|
282
|
+
msg_kwargs["index"] = i
|
|
283
|
+
elif param_name == "edge_index":
|
|
284
|
+
msg_kwargs["edge_index"] = edge_index
|
|
285
|
+
elif param_name.endswith("_i"):
|
|
286
|
+
root = param_name[:-2]
|
|
287
|
+
if root in kwargs:
|
|
288
|
+
val = kwargs[root]
|
|
289
|
+
val = val[1] if isinstance(val, (tuple, list)) else val
|
|
290
|
+
msg_kwargs[param_name] = ops.take(val, i, axis=self.node_dim) if val is not None else None
|
|
291
|
+
elif param_name.endswith("_j"):
|
|
292
|
+
root = param_name[:-2]
|
|
293
|
+
if root in kwargs:
|
|
294
|
+
val = kwargs[root]
|
|
295
|
+
val = val[0] if isinstance(val, (tuple, list)) else val
|
|
296
|
+
msg_kwargs[param_name] = ops.take(val, j, axis=self.node_dim) if val is not None else None
|
|
297
|
+
elif param_name == "index":
|
|
298
|
+
msg_kwargs["index"] = i
|
|
299
|
+
elif param_name in ("dim_size", "size_i"):
|
|
300
|
+
msg_kwargs[param_name] = dim_size
|
|
301
|
+
elif param_name == "size_j":
|
|
302
|
+
msg_kwargs["size_j"] = size[0] if size is not None and size[0] is not None else dim_size
|
|
303
|
+
elif param_name == "edge_index":
|
|
304
|
+
msg_kwargs["edge_index"] = edge_index
|
|
305
|
+
|
|
306
|
+
out = self.message(**msg_kwargs)
|
|
307
|
+
|
|
308
|
+
# Aggregate
|
|
309
|
+
agg_kwargs = {}
|
|
310
|
+
for param_name in self.agg_signature.keys():
|
|
311
|
+
if param_name in kwargs and param_name not in ["inputs", "index", "ptr", "dim_size"]:
|
|
312
|
+
agg_kwargs[param_name] = kwargs[param_name]
|
|
313
|
+
|
|
314
|
+
out = self.aggregate(
|
|
315
|
+
out,
|
|
316
|
+
index=i,
|
|
317
|
+
ptr=kwargs.get("ptr", None),
|
|
318
|
+
dim_size=dim_size,
|
|
319
|
+
**agg_kwargs,
|
|
320
|
+
)
|
|
321
|
+
|
|
322
|
+
return self.update(out, **self._update_kwargs(kwargs))
|
|
323
|
+
|
|
324
|
+
def _update_kwargs(self, kwargs):
|
|
325
|
+
upd_kwargs = {}
|
|
326
|
+
for param_name in self.upd_signature.keys():
|
|
327
|
+
if param_name in kwargs and param_name not in ["inputs", "aggr_out", "embeddings"]:
|
|
328
|
+
upd_kwargs[param_name] = kwargs[param_name]
|
|
329
|
+
elif param_name.endswith("_i"):
|
|
330
|
+
root = param_name[:-2]
|
|
331
|
+
if root in kwargs:
|
|
332
|
+
val = kwargs[root]
|
|
333
|
+
val = val[1] if isinstance(val, (tuple, list)) else val
|
|
334
|
+
upd_kwargs[param_name] = val
|
|
335
|
+
return upd_kwargs
|
|
336
|
+
|
|
337
|
+
#: Set to ``True`` in layers whose ``message`` is ``edge_weight * x_j`` (or ``x_j``) and that
|
|
338
|
+
#: sum the messages; ``propagate`` then uses a sparse matrix product (see ``k3_node.ops.spmm``).
|
|
339
|
+
weighted_sum_message = False
|
|
340
|
+
|
|
341
|
+
def _fused_weighted_sum(self, i, j, dim_size, size, kwargs):
|
|
342
|
+
if not self.weighted_sum_message or keras.config.backend() != "torch":
|
|
343
|
+
return None
|
|
344
|
+
if not (self.aggr in ("add", "sum") and type(self.aggr_module).__name__ == "SumAggregation"):
|
|
345
|
+
return None
|
|
346
|
+
x = kwargs.get("x")
|
|
347
|
+
if x is None or isinstance(x, (tuple, list)) or len(getattr(x, "shape", ())) != 2 or self.node_dim not in (0, -2):
|
|
348
|
+
return None
|
|
349
|
+
if set(kwargs) - {"x", "edge_weight"}:
|
|
350
|
+
return None
|
|
351
|
+
from k3_node.ops.sparse import spmm
|
|
352
|
+
|
|
353
|
+
return spmm(i, j, kwargs.get("edge_weight"), x, dim_size)
|
|
354
|
+
|
|
355
|
+
def message(self, x=None, x_j=None, **kwargs):
|
|
356
|
+
r"""Constructs messages from node :math:`j` to node :math:`i`."""
|
|
357
|
+
if x_j is not None:
|
|
358
|
+
return x_j
|
|
359
|
+
if x is not None and hasattr(self, "index_sources"):
|
|
360
|
+
return self.get_sources(x)
|
|
361
|
+
return x_j
|
|
362
|
+
|
|
363
|
+
def aggregate(
|
|
364
|
+
self,
|
|
365
|
+
inputs=None,
|
|
366
|
+
index=None,
|
|
367
|
+
ptr: Optional[Any] = None,
|
|
368
|
+
dim_size: Optional[int] = None,
|
|
369
|
+
**kwargs,
|
|
370
|
+
):
|
|
371
|
+
r"""Aggregates messages from neighbors as given by :obj:`index`."""
|
|
372
|
+
# Check if called as Spektral: aggregate(messages, **kwargs)
|
|
373
|
+
if inputs is not None and index is None and hasattr(self, "index_targets"):
|
|
374
|
+
index = self.index_targets
|
|
375
|
+
dim_size = self.n_nodes
|
|
376
|
+
if hasattr(self, "agg") and callable(self.agg):
|
|
377
|
+
return self.agg(inputs, index, dim_size)
|
|
378
|
+
|
|
379
|
+
if self.aggr_module is not None:
|
|
380
|
+
return self.aggr_module(
|
|
381
|
+
inputs, index=index, ptr=ptr, dim_size=dim_size, dim=self.node_dim
|
|
382
|
+
)
|
|
383
|
+
return segment_sum(inputs, index, num_segments=dim_size)
|
|
384
|
+
|
|
385
|
+
def update(self, embeddings=None, **kwargs):
|
|
386
|
+
r"""Updates node embeddings."""
|
|
387
|
+
return embeddings
|
|
388
|
+
|
|
389
|
+
def edge_updater(
|
|
390
|
+
self,
|
|
391
|
+
edge_index,
|
|
392
|
+
size: Optional[Tuple[Optional[int], Optional[int]]] = None,
|
|
393
|
+
**kwargs: Any,
|
|
394
|
+
):
|
|
395
|
+
r"""Computes or updates edge-level representations."""
|
|
396
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
397
|
+
if self.flow == "source_to_target":
|
|
398
|
+
i = edge_index[1]
|
|
399
|
+
j = edge_index[0]
|
|
400
|
+
else:
|
|
401
|
+
i = edge_index[0]
|
|
402
|
+
j = edge_index[1]
|
|
403
|
+
|
|
404
|
+
i = ops.cast(i, "int32")
|
|
405
|
+
j = ops.cast(j, "int32")
|
|
406
|
+
|
|
407
|
+
dim_size = self._get_dim_size(kwargs, i, size)
|
|
408
|
+
edge_update_params = inspect.signature(self.edge_update).parameters
|
|
409
|
+
|
|
410
|
+
upd_kwargs = {}
|
|
411
|
+
for param_name in edge_update_params.keys():
|
|
412
|
+
if param_name in kwargs:
|
|
413
|
+
upd_kwargs[param_name] = kwargs[param_name]
|
|
414
|
+
elif param_name.endswith("_i"):
|
|
415
|
+
root = param_name[:-2]
|
|
416
|
+
if root in kwargs:
|
|
417
|
+
val = kwargs[root]
|
|
418
|
+
val = val[1] if isinstance(val, (tuple, list)) else val
|
|
419
|
+
upd_kwargs[param_name] = ops.take(val, i, axis=self.node_dim) if val is not None else None
|
|
420
|
+
elif param_name.endswith("_j"):
|
|
421
|
+
root = param_name[:-2]
|
|
422
|
+
if root in kwargs:
|
|
423
|
+
val = kwargs[root]
|
|
424
|
+
val = val[0] if isinstance(val, (tuple, list)) else val
|
|
425
|
+
upd_kwargs[param_name] = ops.take(val, j, axis=self.node_dim) if val is not None else None
|
|
426
|
+
elif param_name == "index":
|
|
427
|
+
upd_kwargs["index"] = i
|
|
428
|
+
elif param_name == "dim_size":
|
|
429
|
+
upd_kwargs["dim_size"] = dim_size
|
|
430
|
+
elif param_name == "edge_index":
|
|
431
|
+
upd_kwargs["edge_index"] = edge_index
|
|
432
|
+
|
|
433
|
+
return self.edge_update(**upd_kwargs)
|
|
434
|
+
|
|
435
|
+
def edge_update(self, **kwargs):
|
|
436
|
+
r"""Computes or updates edge attributes."""
|
|
437
|
+
raise NotImplementedError
|
|
438
|
+
|
|
439
|
+
def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
|
|
440
|
+
r"""Default call handler supporting both (x, edge_index) and legacy (inputs,) tuples."""
|
|
441
|
+
if edge_index is None and isinstance(inputs, (tuple, list)):
|
|
442
|
+
x, a, e = self.get_inputs(inputs)
|
|
443
|
+
return self.propagate(x, a, e, **kwargs)
|
|
444
|
+
if edge_index is not None:
|
|
445
|
+
if edge_attr is not None:
|
|
446
|
+
kwargs["edge_attr"] = edge_attr
|
|
447
|
+
kwargs["e"] = edge_attr
|
|
448
|
+
return self.propagate(edge_index, x=inputs, **kwargs)
|
|
449
|
+
raise NotImplementedError(
|
|
450
|
+
f"Layer {self.__class__.__name__} does not implement call() with inputs={inputs}, edge_index={edge_index}"
|
|
451
|
+
)
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
from typing import Union, Tuple
|
|
2
|
+
import keras
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.utils import degree
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class MFConv(MessagePassing):
|
|
10
|
+
r"""The molecular fingerprint graph convolutional operator from the
|
|
11
|
+
`"Convolutional Networks on Graphs for Learning Molecular Fingerprints"
|
|
12
|
+
<https://arxiv.org/abs/1509.09292>`_ paper.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
16
|
+
out_channels: Size of each output sample.
|
|
17
|
+
max_degree: The maximum degree of any node. (default: ``10``)
|
|
18
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
19
|
+
(default: ``True``)
|
|
20
|
+
|
|
21
|
+
Example:
|
|
22
|
+
```python
|
|
23
|
+
import numpy as np
|
|
24
|
+
from k3_node.layers import MFConv
|
|
25
|
+
|
|
26
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
27
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
28
|
+
|
|
29
|
+
layer = MFConv(in_channels=8, out_channels=16)
|
|
30
|
+
out = layer(x, edge_index)
|
|
31
|
+
print(tuple(out.shape)) # (10, 16)
|
|
32
|
+
```
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
def __init__(
|
|
36
|
+
self,
|
|
37
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
38
|
+
out_channels: int,
|
|
39
|
+
max_degree: int = 10,
|
|
40
|
+
bias: bool = True,
|
|
41
|
+
**kwargs,
|
|
42
|
+
):
|
|
43
|
+
super().__init__(aggr="add", **kwargs)
|
|
44
|
+
self.in_channels = in_channels
|
|
45
|
+
self.out_channels = out_channels
|
|
46
|
+
self.max_degree = max_degree
|
|
47
|
+
self.use_bias = bias
|
|
48
|
+
|
|
49
|
+
self.lins_l = [layers.Dense(out_channels, use_bias=bias) for _ in range(max_degree + 1)]
|
|
50
|
+
self.lins_r = [layers.Dense(out_channels, use_bias=False) for _ in range(max_degree + 1)]
|
|
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_src = input_shape[0][-1]
|
|
55
|
+
in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
|
|
56
|
+
else:
|
|
57
|
+
in_channels_src = input_shape[-1]
|
|
58
|
+
in_channels_dst = input_shape[-1]
|
|
59
|
+
|
|
60
|
+
for lin_l in self.lins_l:
|
|
61
|
+
lin_l.build((None, in_channels_src))
|
|
62
|
+
for lin_r in self.lins_r:
|
|
63
|
+
lin_r.build((None, in_channels_dst))
|
|
64
|
+
self.built = True
|
|
65
|
+
|
|
66
|
+
def call(self, x, edge_index=None, size=None, **kwargs):
|
|
67
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
68
|
+
x, edge_index = x[0], x[1]
|
|
69
|
+
|
|
70
|
+
if not isinstance(x, (tuple, list)):
|
|
71
|
+
x_src, x_dst = x, x
|
|
72
|
+
else:
|
|
73
|
+
x_src, x_dst = x[0], x[1]
|
|
74
|
+
|
|
75
|
+
target_idx = edge_index[1] if self.flow == "source_to_target" else edge_index[0]
|
|
76
|
+
N = ops.shape(x_dst)[self.node_dim] if x_dst is not None else ops.shape(x_src)[self.node_dim]
|
|
77
|
+
deg = degree(target_idx, num_nodes=N)
|
|
78
|
+
deg = ops.clip(deg, 0, self.max_degree)
|
|
79
|
+
|
|
80
|
+
h = self.propagate(edge_index, x=(x_src, x_dst), size=size)
|
|
81
|
+
|
|
82
|
+
out = ops.zeros((ops.shape(h)[0], self.out_channels), dtype=h.dtype)
|
|
83
|
+
for i, (lin_l, lin_r) in enumerate(zip(self.lins_l, self.lins_r)):
|
|
84
|
+
mask = ops.equal(deg, i)
|
|
85
|
+
mask = ops.expand_dims(ops.cast(mask, h.dtype), -1)
|
|
86
|
+
term = lin_l(h)
|
|
87
|
+
if x_dst is not None:
|
|
88
|
+
term = term + lin_r(x_dst)
|
|
89
|
+
out = out + mask * term
|
|
90
|
+
|
|
91
|
+
return out
|
|
92
|
+
|
|
93
|
+
def message(self, x_j):
|
|
94
|
+
return x_j
|
|
95
|
+
|
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
from typing import Optional, List
|
|
2
|
+
|
|
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.conv.utils import gcn_norm
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MixHopConv(MessagePassing):
|
|
11
|
+
r"""The MixHop graph convolutional operator from the
|
|
12
|
+
`"Higher-Order Graph Convolutional Networks via MixHop"
|
|
13
|
+
<https://arxiv.org/abs/1905.00067>`_ paper.
|
|
14
|
+
|
|
15
|
+
Example:
|
|
16
|
+
```python
|
|
17
|
+
import numpy as np
|
|
18
|
+
from k3_node.layers import MixHopConv
|
|
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
|
+
layer = MixHopConv(in_channels=8, out_channels=16, powers=[0, 1, 2])
|
|
24
|
+
out = layer(x, edge_index)
|
|
25
|
+
print(tuple(out.shape)) # (10, 48)
|
|
26
|
+
```
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
weighted_sum_message = True
|
|
30
|
+
def __init__(
|
|
31
|
+
self,
|
|
32
|
+
in_channels: int,
|
|
33
|
+
out_channels: int,
|
|
34
|
+
powers: Optional[List[int]] = None,
|
|
35
|
+
add_self_loops: bool = True,
|
|
36
|
+
bias: bool = True,
|
|
37
|
+
**kwargs,
|
|
38
|
+
):
|
|
39
|
+
kwargs.setdefault("aggr", "add")
|
|
40
|
+
super().__init__(**kwargs)
|
|
41
|
+
|
|
42
|
+
if powers is None:
|
|
43
|
+
powers = [0, 1, 2]
|
|
44
|
+
|
|
45
|
+
self.in_channels = in_channels
|
|
46
|
+
self.out_channels = out_channels
|
|
47
|
+
self.powers = powers
|
|
48
|
+
self.add_self_loops = add_self_loops
|
|
49
|
+
self.use_bias = bias
|
|
50
|
+
|
|
51
|
+
self.lins = [Dense(out_channels, use_bias=False) for _ in range(max(powers) + 1)]
|
|
52
|
+
|
|
53
|
+
if bias:
|
|
54
|
+
self.bias = self.add_weight(
|
|
55
|
+
shape=(len(powers) * out_channels,),
|
|
56
|
+
initializer="zeros",
|
|
57
|
+
name="bias",
|
|
58
|
+
)
|
|
59
|
+
else:
|
|
60
|
+
self.bias = None
|
|
61
|
+
|
|
62
|
+
def build(self, input_shape=None):
|
|
63
|
+
for lin in self.lins:
|
|
64
|
+
lin.build((None, self.in_channels))
|
|
65
|
+
self.built = True
|
|
66
|
+
|
|
67
|
+
def call(self, inputs, edge_index=None, edge_weight=None, **kwargs):
|
|
68
|
+
if edge_index is None:
|
|
69
|
+
if isinstance(inputs, (list, tuple)):
|
|
70
|
+
if len(inputs) == 3:
|
|
71
|
+
x, edge_index, edge_weight = inputs
|
|
72
|
+
elif len(inputs) == 2:
|
|
73
|
+
x, edge_index = inputs
|
|
74
|
+
else:
|
|
75
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
76
|
+
else:
|
|
77
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
78
|
+
else:
|
|
79
|
+
x = inputs
|
|
80
|
+
|
|
81
|
+
if not self.built:
|
|
82
|
+
self.build()
|
|
83
|
+
|
|
84
|
+
num_nodes = ops.shape(x)[0]
|
|
85
|
+
edge_index, edge_weight = gcn_norm(
|
|
86
|
+
edge_index,
|
|
87
|
+
edge_weight,
|
|
88
|
+
num_nodes=num_nodes,
|
|
89
|
+
add_self_loops=self.add_self_loops,
|
|
90
|
+
dtype=x.dtype,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
outs = [self.lins[0](x)]
|
|
94
|
+
curr_x = x
|
|
95
|
+
for lin in self.lins[1:]:
|
|
96
|
+
curr_x = self.propagate(edge_index, x=curr_x, edge_weight=edge_weight)
|
|
97
|
+
outs.append(lin(curr_x))
|
|
98
|
+
|
|
99
|
+
out = ops.concatenate([outs[p] for p in self.powers], axis=-1)
|
|
100
|
+
|
|
101
|
+
if self.bias is not None:
|
|
102
|
+
out = out + self.bias
|
|
103
|
+
|
|
104
|
+
return out
|
|
105
|
+
|
|
106
|
+
def message(self, x_j, edge_weight=None):
|
|
107
|
+
return x_j if edge_weight is None else ops.expand_dims(edge_weight, -1) * x_j
|
|
108
|
+
|