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,165 @@
|
|
|
1
|
+
from typing import Callable, Optional, Tuple
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from .consecutive import consecutive_cluster
|
|
5
|
+
from .pool import as_mutable_graph, pool_batch, pool_edge, pool_pos
|
|
6
|
+
from k3_node.ops.segment import segment_sum
|
|
7
|
+
from k3_node.ops.host import to_numpy
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _avg_pool_x(cluster, x, size: Optional[int] = None):
|
|
11
|
+
cluster = ops.cast(cluster, dtype="int32")
|
|
12
|
+
if size is None:
|
|
13
|
+
size = int(to_numpy(cluster).max()) + 1 if ops.shape(cluster)[0] > 0 else 0
|
|
14
|
+
sum_x = segment_sum(x, cluster, num_segments=size)
|
|
15
|
+
ones = ops.ones_like(x)
|
|
16
|
+
count = segment_sum(ones, cluster, num_segments=size)
|
|
17
|
+
return sum_x / ops.maximum(count, 1.0)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def avg_pool_x(
|
|
21
|
+
cluster,
|
|
22
|
+
x,
|
|
23
|
+
batch,
|
|
24
|
+
batch_size: Optional[int] = None,
|
|
25
|
+
size: Optional[int] = None,
|
|
26
|
+
) -> Tuple[any, Optional[any]]:
|
|
27
|
+
r"""Average-pools node features according to the clustering defined in `cluster`.
|
|
28
|
+
|
|
29
|
+
Example:
|
|
30
|
+
```python
|
|
31
|
+
import numpy as np
|
|
32
|
+
from k3_node.layers import avg_pool_x
|
|
33
|
+
|
|
34
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
35
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
36
|
+
cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
|
|
37
|
+
|
|
38
|
+
x_pool, batch_pool = avg_pool_x(cluster, x, batch)
|
|
39
|
+
print(tuple(x_pool.shape)) # (5, 8)
|
|
40
|
+
```
|
|
41
|
+
"""
|
|
42
|
+
if size is not None:
|
|
43
|
+
if batch_size is None:
|
|
44
|
+
batch_size = int(to_numpy(batch).max()) + 1
|
|
45
|
+
return _avg_pool_x(cluster, x, batch_size * size), None
|
|
46
|
+
|
|
47
|
+
cluster, perm = consecutive_cluster(cluster)
|
|
48
|
+
x = _avg_pool_x(cluster, x)
|
|
49
|
+
batch = pool_batch(perm, batch)
|
|
50
|
+
return x, batch
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def avg_pool(
|
|
54
|
+
cluster,
|
|
55
|
+
data,
|
|
56
|
+
transform: Optional[Callable] = None,
|
|
57
|
+
edge_index: Optional[any] = None,
|
|
58
|
+
edge_attr: Optional[any] = None,
|
|
59
|
+
batch: Optional[any] = None,
|
|
60
|
+
pos: Optional[any] = None,
|
|
61
|
+
):
|
|
62
|
+
r"""Pools and coarsens a graph given by `data` according to `cluster` using averaging.
|
|
63
|
+
|
|
64
|
+
Example:
|
|
65
|
+
```python
|
|
66
|
+
import numpy as np
|
|
67
|
+
from k3_node.layers import avg_pool
|
|
68
|
+
|
|
69
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
70
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
71
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
72
|
+
cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
|
|
73
|
+
|
|
74
|
+
x_pool, edge_index_pool, batch_pool = avg_pool(cluster, x, edge_index, batch=batch)
|
|
75
|
+
print(tuple(x_pool.shape)) # (5, 8)
|
|
76
|
+
```
|
|
77
|
+
"""
|
|
78
|
+
cluster, perm = consecutive_cluster(cluster)
|
|
79
|
+
|
|
80
|
+
if hasattr(data, "x"):
|
|
81
|
+
data = as_mutable_graph(data)
|
|
82
|
+
x = getattr(data, "x", None)
|
|
83
|
+
if x is not None:
|
|
84
|
+
data.x = _avg_pool_x(cluster, x)
|
|
85
|
+
|
|
86
|
+
edge_index = getattr(data, "edge_index", None)
|
|
87
|
+
edge_attr = getattr(data, "edge_attr", None)
|
|
88
|
+
if edge_index is not None:
|
|
89
|
+
data.edge_index, data.edge_attr = pool_edge(cluster, edge_index, edge_attr, reduce="mean")
|
|
90
|
+
|
|
91
|
+
batch = getattr(data, "batch", None)
|
|
92
|
+
if batch is not None:
|
|
93
|
+
data.batch = pool_batch(perm, batch)
|
|
94
|
+
|
|
95
|
+
pos = getattr(data, "pos", None)
|
|
96
|
+
if pos is not None:
|
|
97
|
+
data.pos = pool_pos(cluster, pos)
|
|
98
|
+
|
|
99
|
+
if transform is not None:
|
|
100
|
+
data = transform(data)
|
|
101
|
+
|
|
102
|
+
return data
|
|
103
|
+
|
|
104
|
+
# Raw tensor mode
|
|
105
|
+
pooled_x = _avg_pool_x(cluster, data)
|
|
106
|
+
pooled_edge_index, pooled_edge_attr = (None, None)
|
|
107
|
+
if edge_index is not None:
|
|
108
|
+
pooled_edge_index, pooled_edge_attr = pool_edge(cluster, edge_index, edge_attr, reduce="mean")
|
|
109
|
+
pooled_batch = pool_batch(perm, batch) if batch is not None else None
|
|
110
|
+
|
|
111
|
+
return pooled_x, pooled_edge_index, pooled_batch
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def avg_pool_neighbor_x(
|
|
115
|
+
data,
|
|
116
|
+
edge_index=None,
|
|
117
|
+
flow: str = "source_to_target",
|
|
118
|
+
):
|
|
119
|
+
r"""Average-pools neighboring node features.
|
|
120
|
+
|
|
121
|
+
Example:
|
|
122
|
+
```python
|
|
123
|
+
import numpy as np
|
|
124
|
+
from k3_node.layers import avg_pool_neighbor_x
|
|
125
|
+
|
|
126
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
127
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
128
|
+
|
|
129
|
+
out = avg_pool_neighbor_x(x, edge_index=edge_index) # pool each node with its neighbors
|
|
130
|
+
print(tuple(out.shape)) # (10, 8)
|
|
131
|
+
```
|
|
132
|
+
"""
|
|
133
|
+
if hasattr(data, "x"):
|
|
134
|
+
x = data.x
|
|
135
|
+
edge_index = data.edge_index
|
|
136
|
+
is_data_obj = True
|
|
137
|
+
else:
|
|
138
|
+
x = data
|
|
139
|
+
is_data_obj = False
|
|
140
|
+
if edge_index is None:
|
|
141
|
+
raise ValueError("edge_index must be provided if data is a tensor.")
|
|
142
|
+
|
|
143
|
+
num_nodes = getattr(data, "num_nodes", None) if is_data_obj else ops.shape(x)[0]
|
|
144
|
+
if num_nodes is None:
|
|
145
|
+
num_nodes = ops.shape(x)[0]
|
|
146
|
+
|
|
147
|
+
# Add self-loops
|
|
148
|
+
loop_idx = ops.arange(num_nodes, dtype=edge_index.dtype)
|
|
149
|
+
loop_edge = ops.stack([loop_idx, loop_idx], axis=0)
|
|
150
|
+
full_edge_index = ops.concatenate([edge_index, loop_edge], axis=1)
|
|
151
|
+
|
|
152
|
+
row = full_edge_index[0]
|
|
153
|
+
col = full_edge_index[1]
|
|
154
|
+
row, col = (row, col) if flow == "source_to_target" else (col, row)
|
|
155
|
+
|
|
156
|
+
col = ops.cast(col, dtype="int32")
|
|
157
|
+
x_src = ops.take(x, row, axis=0)
|
|
158
|
+
sum_x = segment_sum(x_src, col, num_segments=num_nodes)
|
|
159
|
+
ones = ops.ones_like(x_src)
|
|
160
|
+
count = segment_sum(ones, col, num_segments=num_nodes)
|
|
161
|
+
out_x = sum_x / ops.maximum(count, 1.0)
|
|
162
|
+
if is_data_obj:
|
|
163
|
+
data.x = out_x
|
|
164
|
+
return data
|
|
165
|
+
return out_x
|
|
@@ -0,0 +1,168 @@
|
|
|
1
|
+
from typing import NamedTuple, Optional, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
from k3_node.ops.segment import segment_sum
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class UnpoolInfo(NamedTuple):
|
|
8
|
+
edge_index: any
|
|
9
|
+
cluster: any
|
|
10
|
+
batch: any
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ClusterPooling(layers.Layer):
|
|
14
|
+
r"""The cluster pooling operator from the `"Edge-Based Graph Component
|
|
15
|
+
Pooling" <https://arxiv.org/abs/2409.11856>`_ paper.
|
|
16
|
+
|
|
17
|
+
Example:
|
|
18
|
+
```python
|
|
19
|
+
import numpy as np
|
|
20
|
+
from k3_node.layers import ClusterPooling
|
|
21
|
+
|
|
22
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
23
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
24
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
25
|
+
|
|
26
|
+
layer = ClusterPooling(in_channels=8)
|
|
27
|
+
x_pool, edge_index_pool, batch_pool, unpool_info = layer(x, edge_index, batch)
|
|
28
|
+
# The number of clusters depends on the learned edge scores
|
|
29
|
+
print(x_pool.shape[0] <= 10, x_pool.shape[1]) # True 8: fewer nodes, same features
|
|
30
|
+
```
|
|
31
|
+
"""
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
in_channels: int,
|
|
35
|
+
edge_score_method: str = "tanh",
|
|
36
|
+
dropout: float = 0.0,
|
|
37
|
+
threshold: Optional[float] = None,
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
super().__init__(**kwargs)
|
|
41
|
+
assert edge_score_method in ["tanh", "sigmoid", "log_softmax"]
|
|
42
|
+
|
|
43
|
+
if threshold is None:
|
|
44
|
+
threshold = 0.5 if edge_score_method == "sigmoid" else 0.0
|
|
45
|
+
|
|
46
|
+
self.in_channels = in_channels
|
|
47
|
+
self.edge_score_method = edge_score_method
|
|
48
|
+
self.dropout_rate = dropout
|
|
49
|
+
self.drop = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
50
|
+
self.threshold = threshold
|
|
51
|
+
|
|
52
|
+
self.lin = layers.Dense(1, use_bias=True, name="lin")
|
|
53
|
+
|
|
54
|
+
def build(self, input_shape=None):
|
|
55
|
+
self.lin.build((None, 2 * self.in_channels))
|
|
56
|
+
super().build(input_shape)
|
|
57
|
+
|
|
58
|
+
def reset_parameters(self):
|
|
59
|
+
self.lin.reset_parameters()
|
|
60
|
+
|
|
61
|
+
def call(
|
|
62
|
+
self,
|
|
63
|
+
x,
|
|
64
|
+
edge_index,
|
|
65
|
+
batch,
|
|
66
|
+
training: bool = False,
|
|
67
|
+
) -> Tuple[any, any, any, UnpoolInfo]:
|
|
68
|
+
r"""Forward pass."""
|
|
69
|
+
from k3_node.layers.conv.utils import eager_only_placeholder
|
|
70
|
+
if eager_only_placeholder("ClusterPooling", x, edge_index):
|
|
71
|
+
num_nodes = ops.shape(x)[0]
|
|
72
|
+
unpool_info = UnpoolInfo(edge_index, ops.arange(num_nodes, dtype="int32"), batch)
|
|
73
|
+
return x, edge_index, batch, unpool_info
|
|
74
|
+
|
|
75
|
+
edge_index_np = ops.convert_to_numpy(edge_index).astype(np.int64)
|
|
76
|
+
mask = edge_index_np[0] != edge_index_np[1]
|
|
77
|
+
edge_index_filtered = edge_index_np[:, mask]
|
|
78
|
+
|
|
79
|
+
row = ops.convert_to_tensor(edge_index_filtered[0], dtype="int32")
|
|
80
|
+
col = ops.convert_to_tensor(edge_index_filtered[1], dtype="int32")
|
|
81
|
+
|
|
82
|
+
edge_attr = ops.concatenate([ops.take(x, row, axis=0), ops.take(x, col, axis=0)], axis=-1)
|
|
83
|
+
edge_score = ops.reshape(self.lin(edge_attr), (-1,))
|
|
84
|
+
if self.drop is not None:
|
|
85
|
+
edge_score = self.drop(edge_score, training=training)
|
|
86
|
+
|
|
87
|
+
if self.edge_score_method == "tanh":
|
|
88
|
+
edge_score = ops.tanh(edge_score)
|
|
89
|
+
elif self.edge_score_method == "sigmoid":
|
|
90
|
+
edge_score = ops.sigmoid(edge_score)
|
|
91
|
+
else:
|
|
92
|
+
edge_score = ops.log_softmax(edge_score, axis=0)
|
|
93
|
+
|
|
94
|
+
edge_index_tensor = ops.convert_to_tensor(edge_index_filtered, dtype=edge_index.dtype)
|
|
95
|
+
return self._merge_edges(x, edge_index_tensor, batch, edge_score)
|
|
96
|
+
|
|
97
|
+
def _merge_edges(
|
|
98
|
+
self,
|
|
99
|
+
x,
|
|
100
|
+
edge_index,
|
|
101
|
+
batch,
|
|
102
|
+
edge_score,
|
|
103
|
+
) -> Tuple[any, any, any, UnpoolInfo]:
|
|
104
|
+
from scipy.sparse import coo_matrix
|
|
105
|
+
from scipy.sparse.csgraph import connected_components
|
|
106
|
+
|
|
107
|
+
from k3_node.layers.conv.utils import host_callback
|
|
108
|
+
|
|
109
|
+
num_nodes = int(ops.shape(x)[0])
|
|
110
|
+
num_edges = int(ops.shape(edge_index)[1])
|
|
111
|
+
threshold = self.threshold
|
|
112
|
+
|
|
113
|
+
def contract(edge_index_np, edge_score_np, batch_np):
|
|
114
|
+
# Clusters are the weakly connected components of the edges scoring above the threshold.
|
|
115
|
+
edge_index_np = edge_index_np.astype(np.int64)
|
|
116
|
+
edge_contract = edge_index_np[:, edge_score_np > threshold]
|
|
117
|
+
if edge_contract.shape[1] > 0:
|
|
118
|
+
adj = coo_matrix(
|
|
119
|
+
(np.ones(edge_contract.shape[1]), (edge_contract[0], edge_contract[1])),
|
|
120
|
+
shape=(num_nodes, num_nodes),
|
|
121
|
+
)
|
|
122
|
+
_, cluster_np = connected_components(adj, directed=True, connection="weak")
|
|
123
|
+
else:
|
|
124
|
+
cluster_np = np.arange(num_nodes)
|
|
125
|
+
num_clusters = int(np.max(cluster_np)) + 1 if num_nodes > 0 else 0
|
|
126
|
+
|
|
127
|
+
# Nodes without any contracted edge keep their own features (unit diagonal score).
|
|
128
|
+
single = np.ones(num_nodes, dtype=bool)
|
|
129
|
+
single[edge_contract[0]] = False
|
|
130
|
+
single[edge_contract[1]] = False
|
|
131
|
+
|
|
132
|
+
# Coarsened edges between distinct clusters, in (row, col) order.
|
|
133
|
+
pairs = cluster_np[edge_index_np]
|
|
134
|
+
pairs = np.unique(pairs[:, pairs[0] != pairs[1]], axis=1)
|
|
135
|
+
edges_pad = np.zeros((2, num_edges), dtype=np.int64)
|
|
136
|
+
edges_pad[:, : pairs.shape[1]] = pairs
|
|
137
|
+
|
|
138
|
+
batch_pad = np.zeros(num_nodes, dtype=np.int64)
|
|
139
|
+
batch_pad[cluster_np] = batch_np
|
|
140
|
+
return cluster_np, num_clusters, single, edges_pad, pairs.shape[1], batch_pad
|
|
141
|
+
|
|
142
|
+
cluster, num_clusters, single, edges_pad, num_new_edges, batch_pad = host_callback(
|
|
143
|
+
contract,
|
|
144
|
+
[((num_nodes,), "int32"), ((), "int32"), ((num_nodes,), "float32"),
|
|
145
|
+
((2, num_edges), "int32"), ((), "int32"), ((num_nodes,), "int32")],
|
|
146
|
+
edge_index, edge_score, batch,
|
|
147
|
+
)
|
|
148
|
+
num_clusters, num_new_edges = int(num_clusters), int(num_new_edges)
|
|
149
|
+
|
|
150
|
+
# x_out = (S @ C)^T @ x, computed sparsely: every edge (row -> col) adds score * x[row] to
|
|
151
|
+
# cluster(col), and every unmatched node adds its own features to its cluster. The score
|
|
152
|
+
# enters as a tensor so gradients reach the scoring layer.
|
|
153
|
+
row = ops.cast(edge_index[0], "int32")
|
|
154
|
+
col = ops.cast(edge_index[1], "int32")
|
|
155
|
+
msgs = ops.expand_dims(ops.cast(edge_score, x.dtype), -1) * ops.take(x, row, axis=0)
|
|
156
|
+
x_out = segment_sum(msgs, ops.take(cluster, col, axis=0), num_segments=num_clusters)
|
|
157
|
+
x_out = x_out + segment_sum(
|
|
158
|
+
x * ops.expand_dims(ops.cast(single, x.dtype), -1), cluster, num_segments=num_clusters
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
edge_index_out = ops.cast(edges_pad[:, :num_new_edges], edge_index.dtype)
|
|
162
|
+
batch_out = ops.cast(batch_pad[:num_clusters], batch.dtype)
|
|
163
|
+
|
|
164
|
+
unpool_info = UnpoolInfo(edge_index, cluster, batch)
|
|
165
|
+
return x_out, edge_index_out, batch_out, unpool_info
|
|
166
|
+
|
|
167
|
+
def __repr__(self) -> str:
|
|
168
|
+
return f"{self.__class__.__name__}({self.in_channels})"
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import Optional
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from ..select.base import SelectOutput
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass
|
|
9
|
+
class ConnectOutput:
|
|
10
|
+
r"""The output of the :class:`Connect` method, which holds the coarsened
|
|
11
|
+
graph structure, and optional pooled edge features and batch vectors.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
edge_index: The edge indices of the coarsened graph.
|
|
15
|
+
edge_attr: The pooled edge features of the coarsened graph. (default: None)
|
|
16
|
+
batch: The pooled batch vector of the coarsened graph. (default: None)
|
|
17
|
+
"""
|
|
18
|
+
edge_index: any
|
|
19
|
+
edge_attr: Optional[any] = None
|
|
20
|
+
batch: Optional[any] = None
|
|
21
|
+
|
|
22
|
+
def __post_init__(self):
|
|
23
|
+
shape_edge = getattr(self.edge_index, "shape", None)
|
|
24
|
+
if shape_edge is not None:
|
|
25
|
+
if len(shape_edge) != 2:
|
|
26
|
+
raise ValueError(
|
|
27
|
+
f"Expected 'edge_index' to be two-dimensional "
|
|
28
|
+
f"(got {len(shape_edge)} dimensions)"
|
|
29
|
+
)
|
|
30
|
+
if shape_edge[0] is not None and shape_edge[0] != 2:
|
|
31
|
+
raise ValueError(
|
|
32
|
+
f"Expected 'edge_index' to have size '2' in the first dimension "
|
|
33
|
+
f"(got '{shape_edge[0]}')"
|
|
34
|
+
)
|
|
35
|
+
if self.edge_attr is not None:
|
|
36
|
+
shape_attr = getattr(self.edge_attr, "shape", None)
|
|
37
|
+
if (
|
|
38
|
+
shape_edge is not None
|
|
39
|
+
and shape_attr is not None
|
|
40
|
+
and len(shape_edge) == 2
|
|
41
|
+
and len(shape_attr) >= 1
|
|
42
|
+
and shape_edge[1] is not None
|
|
43
|
+
and shape_attr[0] is not None
|
|
44
|
+
and shape_attr[0] != shape_edge[1]
|
|
45
|
+
):
|
|
46
|
+
raise ValueError(
|
|
47
|
+
f"Expected 'edge_index' and 'edge_attr' to hold the same number "
|
|
48
|
+
f"of edges (got {shape_edge[1]} and {shape_attr[0]} edges)"
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
import keras
|
|
53
|
+
|
|
54
|
+
if keras.config.backend() == "jax":
|
|
55
|
+
try:
|
|
56
|
+
import jax
|
|
57
|
+
from jax.tree_util import register_pytree_node
|
|
58
|
+
|
|
59
|
+
register_pytree_node(
|
|
60
|
+
ConnectOutput,
|
|
61
|
+
lambda c: ((c.edge_index, c.edge_attr, c.batch), ()),
|
|
62
|
+
lambda aux, children: ConnectOutput(children[0], children[1], children[2]),
|
|
63
|
+
)
|
|
64
|
+
except Exception:
|
|
65
|
+
pass
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
class Connect(layers.Layer):
|
|
70
|
+
r"""An abstract base class for implementing custom edge connection
|
|
71
|
+
operators as described in the `"Understanding Pooling in Graph Neural
|
|
72
|
+
Networks" <https://arxiv.org/abs/1905.05178>`_ paper.
|
|
73
|
+
"""
|
|
74
|
+
def reset_parameters(self):
|
|
75
|
+
r"""Resets all learnable parameters of the module."""
|
|
76
|
+
pass
|
|
77
|
+
|
|
78
|
+
def __call__(self, *args, **kwargs):
|
|
79
|
+
if len(args) > 0 and isinstance(args[0], SelectOutput):
|
|
80
|
+
return self.call(*args, **kwargs)
|
|
81
|
+
return super().__call__(*args, **kwargs)
|
|
82
|
+
|
|
83
|
+
def call(
|
|
84
|
+
self,
|
|
85
|
+
select_output: SelectOutput,
|
|
86
|
+
edge_index,
|
|
87
|
+
edge_attr: Optional[any] = None,
|
|
88
|
+
batch: Optional[any] = None,
|
|
89
|
+
) -> ConnectOutput:
|
|
90
|
+
raise NotImplementedError
|
|
91
|
+
|
|
92
|
+
@staticmethod
|
|
93
|
+
def get_pooled_batch(
|
|
94
|
+
select_output: SelectOutput,
|
|
95
|
+
batch: Optional[any],
|
|
96
|
+
) -> Optional[any]:
|
|
97
|
+
r"""Returns the batch vector of the coarsened graph."""
|
|
98
|
+
if batch is None:
|
|
99
|
+
return None
|
|
100
|
+
return ops.take(batch, select_output.node_index, axis=0)
|
|
101
|
+
|
|
102
|
+
def __repr__(self) -> str:
|
|
103
|
+
return f'{self.__class__.__name__}()'
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
from typing import Optional, Tuple
|
|
2
|
+
from keras import ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from .base import Connect, ConnectOutput
|
|
6
|
+
from ..select.base import SelectOutput
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
10
|
+
from k3_node.ops.creation import full
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def filter_adj(
|
|
14
|
+
edge_index,
|
|
15
|
+
edge_attr: Optional[any] = None,
|
|
16
|
+
node_index=None,
|
|
17
|
+
cluster_index: Optional[any] = None,
|
|
18
|
+
num_nodes: Optional[int] = None,
|
|
19
|
+
) -> Tuple[any, Optional[any]]:
|
|
20
|
+
r"""Filters out edges if their incident nodes are not in any cluster.
|
|
21
|
+
|
|
22
|
+
Example:
|
|
23
|
+
```python
|
|
24
|
+
import numpy as np
|
|
25
|
+
from k3_node.layers import filter_adj
|
|
26
|
+
|
|
27
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
28
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
29
|
+
|
|
30
|
+
kept_nodes = np.array([0, 2, 4, 6, 8])
|
|
31
|
+
edge_index_new, _ = filter_adj(edge_index, node_index=kept_nodes) # edges between kept nodes, relabeled
|
|
32
|
+
print(edge_index_new.shape[0]) # 2
|
|
33
|
+
```
|
|
34
|
+
"""
|
|
35
|
+
if is_tracing(edge_index) or (node_index is not None and is_tracing(node_index)):
|
|
36
|
+
return edge_index, edge_attr
|
|
37
|
+
|
|
38
|
+
if node_index is None:
|
|
39
|
+
return edge_index, edge_attr
|
|
40
|
+
|
|
41
|
+
edge_index = ops.cast(edge_index, "int32")
|
|
42
|
+
node_index = ops.cast(node_index, "int32")
|
|
43
|
+
if cluster_index is None:
|
|
44
|
+
cluster_index = ops.arange(ops.shape(node_index)[0], dtype="int32")
|
|
45
|
+
else:
|
|
46
|
+
cluster_index = ops.cast(cluster_index, "int32")
|
|
47
|
+
|
|
48
|
+
if num_nodes is None:
|
|
49
|
+
num_nodes = ops.max(node_index) + 1 if ops.shape(node_index)[0] > 0 else 0
|
|
50
|
+
if ops.shape(edge_index)[1] > 0:
|
|
51
|
+
num_nodes = ops.maximum(num_nodes, ops.max(edge_index) + 1)
|
|
52
|
+
try:
|
|
53
|
+
num_nodes = int(num_nodes)
|
|
54
|
+
except (TypeError, ValueError):
|
|
55
|
+
pass
|
|
56
|
+
|
|
57
|
+
mapping = full((num_nodes,), -1, dtype="int32")
|
|
58
|
+
mapping = ops.scatter_update(mapping, ops.expand_dims(node_index, -1), cluster_index)
|
|
59
|
+
|
|
60
|
+
row = ops.take(mapping, edge_index[0], axis=0)
|
|
61
|
+
col = ops.take(mapping, edge_index[1], axis=0)
|
|
62
|
+
mask = (row >= 0) & (col >= 0)
|
|
63
|
+
valid_idx = ops.where(mask)
|
|
64
|
+
if isinstance(valid_idx, (tuple, list)):
|
|
65
|
+
valid_idx = valid_idx[0]
|
|
66
|
+
valid_idx = ops.reshape(valid_idx, (-1,))
|
|
67
|
+
|
|
68
|
+
new_edge_index = ops.stack(
|
|
69
|
+
[ops.take(row, valid_idx, axis=0), ops.take(col, valid_idx, axis=0)], axis=0
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
new_edge_attr = None
|
|
73
|
+
if edge_attr is not None:
|
|
74
|
+
new_edge_attr = ops.take(edge_attr, valid_idx, axis=0)
|
|
75
|
+
|
|
76
|
+
return new_edge_index, new_edge_attr
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
class FilterEdges(Connect):
|
|
80
|
+
r"""Filters out edges if their incident nodes are not in any cluster.
|
|
81
|
+
|
|
82
|
+
Example:
|
|
83
|
+
```python
|
|
84
|
+
import numpy as np
|
|
85
|
+
from k3_node.layers import FilterEdges, SelectTopK
|
|
86
|
+
|
|
87
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
88
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
89
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
90
|
+
|
|
91
|
+
select_output = SelectTopK(in_channels=8, ratio=0.5)(x, batch)
|
|
92
|
+
connect = FilterEdges()
|
|
93
|
+
out = connect(select_output, edge_index, batch=batch) # keep edges between kept nodes
|
|
94
|
+
print(out.edge_index.shape[0]) # 2
|
|
95
|
+
```
|
|
96
|
+
"""
|
|
97
|
+
def call(
|
|
98
|
+
self,
|
|
99
|
+
select_output: SelectOutput,
|
|
100
|
+
edge_index,
|
|
101
|
+
edge_attr: Optional[any] = None,
|
|
102
|
+
batch: Optional[any] = None,
|
|
103
|
+
) -> ConnectOutput:
|
|
104
|
+
new_edge_index, new_edge_attr = filter_adj(
|
|
105
|
+
edge_index,
|
|
106
|
+
edge_attr,
|
|
107
|
+
select_output.node_index,
|
|
108
|
+
select_output.cluster_index,
|
|
109
|
+
num_nodes=select_output.num_nodes,
|
|
110
|
+
)
|
|
111
|
+
new_batch = self.get_pooled_batch(select_output, batch)
|
|
112
|
+
return ConnectOutput(new_edge_index, new_edge_attr, new_batch)
|
|
113
|
+
|
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
from typing import Tuple
|
|
2
|
+
from keras import ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
from k3_node.ops.host import to_numpy
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def consecutive_cluster(src) -> Tuple[any, any]:
|
|
8
|
+
r"""Maps elements in `src` to consecutive integers starting from 0,
|
|
9
|
+
returning the mapped indices and a permutation array of representative indices.
|
|
10
|
+
|
|
11
|
+
Example:
|
|
12
|
+
```python
|
|
13
|
+
import numpy as np
|
|
14
|
+
from k3_node.layers import consecutive_cluster
|
|
15
|
+
|
|
16
|
+
cluster = np.array([4, 4, 9, 2, 9])
|
|
17
|
+
new_cluster, perm = consecutive_cluster(cluster) # relabel cluster ids to 0..num_clusters-1
|
|
18
|
+
print(tuple(new_cluster.shape), tuple(perm.shape)) # (5,) (3,)
|
|
19
|
+
```
|
|
20
|
+
"""
|
|
21
|
+
src_np = to_numpy(src)
|
|
22
|
+
unique, inv = np.unique(src_np, return_inverse=True)
|
|
23
|
+
perm = np.empty(len(unique), dtype=inv.dtype)
|
|
24
|
+
arange = np.arange(len(inv), dtype=inv.dtype)
|
|
25
|
+
perm[inv] = arange
|
|
26
|
+
|
|
27
|
+
inv_tensor = ops.convert_to_tensor(inv, dtype=src.dtype)
|
|
28
|
+
perm_tensor = ops.convert_to_tensor(perm, dtype=src.dtype)
|
|
29
|
+
return inv_tensor, perm_tensor
|
|
30
|
+
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
from typing import Tuple, Union
|
|
2
|
+
from keras import ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
from k3_node.ops.host import _in_shape_inference, to_numpy
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def decimation_indices(
|
|
8
|
+
ptr,
|
|
9
|
+
decimation_factor: Union[int, float],
|
|
10
|
+
) -> Tuple[any, any]:
|
|
11
|
+
r"""Gets indices which downsample each point cloud by a decimation factor.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import decimation_indices
|
|
17
|
+
|
|
18
|
+
ptr = np.array([0, 4, 10]) # two graphs with 4 and 6 nodes
|
|
19
|
+
index, new_ptr = decimation_indices(ptr, decimation_factor=2) # keep every 2nd node per graph
|
|
20
|
+
print(tuple(index.shape), tuple(new_ptr.shape)) # (5,) (3,)
|
|
21
|
+
```
|
|
22
|
+
"""
|
|
23
|
+
if decimation_factor < 1:
|
|
24
|
+
raise ValueError(
|
|
25
|
+
f"The argument `decimation_factor` should be higher than (or "
|
|
26
|
+
f"equal to) 1 for downsampling. (got {decimation_factor})"
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
ptr_np = to_numpy(ptr)
|
|
30
|
+
batch_size = len(ptr_np) - 1
|
|
31
|
+
count = ptr_np[1:] - ptr_np[:-1]
|
|
32
|
+
if _in_shape_inference(): # `ptr` is placeholder zeros: keep one (valid) node per graph
|
|
33
|
+
count = np.maximum(count, 1)
|
|
34
|
+
decim_count = np.maximum(count // int(decimation_factor), 1)
|
|
35
|
+
|
|
36
|
+
decim_indices_list = []
|
|
37
|
+
for i in range(batch_size):
|
|
38
|
+
perm = np.random.permutation(count[i])[:decim_count[i]]
|
|
39
|
+
decim_indices_list.append(ptr_np[i] + perm)
|
|
40
|
+
|
|
41
|
+
decim_indices = np.concatenate(decim_indices_list, axis=0)
|
|
42
|
+
decim_ptr = np.concatenate([[0], np.cumsum(decim_count)])
|
|
43
|
+
|
|
44
|
+
return (
|
|
45
|
+
ops.convert_to_tensor(decim_indices, dtype=ptr.dtype),
|
|
46
|
+
ops.convert_to_tensor(decim_ptr, dtype=ptr.dtype),
|
|
47
|
+
)
|
|
48
|
+
|