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,107 @@
|
|
|
1
|
+
from typing import Callable, Optional, 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.pool.knn import knn_graph
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class EdgeConv(MessagePassing):
|
|
10
|
+
r"""The edge convolutional operator from the `"Dynamic Graph CNN for
|
|
11
|
+
Learning on Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
nn: A neural network :math:`h_{\mathbf{\Theta}}` that maps
|
|
15
|
+
pair-wise node features to new edge representations.
|
|
16
|
+
aggr: The aggregation scheme to use (``"max"``, ``"mean"``, ``"sum"``).
|
|
17
|
+
(default: ``"max"``)
|
|
18
|
+
|
|
19
|
+
Example:
|
|
20
|
+
```python
|
|
21
|
+
import numpy as np
|
|
22
|
+
import keras
|
|
23
|
+
from k3_node.layers import EdgeConv
|
|
24
|
+
|
|
25
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
26
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
27
|
+
|
|
28
|
+
nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
|
|
29
|
+
layer = EdgeConv(nn)
|
|
30
|
+
out = layer(x, edge_index)
|
|
31
|
+
print(tuple(out.shape)) # (10, 16)
|
|
32
|
+
```
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
def __init__(self, nn: Callable, aggr: str = "max", **kwargs):
|
|
36
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
37
|
+
self.nn = nn
|
|
38
|
+
|
|
39
|
+
def build(self, input_shape):
|
|
40
|
+
if hasattr(self.nn, "build") and not getattr(self.nn, "built", False):
|
|
41
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
42
|
+
c = input_shape[0][-1]
|
|
43
|
+
elif isinstance(input_shape, (tuple, list)):
|
|
44
|
+
c = input_shape[-1]
|
|
45
|
+
else:
|
|
46
|
+
c = None
|
|
47
|
+
nn_shape = (None, 2 * c) if c is not None else None
|
|
48
|
+
self.nn.build(nn_shape)
|
|
49
|
+
self.built = True
|
|
50
|
+
|
|
51
|
+
def call(self, x, edge_index=None, training=None, **kwargs):
|
|
52
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
53
|
+
x, edge_index = x[0], x[1]
|
|
54
|
+
|
|
55
|
+
if not isinstance(x, (tuple, list)):
|
|
56
|
+
x_src, x_dst = x, x
|
|
57
|
+
else:
|
|
58
|
+
x_src, x_dst = x[0], x[1]
|
|
59
|
+
|
|
60
|
+
return self.propagate(edge_index, x=(x_src, x_dst), training=training, **kwargs)
|
|
61
|
+
|
|
62
|
+
def message(self, x_i, x_j, training=None):
|
|
63
|
+
h = ops.concatenate([x_i, x_j - x_i], axis=-1)
|
|
64
|
+
# Forward `training` explicitly: Keras does not propagate it to nested layers on JAX.
|
|
65
|
+
return self.nn(h, training=training) if isinstance(self.nn, keras.layers.Layer) else self.nn(h)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
class DynamicEdgeConv(EdgeConv):
|
|
69
|
+
r"""The dynamic edge convolutional operator from the `"Dynamic Graph CNN
|
|
70
|
+
for Learning on Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper,
|
|
71
|
+
which dynamically constructs a graph using :math:`k`-NN at each layer.
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
nn: A neural network :math:`h_{\mathbf{\Theta}}`.
|
|
75
|
+
k: Number of nearest neighbors. (default: ``6``)
|
|
76
|
+
aggr: The aggregation scheme to use (``"max"``, ``"mean"``, ``"sum"``).
|
|
77
|
+
(default: ``"max"``)
|
|
78
|
+
num_workers: Number of workers (ignored in Keras backend).
|
|
79
|
+
|
|
80
|
+
Example:
|
|
81
|
+
```python
|
|
82
|
+
import numpy as np
|
|
83
|
+
import keras
|
|
84
|
+
from k3_node.layers import DynamicEdgeConv
|
|
85
|
+
|
|
86
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
87
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
88
|
+
|
|
89
|
+
nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
|
|
90
|
+
layer = DynamicEdgeConv(nn, k=3)
|
|
91
|
+
out = layer(x) # k-NN graph is built from x
|
|
92
|
+
print(tuple(out.shape)) # (10, 16)
|
|
93
|
+
```
|
|
94
|
+
"""
|
|
95
|
+
|
|
96
|
+
def __init__(self, nn: Callable, k: int = 6, aggr: str = "max", num_workers: int = 1, **kwargs):
|
|
97
|
+
super().__init__(nn=nn, aggr=aggr, **kwargs)
|
|
98
|
+
self.k = k
|
|
99
|
+
|
|
100
|
+
def call(self, x, batch=None, training=None, **kwargs):
|
|
101
|
+
if isinstance(x, (tuple, list)):
|
|
102
|
+
x_src = x[0]
|
|
103
|
+
else:
|
|
104
|
+
x_src = x
|
|
105
|
+
|
|
106
|
+
edge_index = knn_graph(x_src, k=self.k, batch=batch, loop=False, flow=self.flow)
|
|
107
|
+
return super().call(x, edge_index=edge_index, training=training, **kwargs)
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
from typing import Optional, List
|
|
2
|
+
from keras import ops
|
|
3
|
+
from keras.layers import Dense
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.utils import gcn_norm, scatter
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class EGConv(MessagePassing):
|
|
10
|
+
r"""The Efficient Graph Convolution from the `"Adaptive Filters and
|
|
11
|
+
Aggregator Fusion for Efficient Graph Convolutions"
|
|
12
|
+
<https://arxiv.org/abs/2104.01481>`_ paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import EGConv
|
|
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
|
+
|
|
22
|
+
layer = EGConv(in_channels=8, out_channels=16, num_heads=2)
|
|
23
|
+
out = layer(x, edge_index)
|
|
24
|
+
print(tuple(out.shape)) # (10, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: int,
|
|
30
|
+
out_channels: int,
|
|
31
|
+
aggregators: Optional[List[str]] = None,
|
|
32
|
+
num_heads: int = 8,
|
|
33
|
+
num_bases: int = 4,
|
|
34
|
+
cached: bool = False,
|
|
35
|
+
add_self_loops: bool = True,
|
|
36
|
+
bias: bool = True,
|
|
37
|
+
**kwargs,
|
|
38
|
+
):
|
|
39
|
+
super().__init__(node_dim=0, **kwargs)
|
|
40
|
+
|
|
41
|
+
if out_channels % num_heads != 0:
|
|
42
|
+
raise ValueError(
|
|
43
|
+
f"'out_channels' ({out_channels}) must be divisible by num_heads ({num_heads})"
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
self.in_channels = in_channels
|
|
47
|
+
self.out_channels = out_channels
|
|
48
|
+
self.num_heads = num_heads
|
|
49
|
+
self.num_bases = num_bases
|
|
50
|
+
self.cached = cached
|
|
51
|
+
self.add_self_loops = add_self_loops
|
|
52
|
+
self.aggregators = aggregators or ["symnorm"]
|
|
53
|
+
self.use_bias = bias
|
|
54
|
+
|
|
55
|
+
self.bases_lin = Dense(
|
|
56
|
+
(out_channels // num_heads) * num_bases, use_bias=False
|
|
57
|
+
)
|
|
58
|
+
self.comb_lin = Dense(
|
|
59
|
+
num_heads * num_bases * len(self.aggregators), use_bias=True
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
if bias:
|
|
63
|
+
self.bias = self.add_weight(
|
|
64
|
+
shape=(out_channels,),
|
|
65
|
+
initializer="zeros",
|
|
66
|
+
name="bias",
|
|
67
|
+
)
|
|
68
|
+
else:
|
|
69
|
+
self.bias = None
|
|
70
|
+
|
|
71
|
+
def build(self, input_shape=None):
|
|
72
|
+
self.bases_lin.build((None, self.in_channels))
|
|
73
|
+
self.comb_lin.build((None, self.in_channels))
|
|
74
|
+
self.built = True
|
|
75
|
+
|
|
76
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
77
|
+
if edge_index is None:
|
|
78
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
79
|
+
x, edge_index = inputs
|
|
80
|
+
else:
|
|
81
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
82
|
+
else:
|
|
83
|
+
x = inputs
|
|
84
|
+
|
|
85
|
+
if not self.built:
|
|
86
|
+
self.build()
|
|
87
|
+
|
|
88
|
+
num_nodes = ops.shape(x)[0]
|
|
89
|
+
symnorm_weight = None
|
|
90
|
+
if "symnorm" in self.aggregators:
|
|
91
|
+
edge_index, symnorm_weight = gcn_norm(
|
|
92
|
+
edge_index,
|
|
93
|
+
edge_weight=None,
|
|
94
|
+
num_nodes=num_nodes,
|
|
95
|
+
add_self_loops=self.add_self_loops,
|
|
96
|
+
dtype=x.dtype,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
bases = self.bases_lin(x)
|
|
100
|
+
weightings = self.comb_lin(x)
|
|
101
|
+
|
|
102
|
+
aggregated = self.propagate(
|
|
103
|
+
edge_index,
|
|
104
|
+
x=bases,
|
|
105
|
+
symnorm_weight=symnorm_weight,
|
|
106
|
+
size=(num_nodes, num_nodes),
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
weightings = ops.reshape(
|
|
110
|
+
weightings,
|
|
111
|
+
(-1, self.num_heads, self.num_bases * len(self.aggregators)),
|
|
112
|
+
)
|
|
113
|
+
aggregated = ops.reshape(
|
|
114
|
+
aggregated,
|
|
115
|
+
(
|
|
116
|
+
-1,
|
|
117
|
+
len(self.aggregators) * self.num_bases,
|
|
118
|
+
self.out_channels // self.num_heads,
|
|
119
|
+
),
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
out = ops.matmul(weightings, aggregated)
|
|
123
|
+
out = ops.reshape(out, (-1, self.out_channels))
|
|
124
|
+
|
|
125
|
+
if self.bias is not None:
|
|
126
|
+
out = out + self.bias
|
|
127
|
+
|
|
128
|
+
return out
|
|
129
|
+
|
|
130
|
+
def message(self, x_j):
|
|
131
|
+
return x_j
|
|
132
|
+
|
|
133
|
+
def aggregate(self, inputs, edge_index=None, index=None, dim_size=None, symnorm_weight=None, **kwargs):
|
|
134
|
+
if index is None and edge_index is not None:
|
|
135
|
+
index = edge_index[1]
|
|
136
|
+
|
|
137
|
+
outs = []
|
|
138
|
+
for aggr in self.aggregators:
|
|
139
|
+
if aggr == "symnorm":
|
|
140
|
+
inp = inputs if symnorm_weight is None else inputs * ops.expand_dims(symnorm_weight, -1)
|
|
141
|
+
out = scatter(inp, index, dim=0, dim_size=dim_size, reduce="sum")
|
|
142
|
+
elif aggr in ("var", "std"):
|
|
143
|
+
mean = scatter(inputs, index, dim=0, dim_size=dim_size, reduce="mean")
|
|
144
|
+
mean_sq = scatter(inputs * inputs, index, dim=0, dim_size=dim_size, reduce="mean")
|
|
145
|
+
out = mean_sq - mean * mean
|
|
146
|
+
if aggr == "std":
|
|
147
|
+
out = ops.sqrt(ops.maximum(out, 1e-5))
|
|
148
|
+
else:
|
|
149
|
+
out = scatter(inputs, index, dim=0, dim_size=dim_size, reduce=aggr)
|
|
150
|
+
outs.append(out)
|
|
151
|
+
|
|
152
|
+
if len(outs) > 1:
|
|
153
|
+
return ops.stack(outs, axis=1)
|
|
154
|
+
return outs[0]
|
|
155
|
+
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
from keras.layers import Dense
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class FAConv(MessagePassing):
|
|
9
|
+
r"""The Frequency Adaptive Graph Convolution operator from the
|
|
10
|
+
`"Beyond Low-Frequency Information in Graph Convolutional Networks"
|
|
11
|
+
<https://arxiv.org/abs/2101.00797>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import FAConv
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
20
|
+
|
|
21
|
+
x_0 = x # initial node representations
|
|
22
|
+
layer = FAConv(channels=8, eps=0.1)
|
|
23
|
+
out = layer(x, x_0, edge_index)
|
|
24
|
+
print(tuple(out.shape)) # (10, 8)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
channels: int,
|
|
30
|
+
eps: float = 0.1,
|
|
31
|
+
dropout: float = 0.0,
|
|
32
|
+
cached: bool = False,
|
|
33
|
+
add_self_loops: bool = True,
|
|
34
|
+
normalize: bool = True,
|
|
35
|
+
**kwargs,
|
|
36
|
+
):
|
|
37
|
+
kwargs.setdefault("aggr", "add")
|
|
38
|
+
super().__init__(**kwargs)
|
|
39
|
+
|
|
40
|
+
self.channels = channels
|
|
41
|
+
self.eps = eps
|
|
42
|
+
self.dropout_rate = dropout
|
|
43
|
+
self.cached = cached
|
|
44
|
+
self.add_self_loops = add_self_loops
|
|
45
|
+
self.normalize = normalize
|
|
46
|
+
|
|
47
|
+
self.att_l = Dense(1, use_bias=False)
|
|
48
|
+
self.att_r = Dense(1, use_bias=False)
|
|
49
|
+
|
|
50
|
+
def build(self, input_shape=None):
|
|
51
|
+
self.att_l.build((None, self.channels))
|
|
52
|
+
self.att_r.build((None, self.channels))
|
|
53
|
+
self.built = True
|
|
54
|
+
|
|
55
|
+
def call(self, inputs, x_0=None, edge_index=None, edge_weight=None, **kwargs):
|
|
56
|
+
if edge_index is None:
|
|
57
|
+
if isinstance(inputs, (list, tuple)):
|
|
58
|
+
if len(inputs) == 4:
|
|
59
|
+
x, x_0, edge_index, edge_weight = inputs
|
|
60
|
+
elif len(inputs) == 3:
|
|
61
|
+
x, x_0, edge_index = inputs
|
|
62
|
+
elif len(inputs) == 2:
|
|
63
|
+
x, edge_index = inputs
|
|
64
|
+
x_0 = x
|
|
65
|
+
else:
|
|
66
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
67
|
+
else:
|
|
68
|
+
raise ValueError("Expected inputs with edge_index")
|
|
69
|
+
else:
|
|
70
|
+
x = inputs
|
|
71
|
+
if x_0 is None:
|
|
72
|
+
x_0 = x
|
|
73
|
+
|
|
74
|
+
if not self.built:
|
|
75
|
+
self.build()
|
|
76
|
+
|
|
77
|
+
num_nodes = ops.shape(x)[0]
|
|
78
|
+
if self.normalize:
|
|
79
|
+
edge_index, edge_weight = gcn_norm(
|
|
80
|
+
edge_index,
|
|
81
|
+
edge_weight,
|
|
82
|
+
num_nodes=num_nodes,
|
|
83
|
+
add_self_loops=self.add_self_loops,
|
|
84
|
+
dtype=x.dtype,
|
|
85
|
+
)
|
|
86
|
+
|
|
87
|
+
alpha_l = self.att_l(x)
|
|
88
|
+
alpha_r = self.att_r(x)
|
|
89
|
+
|
|
90
|
+
out = self.propagate(
|
|
91
|
+
edge_index,
|
|
92
|
+
x=x,
|
|
93
|
+
alpha=(alpha_l, alpha_r),
|
|
94
|
+
edge_weight=edge_weight,
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
if self.eps != 0.0:
|
|
98
|
+
out = out + self.eps * x_0
|
|
99
|
+
|
|
100
|
+
return out
|
|
101
|
+
|
|
102
|
+
def message(self, x_j, alpha_j, alpha_i, edge_weight=None):
|
|
103
|
+
alpha = ops.tanh(alpha_j + alpha_i)
|
|
104
|
+
if edge_weight is not None:
|
|
105
|
+
alpha = alpha * ops.expand_dims(edge_weight, -1)
|
|
106
|
+
return x_j * alpha
|
|
107
|
+
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
from keras.layers import Dense
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import (
|
|
6
|
+
add_self_loops,
|
|
7
|
+
degree,
|
|
8
|
+
extend_mask_for_self_loops,
|
|
9
|
+
remove_self_loops_masked,
|
|
10
|
+
)
|
|
11
|
+
from k3_node.ops.segment import segment_sum
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class FeaStConv(MessagePassing):
|
|
15
|
+
r"""The (fault-tolerant) feature-steered graph convolution operator from
|
|
16
|
+
the `"FeaStNet: Feature-Steered Graph Convolutions for 3D Shape Analysis"
|
|
17
|
+
<https://arxiv.org/abs/1706.05206>`_ paper.
|
|
18
|
+
|
|
19
|
+
Example:
|
|
20
|
+
```python
|
|
21
|
+
import numpy as np
|
|
22
|
+
from k3_node.layers import FeaStConv
|
|
23
|
+
|
|
24
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
25
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
26
|
+
|
|
27
|
+
layer = FeaStConv(in_channels=8, out_channels=16, heads=2)
|
|
28
|
+
out = layer(x, edge_index)
|
|
29
|
+
print(tuple(out.shape)) # (10, 16)
|
|
30
|
+
```
|
|
31
|
+
"""
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
in_channels: int,
|
|
35
|
+
out_channels: int,
|
|
36
|
+
heads: int = 1,
|
|
37
|
+
add_self_loops: bool = True,
|
|
38
|
+
bias: bool = True,
|
|
39
|
+
**kwargs,
|
|
40
|
+
):
|
|
41
|
+
kwargs.setdefault("aggr", "mean")
|
|
42
|
+
super().__init__(node_dim=0, **kwargs)
|
|
43
|
+
|
|
44
|
+
self.in_channels = in_channels
|
|
45
|
+
self.out_channels = out_channels
|
|
46
|
+
self.heads = heads
|
|
47
|
+
self.add_self_loops = add_self_loops
|
|
48
|
+
self.use_bias = bias
|
|
49
|
+
|
|
50
|
+
self.lin = Dense(heads * out_channels, use_bias=False)
|
|
51
|
+
self.u = self.add_weight(
|
|
52
|
+
shape=(in_channels, heads),
|
|
53
|
+
initializer="glorot_uniform",
|
|
54
|
+
name="u",
|
|
55
|
+
)
|
|
56
|
+
self.c = self.add_weight(
|
|
57
|
+
shape=(heads,),
|
|
58
|
+
initializer="zeros",
|
|
59
|
+
name="c",
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
if bias:
|
|
63
|
+
self.bias = self.add_weight(
|
|
64
|
+
shape=(out_channels,),
|
|
65
|
+
initializer="zeros",
|
|
66
|
+
name="bias",
|
|
67
|
+
)
|
|
68
|
+
else:
|
|
69
|
+
self.bias = None
|
|
70
|
+
|
|
71
|
+
def build(self, input_shape=None):
|
|
72
|
+
self.lin.build((None, self.in_channels))
|
|
73
|
+
self.built = True
|
|
74
|
+
|
|
75
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
76
|
+
if edge_index is None:
|
|
77
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
78
|
+
x, edge_index = inputs
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
81
|
+
else:
|
|
82
|
+
x = inputs
|
|
83
|
+
|
|
84
|
+
if not self.built:
|
|
85
|
+
self.build()
|
|
86
|
+
|
|
87
|
+
if isinstance(x, (list, tuple)):
|
|
88
|
+
x_src, x_dst = x
|
|
89
|
+
else:
|
|
90
|
+
x_src = x_dst = x
|
|
91
|
+
|
|
92
|
+
num_nodes = ops.shape(x_dst)[0]
|
|
93
|
+
keep_mask = None
|
|
94
|
+
if self.add_self_loops:
|
|
95
|
+
edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
|
|
96
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
97
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
98
|
+
|
|
99
|
+
out = self.propagate(
|
|
100
|
+
edge_index,
|
|
101
|
+
x=(x_src, x_dst),
|
|
102
|
+
keep_mask=keep_mask,
|
|
103
|
+
size=(ops.shape(x_src)[0], num_nodes),
|
|
104
|
+
)
|
|
105
|
+
if keep_mask is not None and self.aggr == "mean":
|
|
106
|
+
# Masked messages are zero but still counted by the mean; rescale to the kept count.
|
|
107
|
+
col = ops.cast(edge_index[1], "int32")
|
|
108
|
+
count_all = degree(col, num_nodes=num_nodes, dtype=out.dtype)
|
|
109
|
+
count_kept = segment_sum(ops.cast(keep_mask, out.dtype), col, num_segments=num_nodes)
|
|
110
|
+
out = out * ops.expand_dims(count_all / ops.maximum(count_kept, 1.0), -1)
|
|
111
|
+
|
|
112
|
+
if self.bias is not None:
|
|
113
|
+
out = out + self.bias
|
|
114
|
+
|
|
115
|
+
return out
|
|
116
|
+
|
|
117
|
+
def message(self, x_i, x_j, keep_mask=None):
|
|
118
|
+
q = ops.matmul(x_j - x_i, self.u) + self.c
|
|
119
|
+
q = ops.softmax(q, axis=-1)
|
|
120
|
+
x_j_mapped = ops.reshape(self.lin(x_j), (-1, self.heads, self.out_channels))
|
|
121
|
+
msg = ops.sum(x_j_mapped * ops.expand_dims(q, -1), axis=1)
|
|
122
|
+
# For max/min a duplicated self-loop message is harmless, so only sum/mean need masking.
|
|
123
|
+
if keep_mask is not None and self.aggr in ("add", "sum", "mean"):
|
|
124
|
+
msg = msg * ops.expand_dims(ops.cast(keep_mask, msg.dtype), -1)
|
|
125
|
+
return msg
|
|
126
|
+
|
|
@@ -0,0 +1,143 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Optional, Union, Tuple, Callable
|
|
3
|
+
from keras import ops, activations
|
|
4
|
+
from keras.layers import Dense
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class FiLMConv(MessagePassing):
|
|
10
|
+
r"""The FiLM graph convolutional operator from the
|
|
11
|
+
`"GNN-FiLM: Graph Neural Networks with Feature-wise Linear Modulation"
|
|
12
|
+
<https://arxiv.org/abs/1906.12192>`_ paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import FiLMConv
|
|
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
|
+
|
|
22
|
+
layer = FiLMConv(in_channels=8, out_channels=16)
|
|
23
|
+
out = layer(x, edge_index)
|
|
24
|
+
print(tuple(out.shape)) # (10, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
30
|
+
out_channels: int,
|
|
31
|
+
num_relations: int = 1,
|
|
32
|
+
nn: Optional[Callable] = None,
|
|
33
|
+
act: Optional[Union[str, Callable]] = "relu",
|
|
34
|
+
aggr: str = "mean",
|
|
35
|
+
**kwargs,
|
|
36
|
+
):
|
|
37
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
38
|
+
|
|
39
|
+
self.in_channels = in_channels
|
|
40
|
+
self.out_channels = out_channels
|
|
41
|
+
self.num_relations = max(num_relations, 1)
|
|
42
|
+
self.act = activations.get(act) if act is not None else None
|
|
43
|
+
|
|
44
|
+
if isinstance(in_channels, int):
|
|
45
|
+
self.in_channels_l = in_channels
|
|
46
|
+
self.in_channels_r = in_channels
|
|
47
|
+
else:
|
|
48
|
+
self.in_channels_l, self.in_channels_r = in_channels
|
|
49
|
+
|
|
50
|
+
self.lins = [
|
|
51
|
+
Dense(out_channels, use_bias=False) for _ in range(self.num_relations)
|
|
52
|
+
]
|
|
53
|
+
self.films = []
|
|
54
|
+
for _ in range(self.num_relations):
|
|
55
|
+
if nn is None:
|
|
56
|
+
self.films.append(Dense(2 * out_channels, use_bias=True))
|
|
57
|
+
else:
|
|
58
|
+
self.films.append(copy.deepcopy(nn))
|
|
59
|
+
|
|
60
|
+
self.lin_skip = Dense(out_channels, use_bias=False)
|
|
61
|
+
if nn is None:
|
|
62
|
+
self.film_skip = Dense(2 * out_channels, use_bias=False)
|
|
63
|
+
else:
|
|
64
|
+
self.film_skip = copy.deepcopy(nn)
|
|
65
|
+
|
|
66
|
+
def build(self, input_shape=None):
|
|
67
|
+
for lin in self.lins:
|
|
68
|
+
lin.build((None, self.in_channels_l))
|
|
69
|
+
for film in self.films:
|
|
70
|
+
film.build((None, self.in_channels_r))
|
|
71
|
+
self.lin_skip.build((None, self.in_channels_r))
|
|
72
|
+
self.film_skip.build((None, self.in_channels_r))
|
|
73
|
+
self.built = True
|
|
74
|
+
|
|
75
|
+
def call(self, inputs, edge_index=None, edge_type=None, **kwargs):
|
|
76
|
+
if edge_index is None:
|
|
77
|
+
if isinstance(inputs, (list, tuple)):
|
|
78
|
+
if len(inputs) == 3:
|
|
79
|
+
x, edge_index, edge_type = inputs
|
|
80
|
+
elif len(inputs) == 2:
|
|
81
|
+
x, edge_index = inputs
|
|
82
|
+
else:
|
|
83
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
84
|
+
else:
|
|
85
|
+
raise ValueError("Expected inputs with edge_index")
|
|
86
|
+
else:
|
|
87
|
+
x = inputs
|
|
88
|
+
|
|
89
|
+
if not self.built:
|
|
90
|
+
self.build()
|
|
91
|
+
|
|
92
|
+
if isinstance(x, (list, tuple)):
|
|
93
|
+
x_l, x_r = x
|
|
94
|
+
else:
|
|
95
|
+
x_l = x_r = x
|
|
96
|
+
|
|
97
|
+
edge_index = ops.cast(edge_index, "int32")
|
|
98
|
+
if edge_type is not None:
|
|
99
|
+
edge_type = ops.cast(edge_type, "int32")
|
|
100
|
+
|
|
101
|
+
# Skip connection
|
|
102
|
+
film_s = self.film_skip(x_r)
|
|
103
|
+
beta_s, gamma_s = ops.split(film_s, 2, axis=-1)
|
|
104
|
+
out = gamma_s * self.lin_skip(x_r) + beta_s
|
|
105
|
+
if self.act is not None:
|
|
106
|
+
out = self.act(out)
|
|
107
|
+
|
|
108
|
+
# Message passing per relation
|
|
109
|
+
num_nodes = ops.shape(x_r)[0]
|
|
110
|
+
size = (ops.shape(x_l)[0], num_nodes)
|
|
111
|
+
|
|
112
|
+
for i in range(self.num_relations):
|
|
113
|
+
if edge_type is not None:
|
|
114
|
+
mask = ops.equal(edge_type, i)
|
|
115
|
+
where_mask = ops.where(mask)
|
|
116
|
+
idx = where_mask[0] if isinstance(where_mask, (list, tuple)) else where_mask
|
|
117
|
+
idx = ops.reshape(idx, (-1,))
|
|
118
|
+
idx = ops.cast(idx, "int32")
|
|
119
|
+
if ops.shape(idx)[0] == 0:
|
|
120
|
+
continue
|
|
121
|
+
edge_index_i = ops.take(edge_index, idx, axis=1)
|
|
122
|
+
else:
|
|
123
|
+
edge_index_i = edge_index
|
|
124
|
+
|
|
125
|
+
film_val = self.films[i](x_r)
|
|
126
|
+
beta, gamma = ops.split(film_val, 2, axis=-1)
|
|
127
|
+
h = self.propagate(
|
|
128
|
+
edge_index_i,
|
|
129
|
+
x=self.lins[i](x_l),
|
|
130
|
+
beta=beta,
|
|
131
|
+
gamma=gamma,
|
|
132
|
+
size=size,
|
|
133
|
+
)
|
|
134
|
+
out = out + h
|
|
135
|
+
|
|
136
|
+
return out
|
|
137
|
+
|
|
138
|
+
def message(self, x_j, beta_i, gamma_i):
|
|
139
|
+
out = gamma_i * x_j + beta_i
|
|
140
|
+
if self.act is not None:
|
|
141
|
+
out = self.act(out)
|
|
142
|
+
return out
|
|
143
|
+
|