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,154 @@
|
|
|
1
|
+
from typing import Optional, Union, Tuple, List
|
|
2
|
+
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.aggr.base import Aggregation
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class SAGEConv(MessagePassing):
|
|
10
|
+
r"""The GraphSAGE operator from the `"Inductive Representation Learning on
|
|
11
|
+
Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
15
|
+
out_channels: Size of each output sample.
|
|
16
|
+
aggr: The aggregation scheme to use (``"mean"``, ``"max"``,
|
|
17
|
+
``"lstm"``, etc.). (default: ``"mean"``)
|
|
18
|
+
normalize: If set to :obj:`True`, output features will be
|
|
19
|
+
:math:`\ell_2`-normalized. (default: :obj:`False`)
|
|
20
|
+
root_weight: If set to :obj:`False`, the layer will not add
|
|
21
|
+
the transformed root node features. (default: :obj:`True`)
|
|
22
|
+
project: If set to :obj:`True`, the layer will apply a linear
|
|
23
|
+
transformation followed by an activation to source node features.
|
|
24
|
+
(default: :obj:`False`)
|
|
25
|
+
bias: If set to :obj:`False`, the layer will not learn
|
|
26
|
+
an additive bias. (default: :obj:`True`)
|
|
27
|
+
|
|
28
|
+
Example:
|
|
29
|
+
```python
|
|
30
|
+
import numpy as np
|
|
31
|
+
from k3_node.layers import SAGEConv
|
|
32
|
+
|
|
33
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
34
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
35
|
+
|
|
36
|
+
layer = SAGEConv(in_channels=8, out_channels=16)
|
|
37
|
+
out = layer(x, edge_index)
|
|
38
|
+
print(tuple(out.shape)) # (10, 16)
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
def __init__(
|
|
43
|
+
self,
|
|
44
|
+
in_channels: Union[int, Tuple[int, int], None],
|
|
45
|
+
out_channels: Optional[int] = None,
|
|
46
|
+
aggr: Optional[Union[str, List[str], Aggregation]] = "mean",
|
|
47
|
+
normalize: bool = False,
|
|
48
|
+
root_weight: bool = True,
|
|
49
|
+
project: bool = False,
|
|
50
|
+
bias: bool = True,
|
|
51
|
+
**kwargs,
|
|
52
|
+
):
|
|
53
|
+
if out_channels is None:
|
|
54
|
+
# Backward compatibility: SAGEConv(out_channels, ...)
|
|
55
|
+
out_channels = in_channels
|
|
56
|
+
in_channels = None
|
|
57
|
+
|
|
58
|
+
self.in_channels = in_channels
|
|
59
|
+
self.out_channels = out_channels
|
|
60
|
+
self.normalize = normalize
|
|
61
|
+
self.root_weight = root_weight
|
|
62
|
+
self.project = project
|
|
63
|
+
self.use_bias = bias
|
|
64
|
+
|
|
65
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
66
|
+
|
|
67
|
+
if self.project:
|
|
68
|
+
self.lin_proj = layers.Dense(
|
|
69
|
+
in_channels[0] if isinstance(in_channels, (tuple, list)) else in_channels,
|
|
70
|
+
activation="relu",
|
|
71
|
+
use_bias=True,
|
|
72
|
+
)
|
|
73
|
+
else:
|
|
74
|
+
self.lin_proj = None
|
|
75
|
+
|
|
76
|
+
self.lin_l = layers.Dense(out_channels, use_bias=bias)
|
|
77
|
+
if self.root_weight:
|
|
78
|
+
self.lin_r = layers.Dense(out_channels, use_bias=False)
|
|
79
|
+
else:
|
|
80
|
+
self.lin_r = None
|
|
81
|
+
|
|
82
|
+
def build(self, input_shape):
|
|
83
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
84
|
+
in_channels_l = input_shape[0][-1]
|
|
85
|
+
in_channels_r = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_l
|
|
86
|
+
else:
|
|
87
|
+
in_channels_l = input_shape[-1]
|
|
88
|
+
in_channels_r = input_shape[-1]
|
|
89
|
+
|
|
90
|
+
self.lin_l.build((None, self.out_channels if self.project else in_channels_l))
|
|
91
|
+
if self.lin_r is not None:
|
|
92
|
+
self.lin_r.build((None, in_channels_r))
|
|
93
|
+
if self.lin_proj is not None:
|
|
94
|
+
self.lin_proj.build((None, in_channels_l))
|
|
95
|
+
self.built = True
|
|
96
|
+
|
|
97
|
+
def call(self, x, edge_index=None, size=None, **kwargs):
|
|
98
|
+
# Handle legacy calling: conv(x, adj) where adj is [N, N]
|
|
99
|
+
shape = getattr(edge_index, "shape", None)
|
|
100
|
+
if (
|
|
101
|
+
shape is not None
|
|
102
|
+
and len(shape) == 2
|
|
103
|
+
and shape[0] is not None
|
|
104
|
+
and shape[1] is not None
|
|
105
|
+
and shape[0] > 2
|
|
106
|
+
and shape[0] == shape[1]
|
|
107
|
+
):
|
|
108
|
+
where_adj = ops.where(edge_index != 0)
|
|
109
|
+
where_adj = where_adj if not isinstance(where_adj, list) else where_adj
|
|
110
|
+
edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
|
|
111
|
+
|
|
112
|
+
# Handle legacy calling: conv((x, adj))
|
|
113
|
+
if edge_index is None and isinstance(x, (tuple, list)) and len(x) == 2:
|
|
114
|
+
arg0, arg1 = x[0], x[1]
|
|
115
|
+
s1 = getattr(arg1, "shape", None)
|
|
116
|
+
if (
|
|
117
|
+
s1 is not None
|
|
118
|
+
and len(s1) == 2
|
|
119
|
+
and s1[0] is not None
|
|
120
|
+
and s1[1] is not None
|
|
121
|
+
and s1[0] > 2
|
|
122
|
+
and s1[0] == s1[1]
|
|
123
|
+
):
|
|
124
|
+
where_adj = ops.where(arg1 != 0)
|
|
125
|
+
where_adj = where_adj if not isinstance(where_adj, list) else where_adj
|
|
126
|
+
edge_index = ops.stack([where_adj[0], where_adj[1]], axis=0)
|
|
127
|
+
x = arg0
|
|
128
|
+
elif s1 is not None and len(s1) >= 1 and s1[0] == 2:
|
|
129
|
+
edge_index = arg1
|
|
130
|
+
x = arg0
|
|
131
|
+
|
|
132
|
+
if not isinstance(x, (tuple, list)):
|
|
133
|
+
x_src = x
|
|
134
|
+
x_dst = x
|
|
135
|
+
else:
|
|
136
|
+
x_src, x_dst = x[0], x[1]
|
|
137
|
+
|
|
138
|
+
if self.project and self.lin_proj is not None:
|
|
139
|
+
x_src = self.lin_proj(x_src)
|
|
140
|
+
|
|
141
|
+
out = self.propagate(edge_index, x=(x_src, x_dst), size=size)
|
|
142
|
+
out = self.lin_l(out)
|
|
143
|
+
|
|
144
|
+
if self.root_weight and self.lin_r is not None and x_dst is not None:
|
|
145
|
+
out = out + self.lin_r(x_dst)
|
|
146
|
+
|
|
147
|
+
if self.normalize:
|
|
148
|
+
norm = ops.sqrt(ops.sum(ops.square(out), axis=-1, keepdims=True) + 1e-12)
|
|
149
|
+
out = out / norm
|
|
150
|
+
|
|
151
|
+
return out
|
|
152
|
+
|
|
153
|
+
def message(self, x_j):
|
|
154
|
+
return x_j
|
|
@@ -0,0 +1,96 @@
|
|
|
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 SGConv(MessagePassing):
|
|
8
|
+
r"""The simple graph convolutional operator from the `"Simplifying Graph
|
|
9
|
+
Convolutional Networks" <https://arxiv.org/abs/1902.07153>`_ paper.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
in_channels: Size of each input sample.
|
|
13
|
+
out_channels: Size of each output sample.
|
|
14
|
+
K: Number of hops :math:`K`. (default: ``1``)
|
|
15
|
+
cached: If set to :obj:`True`, the layer will cache the computation of
|
|
16
|
+
:math:`\mathbf{\hat{D}}^{-1/2} \mathbf{\hat{A}} \mathbf{\hat{D}}^{-1/2}`.
|
|
17
|
+
(default: ``False``)
|
|
18
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops.
|
|
19
|
+
(default: ``True``)
|
|
20
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
21
|
+
(default: ``True``)
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
```python
|
|
25
|
+
import numpy as np
|
|
26
|
+
from k3_node.layers import SGConv
|
|
27
|
+
|
|
28
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
29
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
30
|
+
|
|
31
|
+
layer = SGConv(in_channels=8, out_channels=16, K=2)
|
|
32
|
+
out = layer(x, edge_index)
|
|
33
|
+
print(tuple(out.shape)) # (10, 16)
|
|
34
|
+
```
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
weighted_sum_message = True
|
|
38
|
+
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
in_channels: int,
|
|
42
|
+
out_channels: int,
|
|
43
|
+
K: int = 1,
|
|
44
|
+
cached: bool = False,
|
|
45
|
+
add_self_loops: bool = True,
|
|
46
|
+
bias: bool = True,
|
|
47
|
+
**kwargs,
|
|
48
|
+
):
|
|
49
|
+
super().__init__(aggr="add", **kwargs)
|
|
50
|
+
self.in_channels = in_channels
|
|
51
|
+
self.out_channels = out_channels
|
|
52
|
+
self.K = K
|
|
53
|
+
self.cached = cached
|
|
54
|
+
self.add_self_loops = add_self_loops
|
|
55
|
+
self.use_bias = bias
|
|
56
|
+
|
|
57
|
+
self.lin = layers.Dense(out_channels, use_bias=bias)
|
|
58
|
+
self._cached_edge_index = None
|
|
59
|
+
self._cached_norm = None
|
|
60
|
+
|
|
61
|
+
def build(self, input_shape):
|
|
62
|
+
feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
|
|
63
|
+
self.lin.build(feat_shape)
|
|
64
|
+
self.built = True
|
|
65
|
+
|
|
66
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
67
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
68
|
+
x, edge_index = x[0], x[1]
|
|
69
|
+
|
|
70
|
+
if self.cached and self._cached_edge_index is not None:
|
|
71
|
+
edge_index = self._cached_edge_index
|
|
72
|
+
edge_weight = self._cached_norm
|
|
73
|
+
else:
|
|
74
|
+
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]
|
|
75
|
+
edge_index, edge_weight = gcn_norm(
|
|
76
|
+
edge_index,
|
|
77
|
+
edge_weight,
|
|
78
|
+
num_nodes=num_nodes,
|
|
79
|
+
add_self_loops=self.add_self_loops,
|
|
80
|
+
flow=self.flow,
|
|
81
|
+
dtype=x.dtype,
|
|
82
|
+
)
|
|
83
|
+
if self.cached and not is_tracing(edge_index):
|
|
84
|
+
self._cached_edge_index = edge_index
|
|
85
|
+
self._cached_norm = edge_weight
|
|
86
|
+
|
|
87
|
+
for _ in range(self.K):
|
|
88
|
+
x = self.propagate(edge_index, x=x, edge_weight=edge_weight)
|
|
89
|
+
|
|
90
|
+
return self.lin(x)
|
|
91
|
+
|
|
92
|
+
def message(self, x_j, edge_weight=None):
|
|
93
|
+
if edge_weight is None:
|
|
94
|
+
return x_j
|
|
95
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
96
|
+
|
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
from keras.layers import Dense
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class SignedConv(MessagePassing):
|
|
8
|
+
r"""The signed graph convolutional operator from the `"Signed Graph
|
|
9
|
+
Convolutional Network" <https://arxiv.org/abs/1808.06354>`_ paper.
|
|
10
|
+
|
|
11
|
+
Example:
|
|
12
|
+
```python
|
|
13
|
+
import numpy as np
|
|
14
|
+
from k3_node.layers import SignedConv
|
|
15
|
+
|
|
16
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
17
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
18
|
+
|
|
19
|
+
pos_edge_index = edge_index[:, :15] # positive (e.g. "trust") edges
|
|
20
|
+
neg_edge_index = edge_index[:, 15:] # negative (e.g. "distrust") edges
|
|
21
|
+
layer = SignedConv(in_channels=8, out_channels=16, first_aggr=True)
|
|
22
|
+
out = layer(x, pos_edge_index, neg_edge_index)
|
|
23
|
+
print(tuple(out.shape)) # (10, 32): positive and negative embeddings, concatenated
|
|
24
|
+
```
|
|
25
|
+
"""
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
in_channels: int,
|
|
29
|
+
out_channels: int,
|
|
30
|
+
first_aggr: bool,
|
|
31
|
+
bias: bool = True,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
kwargs.setdefault("aggr", "mean")
|
|
35
|
+
super().__init__(**kwargs)
|
|
36
|
+
|
|
37
|
+
self.in_channels = in_channels
|
|
38
|
+
self.out_channels = out_channels
|
|
39
|
+
self.first_aggr = first_aggr
|
|
40
|
+
self.use_bias = bias
|
|
41
|
+
|
|
42
|
+
in_pos = in_channels if first_aggr else 2 * in_channels
|
|
43
|
+
self.lin_pos_l = Dense(out_channels, use_bias=False)
|
|
44
|
+
self.lin_pos_r = Dense(out_channels, use_bias=bias)
|
|
45
|
+
self.lin_neg_l = Dense(out_channels, use_bias=False)
|
|
46
|
+
self.lin_neg_r = Dense(out_channels, use_bias=bias)
|
|
47
|
+
|
|
48
|
+
def build(self, input_shape=None):
|
|
49
|
+
in_dim = self.in_channels if self.first_aggr else 2 * self.in_channels
|
|
50
|
+
self.lin_pos_l.build((None, in_dim))
|
|
51
|
+
self.lin_pos_r.build((None, self.in_channels))
|
|
52
|
+
self.lin_neg_l.build((None, in_dim))
|
|
53
|
+
self.lin_neg_r.build((None, self.in_channels))
|
|
54
|
+
self.built = True
|
|
55
|
+
|
|
56
|
+
def call(self, inputs, pos_edge_index=None, neg_edge_index=None, **kwargs):
|
|
57
|
+
if pos_edge_index is None:
|
|
58
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 3:
|
|
59
|
+
x, pos_edge_index, neg_edge_index = inputs
|
|
60
|
+
else:
|
|
61
|
+
raise ValueError("Expected (x, pos_edge_index, neg_edge_index)")
|
|
62
|
+
else:
|
|
63
|
+
x = inputs
|
|
64
|
+
|
|
65
|
+
if not self.built:
|
|
66
|
+
self.build()
|
|
67
|
+
|
|
68
|
+
if isinstance(x, (list, tuple)):
|
|
69
|
+
x_src, x_dst = x
|
|
70
|
+
else:
|
|
71
|
+
x_src = x_dst = x
|
|
72
|
+
|
|
73
|
+
if self.first_aggr:
|
|
74
|
+
out_pos = self.propagate(pos_edge_index, x=(x_src, x_dst))
|
|
75
|
+
out_pos = self.lin_pos_l(out_pos) + self.lin_pos_r(x_dst)
|
|
76
|
+
|
|
77
|
+
out_neg = self.propagate(neg_edge_index, x=(x_src, x_dst))
|
|
78
|
+
out_neg = self.lin_neg_l(out_neg) + self.lin_neg_r(x_dst)
|
|
79
|
+
|
|
80
|
+
return ops.concatenate([out_pos, out_neg], axis=-1)
|
|
81
|
+
else:
|
|
82
|
+
F_in = self.in_channels
|
|
83
|
+
x_src_1, x_src_2 = x_src[..., :F_in], x_src[..., F_in:]
|
|
84
|
+
x_dst_1, x_dst_2 = x_dst[..., :F_in], x_dst[..., F_in:]
|
|
85
|
+
|
|
86
|
+
out_pos1 = self.propagate(pos_edge_index, x=(x_src_1, x_dst_1))
|
|
87
|
+
out_pos2 = self.propagate(neg_edge_index, x=(x_src_2, x_dst_2))
|
|
88
|
+
out_pos = ops.concatenate([out_pos1, out_pos2], axis=-1)
|
|
89
|
+
out_pos = self.lin_pos_l(out_pos) + self.lin_pos_r(x_dst_1)
|
|
90
|
+
|
|
91
|
+
out_neg1 = self.propagate(pos_edge_index, x=(x_src_2, x_dst_2))
|
|
92
|
+
out_neg2 = self.propagate(neg_edge_index, x=(x_src_1, x_dst_1))
|
|
93
|
+
out_neg = ops.concatenate([out_neg1, out_neg2], axis=-1)
|
|
94
|
+
out_neg = self.lin_neg_l(out_neg) + self.lin_neg_r(x_dst_2)
|
|
95
|
+
|
|
96
|
+
return ops.concatenate([out_pos, out_neg], axis=-1)
|
|
97
|
+
|
|
98
|
+
def message(self, x_j):
|
|
99
|
+
return x_j
|
|
100
|
+
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
from typing import Optional, Union, List
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.aggr.base import Aggregation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class SimpleConv(MessagePassing):
|
|
9
|
+
r"""A simple, parameter-free message passing operator.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
aggr: The aggregation scheme to use (``"sum"``, ``"mean"``,
|
|
13
|
+
``"min"``, ``"max"``, ``"mul"``). (default: ``"sum"``)
|
|
14
|
+
combine_root: The way to combine root node features with the
|
|
15
|
+
aggregated output (``"sum"``, ``"cat"``, ``"self_loop"``,
|
|
16
|
+
or :obj:`None`). (default: :obj:`None`)
|
|
17
|
+
|
|
18
|
+
Example:
|
|
19
|
+
```python
|
|
20
|
+
import numpy as np
|
|
21
|
+
from k3_node.layers import SimpleConv
|
|
22
|
+
|
|
23
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
24
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
25
|
+
|
|
26
|
+
layer = SimpleConv(aggr="mean")
|
|
27
|
+
out = layer(x, edge_index)
|
|
28
|
+
print(tuple(out.shape)) # (10, 8)
|
|
29
|
+
```
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
aggr: Union[str, List[str], Aggregation, None] = "sum",
|
|
35
|
+
combine_root: Optional[str] = None,
|
|
36
|
+
**kwargs,
|
|
37
|
+
):
|
|
38
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
39
|
+
self.combine_root = combine_root
|
|
40
|
+
if combine_root is not None and combine_root not in ["sum", "cat", "self_loop"]:
|
|
41
|
+
raise ValueError(
|
|
42
|
+
f"combine_root must be 'sum', 'cat', 'self_loop', or None, got {combine_root}"
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
def build(self, input_shape):
|
|
46
|
+
self.built = True
|
|
47
|
+
|
|
48
|
+
def call(self, x, edge_index=None, edge_weight=None, size=None, **kwargs):
|
|
49
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
50
|
+
x, edge_index = x[0], x[1]
|
|
51
|
+
|
|
52
|
+
if not isinstance(x, (tuple, list)):
|
|
53
|
+
x_src, x_dst = x, x
|
|
54
|
+
else:
|
|
55
|
+
x_src, x_dst = x[0], x[1]
|
|
56
|
+
|
|
57
|
+
if self.combine_root == "self_loop":
|
|
58
|
+
from k3_node.layers.conv.utils import add_self_loops
|
|
59
|
+
num_nodes = ops.shape(x_src)[0]
|
|
60
|
+
edge_index, edge_weight = add_self_loops(edge_index, edge_weight, num_nodes=num_nodes)
|
|
61
|
+
|
|
62
|
+
out = self.propagate(edge_index, x=(x_src, x_dst), edge_weight=edge_weight, size=size)
|
|
63
|
+
|
|
64
|
+
if self.combine_root is not None and x_dst is not None and self.combine_root != "self_loop":
|
|
65
|
+
if self.combine_root == "sum":
|
|
66
|
+
out = out + x_dst
|
|
67
|
+
elif self.combine_root == "cat":
|
|
68
|
+
out = ops.concatenate([out, x_dst], axis=-1)
|
|
69
|
+
|
|
70
|
+
return out
|
|
71
|
+
|
|
72
|
+
def message(self, x_j, edge_weight=None):
|
|
73
|
+
if edge_weight is None:
|
|
74
|
+
return x_j
|
|
75
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
@@ -0,0 +1,182 @@
|
|
|
1
|
+
from typing import List, Tuple, Union
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class SplineConv(MessagePassing):
|
|
8
|
+
r"""The spline-based convolutional operator from the `"SplineCNN: Fast
|
|
9
|
+
Geometric Deep Learning with Continuous B-Spline Kernels"
|
|
10
|
+
<https://arxiv.org/abs/1711.08920>`_ paper.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
in_channels (int or tuple): Size of each input sample.
|
|
14
|
+
out_channels (int): Size of each output sample.
|
|
15
|
+
dim (int): Pseudo-coordinate dimensionality.
|
|
16
|
+
kernel_size (int or List[int]): Size of the convolving kernel.
|
|
17
|
+
is_open_spline (bool or List[bool], optional): If set to :obj:`False`,
|
|
18
|
+
uses closed B-spline basis. (default: :obj:`True`)
|
|
19
|
+
degree (int, optional): B-spline basis degree. (default: :obj:`1`)
|
|
20
|
+
aggr (str, optional): The aggregation scheme to use (:obj:`"mean"`,
|
|
21
|
+
:obj:`"add"`, :obj:`"max"`). (default: :obj:`"mean"`)
|
|
22
|
+
root_weight (bool, optional): Whether to add transformed root node
|
|
23
|
+
features. (default: :obj:`True`)
|
|
24
|
+
bias (bool, optional): Whether to learn an additive bias. (default: :obj:`True`)
|
|
25
|
+
|
|
26
|
+
Example:
|
|
27
|
+
```python
|
|
28
|
+
import numpy as np
|
|
29
|
+
from k3_node.layers import SplineConv
|
|
30
|
+
|
|
31
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
32
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
33
|
+
pseudo = np.random.rand(30, 2).astype("float32") # pseudo-coordinates in [0, 1]
|
|
34
|
+
|
|
35
|
+
layer = SplineConv(in_channels=8, out_channels=16, dim=2, kernel_size=3)
|
|
36
|
+
out = layer(x, edge_index, pseudo)
|
|
37
|
+
print(tuple(out.shape)) # (10, 16)
|
|
38
|
+
```
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(
|
|
42
|
+
self,
|
|
43
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
44
|
+
out_channels: int,
|
|
45
|
+
dim: int,
|
|
46
|
+
kernel_size: Union[int, List[int]],
|
|
47
|
+
is_open_spline: Union[bool, List[bool]] = True,
|
|
48
|
+
degree: int = 1,
|
|
49
|
+
aggr: str = "mean",
|
|
50
|
+
root_weight: bool = True,
|
|
51
|
+
bias: bool = True,
|
|
52
|
+
**kwargs,
|
|
53
|
+
):
|
|
54
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
55
|
+
|
|
56
|
+
if isinstance(in_channels, int):
|
|
57
|
+
in_channels = (in_channels, in_channels)
|
|
58
|
+
|
|
59
|
+
if isinstance(kernel_size, int):
|
|
60
|
+
kernel_size = [kernel_size] * dim
|
|
61
|
+
assert len(kernel_size) == dim
|
|
62
|
+
|
|
63
|
+
self.in_channels = in_channels
|
|
64
|
+
self.out_channels = out_channels
|
|
65
|
+
self.dim = dim
|
|
66
|
+
self.kernel_size = kernel_size
|
|
67
|
+
self.degree = degree
|
|
68
|
+
self.root_weight = root_weight
|
|
69
|
+
self.use_bias = bias
|
|
70
|
+
|
|
71
|
+
# Total number of spline basis elements
|
|
72
|
+
k_prod = 1
|
|
73
|
+
for k in kernel_size:
|
|
74
|
+
k_prod *= k
|
|
75
|
+
self.K = k_prod
|
|
76
|
+
|
|
77
|
+
# Compute strides for multidimensional indexing
|
|
78
|
+
strides = []
|
|
79
|
+
stride = 1
|
|
80
|
+
for k in reversed(kernel_size):
|
|
81
|
+
strides.insert(0, stride)
|
|
82
|
+
stride *= k
|
|
83
|
+
self.strides = strides
|
|
84
|
+
|
|
85
|
+
def build(self, input_shape=None):
|
|
86
|
+
in_dim = self.in_channels[0]
|
|
87
|
+
if in_dim == -1 and input_shape is not None:
|
|
88
|
+
if isinstance(input_shape, (list, tuple)) and isinstance(input_shape[0], (list, tuple)):
|
|
89
|
+
in_dim = input_shape[0][-1]
|
|
90
|
+
elif isinstance(input_shape, (list, tuple)):
|
|
91
|
+
in_dim = input_shape[-1]
|
|
92
|
+
self.in_channels = (in_dim, self.in_channels[1] if self.in_channels[1] != -1 else in_dim)
|
|
93
|
+
|
|
94
|
+
self.weight = self.add_weight(
|
|
95
|
+
shape=(self.K, self.in_channels[0], self.out_channels),
|
|
96
|
+
initializer="glorot_uniform",
|
|
97
|
+
trainable=True,
|
|
98
|
+
name="weight",
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
if self.root_weight:
|
|
102
|
+
self.root_lin = keras.layers.Dense(self.out_channels, use_bias=False)
|
|
103
|
+
self.root_lin.build((None, self.in_channels[1]))
|
|
104
|
+
|
|
105
|
+
if self.use_bias:
|
|
106
|
+
self.bias = self.add_weight(
|
|
107
|
+
shape=(self.out_channels,),
|
|
108
|
+
initializer="zeros",
|
|
109
|
+
trainable=True,
|
|
110
|
+
name="bias",
|
|
111
|
+
)
|
|
112
|
+
|
|
113
|
+
super().build(input_shape)
|
|
114
|
+
|
|
115
|
+
def _spline_basis_1d(self, e_d, k_d):
|
|
116
|
+
# e_d: (E,) in [0, 1]
|
|
117
|
+
u = ops.clip(e_d * (k_d - 1), 0.0, float(k_d - 1))
|
|
118
|
+
i0 = ops.cast(ops.floor(u), "int32")
|
|
119
|
+
i0 = ops.clip(i0, 0, k_d - 1)
|
|
120
|
+
i1 = ops.clip(i0 + 1, 0, k_d - 1)
|
|
121
|
+
w1 = u - ops.cast(i0, u.dtype)
|
|
122
|
+
w0 = 1.0 - w1
|
|
123
|
+
return [(i0, w0), (i1, w1)]
|
|
124
|
+
|
|
125
|
+
def _compute_kernel(self, edge_attr):
|
|
126
|
+
# edge_attr: (E, D)
|
|
127
|
+
E = ops.shape(edge_attr)[0]
|
|
128
|
+
# Basis product over dimensions
|
|
129
|
+
dim_bases = [
|
|
130
|
+
self._spline_basis_1d(edge_attr[:, d], self.kernel_size[d])
|
|
131
|
+
for d in range(self.dim)
|
|
132
|
+
]
|
|
133
|
+
|
|
134
|
+
# Cartesian product of basis across dimensions
|
|
135
|
+
basis_combinations = [([], 1.0)]
|
|
136
|
+
for d in range(self.dim):
|
|
137
|
+
new_combinations = []
|
|
138
|
+
stride = self.strides[d]
|
|
139
|
+
for curr_indices, curr_weight in basis_combinations:
|
|
140
|
+
for idx, w in dim_bases[d]:
|
|
141
|
+
new_idx = curr_indices + [idx * stride]
|
|
142
|
+
new_w = curr_weight * w
|
|
143
|
+
new_combinations.append((new_idx, new_w))
|
|
144
|
+
basis_combinations = new_combinations
|
|
145
|
+
|
|
146
|
+
# Construct edge-wise weight matrix W_eff: (E, C_in, C_out)
|
|
147
|
+
W_eff = ops.zeros((E, self.in_channels[0], self.out_channels), dtype=self.weight.dtype)
|
|
148
|
+
for idx_parts, basis_w in basis_combinations:
|
|
149
|
+
total_idx = idx_parts[0]
|
|
150
|
+
for part in idx_parts[1:]:
|
|
151
|
+
total_idx = total_idx + part
|
|
152
|
+
# total_idx: (E,)
|
|
153
|
+
basis_weights = ops.take(self.weight, total_idx, axis=0) # (E, C_in, C_out)
|
|
154
|
+
basis_w_expanded = ops.expand_dims(ops.expand_dims(basis_w, axis=-1), axis=-1)
|
|
155
|
+
W_eff = W_eff + basis_w_expanded * basis_weights
|
|
156
|
+
|
|
157
|
+
return W_eff
|
|
158
|
+
|
|
159
|
+
def call(self, x, edge_index, edge_attr=None, **kwargs):
|
|
160
|
+
if isinstance(x, (list, tuple)):
|
|
161
|
+
x_src, x_dst = x[0], x[1]
|
|
162
|
+
else:
|
|
163
|
+
x_src = x_dst = x
|
|
164
|
+
|
|
165
|
+
out = self.propagate(edge_index, x=(x_src, x_dst), edge_attr=edge_attr)
|
|
166
|
+
|
|
167
|
+
if self.root_weight and x_dst is not None:
|
|
168
|
+
out = out + self.root_lin(x_dst)
|
|
169
|
+
|
|
170
|
+
if self.use_bias:
|
|
171
|
+
out = out + self.bias
|
|
172
|
+
|
|
173
|
+
return out
|
|
174
|
+
|
|
175
|
+
def message(self, x_j, edge_attr):
|
|
176
|
+
if edge_attr is None:
|
|
177
|
+
return x_j
|
|
178
|
+
W_eff = self._compute_kernel(edge_attr) # (E, C_in, C_out)
|
|
179
|
+
# x_j: (E, C_in)
|
|
180
|
+
# m = sum_c x_j * W_eff
|
|
181
|
+
return ops.einsum("ec,ecd->ed", x_j, W_eff)
|
|
182
|
+
|