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,119 @@
|
|
|
1
|
+
from typing import Optional, Tuple
|
|
2
|
+
from keras import ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
5
|
+
from k3_node.ops.host import to_numpy
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def pool_edge(
|
|
9
|
+
cluster,
|
|
10
|
+
edge_index,
|
|
11
|
+
edge_attr: Optional[any] = None,
|
|
12
|
+
reduce: str = "sum",
|
|
13
|
+
) -> Tuple[any, Optional[any]]:
|
|
14
|
+
r"""Pools edge indices and attributes based on cluster assignments.
|
|
15
|
+
|
|
16
|
+
Example:
|
|
17
|
+
```python
|
|
18
|
+
import numpy as np
|
|
19
|
+
from k3_node.layers import pool_edge
|
|
20
|
+
|
|
21
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
22
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
23
|
+
cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
|
|
24
|
+
|
|
25
|
+
edge_attr = np.random.rand(30, 3).astype("float32")
|
|
26
|
+
edge_index_pool, edge_attr_pool = pool_edge(cluster, edge_index, edge_attr) # coarsened, deduplicated edges
|
|
27
|
+
print(edge_index_pool.shape[0], edge_attr_pool.shape[1]) # 2 3
|
|
28
|
+
```
|
|
29
|
+
"""
|
|
30
|
+
cluster_np = to_numpy(cluster)
|
|
31
|
+
edge_index_np = to_numpy(edge_index)
|
|
32
|
+
|
|
33
|
+
row = cluster_np[edge_index_np[0]]
|
|
34
|
+
col = cluster_np[edge_index_np[1]]
|
|
35
|
+
|
|
36
|
+
# Remove self-loops
|
|
37
|
+
non_loop = row != col
|
|
38
|
+
row = row[non_loop]
|
|
39
|
+
col = col[non_loop]
|
|
40
|
+
|
|
41
|
+
if len(row) == 0:
|
|
42
|
+
empty_ei = ops.zeros((2, 0), dtype=edge_index.dtype)
|
|
43
|
+
empty_ea = None if edge_attr is None else ops.zeros((0, *ops.shape(edge_attr)[1:]), dtype=edge_attr.dtype)
|
|
44
|
+
return empty_ei, empty_ea
|
|
45
|
+
|
|
46
|
+
edges = np.stack([row, col], axis=0)
|
|
47
|
+
# Coalesce duplicate edges
|
|
48
|
+
unique_edges, inv = np.unique(edges, axis=1, return_inverse=True)
|
|
49
|
+
|
|
50
|
+
out_edge_index = ops.convert_to_tensor(unique_edges, dtype=edge_index.dtype)
|
|
51
|
+
|
|
52
|
+
out_edge_attr = None
|
|
53
|
+
if edge_attr is not None:
|
|
54
|
+
ea_np = to_numpy(edge_attr)[non_loop]
|
|
55
|
+
num_unique = unique_edges.shape[1]
|
|
56
|
+
ea_tensor = ops.convert_to_tensor(ea_np, dtype=edge_attr.dtype)
|
|
57
|
+
inv_tensor = ops.convert_to_tensor(inv, dtype="int32")
|
|
58
|
+
if reduce == "sum":
|
|
59
|
+
out_edge_attr = segment_sum(ea_tensor, inv_tensor, num_segments=num_unique)
|
|
60
|
+
elif reduce == "mean":
|
|
61
|
+
sum_ea = segment_sum(ea_tensor, inv_tensor, num_segments=num_unique)
|
|
62
|
+
count = segment_sum(ops.ones_like(ea_tensor), inv_tensor, num_segments=num_unique)
|
|
63
|
+
out_edge_attr = sum_ea / ops.maximum(count, 1.0)
|
|
64
|
+
elif reduce == "max":
|
|
65
|
+
out_edge_attr = segment_max(ea_tensor, inv_tensor, num_segments=num_unique)
|
|
66
|
+
|
|
67
|
+
return out_edge_index, out_edge_attr
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def pool_batch(perm, batch):
|
|
71
|
+
r"""Pools batch vector given representative indices `perm`.
|
|
72
|
+
|
|
73
|
+
Example:
|
|
74
|
+
```python
|
|
75
|
+
import numpy as np
|
|
76
|
+
from k3_node.layers import pool_batch
|
|
77
|
+
|
|
78
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
79
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
80
|
+
|
|
81
|
+
perm = np.array([0, 3, 6, 9]) # nodes kept after pooling
|
|
82
|
+
print(tuple(pool_batch(perm, batch).shape)) # (4,): graph id of each kept node
|
|
83
|
+
```
|
|
84
|
+
"""
|
|
85
|
+
return ops.take(batch, perm, axis=0)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def pool_pos(cluster, pos):
|
|
89
|
+
r"""Pools node positions by computing average coordinates within each cluster.
|
|
90
|
+
|
|
91
|
+
Example:
|
|
92
|
+
```python
|
|
93
|
+
import numpy as np
|
|
94
|
+
from k3_node.layers import pool_pos
|
|
95
|
+
|
|
96
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
97
|
+
pos = np.random.rand(10, 3).astype("float32") # 3D positions
|
|
98
|
+
cluster = np.repeat(np.arange(5), 2) # merge nodes pairwise into 5 clusters
|
|
99
|
+
|
|
100
|
+
print(tuple(pool_pos(cluster, pos).shape)) # (5, 3): mean position of every cluster
|
|
101
|
+
```
|
|
102
|
+
"""
|
|
103
|
+
cluster = ops.cast(cluster, dtype="int32")
|
|
104
|
+
num_clusters = int(to_numpy(cluster).max()) + 1 if ops.shape(cluster)[0] > 0 else 0
|
|
105
|
+
sum_pos = segment_sum(pos, cluster, num_segments=num_clusters)
|
|
106
|
+
count = segment_sum(ops.ones_like(pos), cluster, num_segments=num_clusters)
|
|
107
|
+
return sum_pos / ops.maximum(count, 1.0)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def as_mutable_graph(data):
|
|
111
|
+
"""Turns a (read-only) loader batch into a ``Data`` object that pooling can update in place.
|
|
112
|
+
|
|
113
|
+
``ptr`` is dropped because it no longer matches the nodes once they are pooled.
|
|
114
|
+
"""
|
|
115
|
+
if hasattr(data, "_asdict"):
|
|
116
|
+
from k3_node.data import Data
|
|
117
|
+
|
|
118
|
+
return Data(**{k: v for k, v in data._asdict().items() if k != "ptr" and v is not None})
|
|
119
|
+
return data
|
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
from typing import Callable, Optional, Tuple, Union
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from .connect.filter_edges import FilterEdges
|
|
5
|
+
from .select.topk import SelectTopK
|
|
6
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class GraphConv(layers.Layer):
|
|
10
|
+
r"""Basic GraphConv layer for projection scoring in pooling layers."""
|
|
11
|
+
def __init__(
|
|
12
|
+
self,
|
|
13
|
+
in_channels: int,
|
|
14
|
+
out_channels: int,
|
|
15
|
+
aggr: str = "add",
|
|
16
|
+
bias: bool = True,
|
|
17
|
+
**kwargs,
|
|
18
|
+
):
|
|
19
|
+
super().__init__(**kwargs)
|
|
20
|
+
self.in_channels = in_channels
|
|
21
|
+
self.out_channels = out_channels
|
|
22
|
+
self.aggr = aggr
|
|
23
|
+
|
|
24
|
+
self.lin_rel = layers.Dense(out_channels, use_bias=bias, name="lin_rel")
|
|
25
|
+
self.lin_root = layers.Dense(out_channels, use_bias=False, name="lin_root")
|
|
26
|
+
|
|
27
|
+
def reset_parameters(self):
|
|
28
|
+
pass
|
|
29
|
+
|
|
30
|
+
def build(self, input_shape=None):
|
|
31
|
+
if hasattr(self.lin_rel, "built") and not self.lin_rel.built:
|
|
32
|
+
self.lin_rel.build((None, self.in_channels))
|
|
33
|
+
if hasattr(self.lin_root, "built") and not self.lin_root.built:
|
|
34
|
+
self.lin_root.build((None, self.in_channels))
|
|
35
|
+
self.built = True
|
|
36
|
+
|
|
37
|
+
def call(self, x, edge_index, edge_weight: Optional[any] = None):
|
|
38
|
+
row = ops.cast(edge_index[0], dtype="int32")
|
|
39
|
+
col = ops.cast(edge_index[1], dtype="int32")
|
|
40
|
+
num_nodes = ops.shape(x)[0]
|
|
41
|
+
|
|
42
|
+
msg = ops.take(x, row, axis=0)
|
|
43
|
+
if edge_weight is not None:
|
|
44
|
+
msg = msg * ops.reshape(edge_weight, (-1, 1))
|
|
45
|
+
|
|
46
|
+
if self.aggr == "add":
|
|
47
|
+
aggr_out = segment_sum(msg, col, num_segments=num_nodes)
|
|
48
|
+
elif self.aggr == "mean":
|
|
49
|
+
sum_val = segment_sum(msg, col, num_segments=num_nodes)
|
|
50
|
+
count = segment_sum(ops.ones_like(msg), col, num_segments=num_nodes)
|
|
51
|
+
aggr_out = sum_val / ops.maximum(count, 1.0)
|
|
52
|
+
elif self.aggr == "max":
|
|
53
|
+
aggr_out = segment_max(msg, col, num_segments=num_nodes)
|
|
54
|
+
else:
|
|
55
|
+
aggr_out = segment_sum(msg, col, num_segments=num_nodes)
|
|
56
|
+
|
|
57
|
+
return self.lin_rel(aggr_out) + self.lin_root(x)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
class SAGPooling(layers.Layer):
|
|
61
|
+
r"""The self-attention pooling operator from the `"Self-Attention Graph
|
|
62
|
+
Pooling" <https://arxiv.org/abs/1904.08082>`_ and `"Understanding
|
|
63
|
+
Attention and Generalization in Graph Neural Networks"
|
|
64
|
+
<https://arxiv.org/abs/1905.02850>`_ papers.
|
|
65
|
+
|
|
66
|
+
Example:
|
|
67
|
+
```python
|
|
68
|
+
import numpy as np
|
|
69
|
+
from k3_node.layers import SAGPooling
|
|
70
|
+
|
|
71
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
72
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
73
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
74
|
+
|
|
75
|
+
layer = SAGPooling(in_channels=8, ratio=0.5) # keep half of the nodes of each graph
|
|
76
|
+
out = layer(x, edge_index, batch=batch)
|
|
77
|
+
x_pool, edge_index_pool, edge_attr_pool, batch_pool = out[0], out[1], out[2], out[3]
|
|
78
|
+
print(tuple(x_pool.shape)) # (6, 8)
|
|
79
|
+
```
|
|
80
|
+
"""
|
|
81
|
+
def __init__(
|
|
82
|
+
self,
|
|
83
|
+
in_channels: int,
|
|
84
|
+
ratio: Union[float, int] = 0.5,
|
|
85
|
+
GNN: Optional[any] = None,
|
|
86
|
+
min_score: Optional[float] = None,
|
|
87
|
+
multiplier: float = 1.0,
|
|
88
|
+
nonlinearity: Union[str, Callable] = "tanh",
|
|
89
|
+
**kwargs,
|
|
90
|
+
):
|
|
91
|
+
super().__init__()
|
|
92
|
+
|
|
93
|
+
self.in_channels = in_channels
|
|
94
|
+
self.ratio = ratio
|
|
95
|
+
self.min_score = min_score
|
|
96
|
+
self.multiplier = multiplier
|
|
97
|
+
self.nonlinearity = nonlinearity
|
|
98
|
+
|
|
99
|
+
if GNN is None:
|
|
100
|
+
self.gnn = GraphConv(in_channels, 1, **kwargs)
|
|
101
|
+
elif callable(GNN):
|
|
102
|
+
self.gnn = GNN(in_channels, 1, **kwargs)
|
|
103
|
+
else:
|
|
104
|
+
self.gnn = GNN
|
|
105
|
+
|
|
106
|
+
self.select = SelectTopK(1, ratio, min_score, nonlinearity)
|
|
107
|
+
self.connect = FilterEdges()
|
|
108
|
+
|
|
109
|
+
def reset_parameters(self):
|
|
110
|
+
r"""Resets all learnable parameters of the module."""
|
|
111
|
+
if hasattr(self.gnn, "reset_parameters"):
|
|
112
|
+
self.gnn.reset_parameters()
|
|
113
|
+
self.select.reset_parameters()
|
|
114
|
+
|
|
115
|
+
def build(self, input_shape=None):
|
|
116
|
+
if hasattr(self.gnn, "built") and not self.gnn.built:
|
|
117
|
+
self.gnn.build(input_shape)
|
|
118
|
+
if hasattr(self.select, "built") and not self.select.built:
|
|
119
|
+
self.select.build(None)
|
|
120
|
+
if hasattr(self.connect, "built") and not self.connect.built:
|
|
121
|
+
self.connect.build(None)
|
|
122
|
+
self.built = True
|
|
123
|
+
|
|
124
|
+
def call(
|
|
125
|
+
self,
|
|
126
|
+
x,
|
|
127
|
+
edge_index,
|
|
128
|
+
edge_attr: Optional[any] = None,
|
|
129
|
+
batch: Optional[any] = None,
|
|
130
|
+
attn: Optional[any] = None,
|
|
131
|
+
) -> Tuple[any, any, Optional[any], Optional[any], any, any]:
|
|
132
|
+
r"""Forward pass."""
|
|
133
|
+
num_nodes = ops.shape(x)[0]
|
|
134
|
+
if batch is None:
|
|
135
|
+
batch = ops.zeros((num_nodes,), dtype="int32")
|
|
136
|
+
|
|
137
|
+
if attn is None:
|
|
138
|
+
attn = x
|
|
139
|
+
if len(ops.shape(attn)) == 1:
|
|
140
|
+
attn = ops.expand_dims(attn, axis=-1)
|
|
141
|
+
|
|
142
|
+
attn = self.gnn(attn, edge_index)
|
|
143
|
+
|
|
144
|
+
select_out = self.select(attn, batch)
|
|
145
|
+
|
|
146
|
+
perm = select_out.node_index
|
|
147
|
+
score = select_out.weight
|
|
148
|
+
|
|
149
|
+
x_pooled = ops.take(x, perm, axis=0) * ops.expand_dims(score, axis=-1)
|
|
150
|
+
if self.multiplier != 1.0:
|
|
151
|
+
x_pooled = x_pooled * self.multiplier
|
|
152
|
+
|
|
153
|
+
connect_out = self.connect(select_out, edge_index, edge_attr, batch)
|
|
154
|
+
|
|
155
|
+
return (
|
|
156
|
+
x_pooled,
|
|
157
|
+
connect_out.edge_index,
|
|
158
|
+
connect_out.edge_attr,
|
|
159
|
+
connect_out.batch,
|
|
160
|
+
perm,
|
|
161
|
+
score,
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
def __repr__(self) -> str:
|
|
165
|
+
if self.min_score is None:
|
|
166
|
+
ratio = f"ratio={self.ratio}"
|
|
167
|
+
else:
|
|
168
|
+
ratio = f"min_score={self.min_score}"
|
|
169
|
+
gnn_name = self.gnn.__class__.__name__
|
|
170
|
+
return (
|
|
171
|
+
f"{self.__class__.__name__}({gnn_name}, {self.in_channels}, "
|
|
172
|
+
f"{ratio}, multiplier={self.multiplier})"
|
|
173
|
+
)
|
|
174
|
+
|
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from typing import Optional
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
@dataclass
|
|
7
|
+
class SelectOutput:
|
|
8
|
+
r"""The output of the :class:`Select` method, which holds an assignment
|
|
9
|
+
from selected nodes to their respective cluster(s).
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
node_index: The indices of the selected nodes.
|
|
13
|
+
num_nodes: The number of nodes.
|
|
14
|
+
cluster_index: The indices of the clusters each node in
|
|
15
|
+
:obj:`node_index` is assigned to.
|
|
16
|
+
num_clusters: The number of clusters.
|
|
17
|
+
weight (optional): A weight vector, denoting the strength
|
|
18
|
+
of the assignment of a node to its cluster. (default: :obj:`None`)
|
|
19
|
+
"""
|
|
20
|
+
node_index: any
|
|
21
|
+
num_nodes: int
|
|
22
|
+
cluster_index: any
|
|
23
|
+
num_clusters: int
|
|
24
|
+
weight: Optional[any] = None
|
|
25
|
+
|
|
26
|
+
def __post_init__(self):
|
|
27
|
+
shape_node = getattr(self.node_index, "shape", None)
|
|
28
|
+
shape_cluster = getattr(self.cluster_index, "shape", None)
|
|
29
|
+
if shape_node is not None and len(shape_node) != 1:
|
|
30
|
+
raise ValueError(
|
|
31
|
+
f"Expected 'node_index' to be one-dimensional "
|
|
32
|
+
f"(got {len(shape_node)} dimensions)"
|
|
33
|
+
)
|
|
34
|
+
if shape_cluster is not None and len(shape_cluster) != 1:
|
|
35
|
+
raise ValueError(
|
|
36
|
+
f"Expected 'cluster_index' to be one-dimensional "
|
|
37
|
+
f"(got {len(shape_cluster)} dimensions)"
|
|
38
|
+
)
|
|
39
|
+
if (
|
|
40
|
+
shape_node is not None
|
|
41
|
+
and shape_cluster is not None
|
|
42
|
+
and len(shape_node) > 0
|
|
43
|
+
and len(shape_cluster) > 0
|
|
44
|
+
and shape_node[0] is not None
|
|
45
|
+
and shape_cluster[0] is not None
|
|
46
|
+
and shape_node[0] != shape_cluster[0]
|
|
47
|
+
):
|
|
48
|
+
raise ValueError(
|
|
49
|
+
f"Expected 'node_index' and 'cluster_index' to hold the same "
|
|
50
|
+
f"number of values (got {shape_node[0]} and "
|
|
51
|
+
f"{shape_cluster[0]} values)"
|
|
52
|
+
)
|
|
53
|
+
if self.weight is not None:
|
|
54
|
+
shape_weight = getattr(self.weight, "shape", None)
|
|
55
|
+
if shape_weight is not None and len(shape_weight) != 1:
|
|
56
|
+
raise ValueError(
|
|
57
|
+
f"Expected 'weight' vector to be one-dimensional "
|
|
58
|
+
f"(got {len(shape_weight)} dimensions)"
|
|
59
|
+
)
|
|
60
|
+
if (
|
|
61
|
+
shape_weight is not None
|
|
62
|
+
and shape_node is not None
|
|
63
|
+
and len(shape_weight) > 0
|
|
64
|
+
and len(shape_node) > 0
|
|
65
|
+
and shape_weight[0] is not None
|
|
66
|
+
and shape_node[0] is not None
|
|
67
|
+
and shape_weight[0] != shape_node[0]
|
|
68
|
+
):
|
|
69
|
+
raise ValueError(
|
|
70
|
+
f"Expected 'weight' to hold {shape_node[0]} "
|
|
71
|
+
f"values (got {shape_weight[0]} values)"
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
import keras
|
|
76
|
+
|
|
77
|
+
if keras.config.backend() == "jax":
|
|
78
|
+
try:
|
|
79
|
+
import jax
|
|
80
|
+
from jax.tree_util import register_pytree_node
|
|
81
|
+
|
|
82
|
+
register_pytree_node(
|
|
83
|
+
SelectOutput,
|
|
84
|
+
lambda s: (
|
|
85
|
+
(s.node_index, s.cluster_index, s.weight),
|
|
86
|
+
(s.num_nodes, s.num_clusters),
|
|
87
|
+
),
|
|
88
|
+
lambda aux, children: SelectOutput(
|
|
89
|
+
children[0], aux[0], children[1], aux[1], children[2]
|
|
90
|
+
),
|
|
91
|
+
)
|
|
92
|
+
except Exception:
|
|
93
|
+
pass
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class Select(layers.Layer):
|
|
98
|
+
r"""An abstract base class for implementing custom node selections as
|
|
99
|
+
described in the `"Understanding Pooling in Graph Neural Networks"
|
|
100
|
+
<https://arxiv.org/abs/1905.05178>`_ paper, which maps the nodes of an
|
|
101
|
+
input graph to supernodes in the coarsened graph.
|
|
102
|
+
"""
|
|
103
|
+
def reset_parameters(self):
|
|
104
|
+
r"""Resets all learnable parameters of the module."""
|
|
105
|
+
pass
|
|
106
|
+
|
|
107
|
+
def call(self, *args, **kwargs) -> SelectOutput:
|
|
108
|
+
raise NotImplementedError
|
|
109
|
+
|
|
110
|
+
def __repr__(self) -> str:
|
|
111
|
+
return f'{self.__class__.__name__}()'
|
|
112
|
+
|
|
@@ -0,0 +1,206 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from typing import Callable, Optional, Union
|
|
3
|
+
from keras import initializers, layers, ops
|
|
4
|
+
|
|
5
|
+
from .base import Select, SelectOutput
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
9
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
10
|
+
from k3_node.ops.creation import full
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def topk(
|
|
14
|
+
x,
|
|
15
|
+
ratio: Optional[Union[float, int]],
|
|
16
|
+
batch,
|
|
17
|
+
min_score: Optional[float] = None,
|
|
18
|
+
tol: float = 1e-7,
|
|
19
|
+
):
|
|
20
|
+
r"""Selects top-k items according to score and batch assignment.
|
|
21
|
+
|
|
22
|
+
Example:
|
|
23
|
+
```python
|
|
24
|
+
import numpy as np
|
|
25
|
+
from k3_node.layers import topk
|
|
26
|
+
|
|
27
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
28
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
29
|
+
|
|
30
|
+
score = np.random.rand(10).astype("float32")
|
|
31
|
+
perm = topk(score, ratio=0.5, batch=batch) # indices of the top 50% nodes per graph
|
|
32
|
+
print(tuple(perm.shape)) # (6,)
|
|
33
|
+
```
|
|
34
|
+
"""
|
|
35
|
+
# Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
|
|
36
|
+
x = ops.convert_to_tensor(x) if isinstance(x, np.ndarray) else x
|
|
37
|
+
batch = ops.convert_to_tensor(batch) if isinstance(batch, np.ndarray) else batch
|
|
38
|
+
if is_tracing(x) or is_tracing(batch):
|
|
39
|
+
return ops.arange(ops.shape(x)[0], dtype="int32")
|
|
40
|
+
|
|
41
|
+
batch = ops.cast(batch, dtype="int32")
|
|
42
|
+
num_nodes = ops.shape(x)[0]
|
|
43
|
+
|
|
44
|
+
num_graphs = ops.max(batch) + 1 if ops.shape(batch)[0] > 0 else 0
|
|
45
|
+
try:
|
|
46
|
+
num_graphs = int(num_graphs)
|
|
47
|
+
except (TypeError, ValueError):
|
|
48
|
+
pass
|
|
49
|
+
|
|
50
|
+
if min_score is not None:
|
|
51
|
+
scores_max = segment_max(x, batch, num_segments=num_graphs)
|
|
52
|
+
scores_max_expanded = ops.take(scores_max, batch, axis=0) - tol
|
|
53
|
+
scores_min = ops.minimum(scores_max_expanded, min_score)
|
|
54
|
+
mask = x > scores_min
|
|
55
|
+
perm = ops.where(mask)
|
|
56
|
+
if isinstance(perm, (tuple, list)):
|
|
57
|
+
perm = perm[0]
|
|
58
|
+
perm = ops.reshape(perm, (-1,))
|
|
59
|
+
return ops.cast(perm, "int32")
|
|
60
|
+
|
|
61
|
+
if ratio is not None:
|
|
62
|
+
ones = ops.ones((num_nodes,), dtype="int32")
|
|
63
|
+
num_nodes_per_graph = segment_sum(ones, batch, num_segments=num_graphs)
|
|
64
|
+
|
|
65
|
+
if ratio >= 1:
|
|
66
|
+
k = full(ops.shape(num_nodes_per_graph), int(ratio), dtype="int32")
|
|
67
|
+
else:
|
|
68
|
+
k = ops.cast(
|
|
69
|
+
ops.ceil(ratio * ops.cast(num_nodes_per_graph, x.dtype)),
|
|
70
|
+
dtype="int32",
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
# Composite key: sorts by batch ascending, then by score descending
|
|
74
|
+
score_span = ops.max(x) - ops.min(x) + 1.0
|
|
75
|
+
key = ops.cast(batch, x.dtype) * (score_span * 2.0) - x
|
|
76
|
+
perm = ops.argsort(key)
|
|
77
|
+
|
|
78
|
+
batch_sorted = ops.take(batch, perm, axis=0)
|
|
79
|
+
# ptr for cumsum
|
|
80
|
+
ptr = ops.concatenate(
|
|
81
|
+
[ops.zeros((1,), dtype="int32"), ops.cumsum(num_nodes_per_graph)[:-1]],
|
|
82
|
+
axis=0,
|
|
83
|
+
)
|
|
84
|
+
rank_in_graph = ops.arange(num_nodes, dtype="int32") - ops.take(ptr, batch_sorted, axis=0)
|
|
85
|
+
mask = rank_in_graph < ops.take(k, batch_sorted, axis=0)
|
|
86
|
+
valid_idx = ops.where(mask)
|
|
87
|
+
if isinstance(valid_idx, (tuple, list)):
|
|
88
|
+
valid_idx = valid_idx[0]
|
|
89
|
+
valid_idx = ops.reshape(valid_idx, (-1,))
|
|
90
|
+
return ops.cast(ops.take(perm, valid_idx, axis=0), "int32")
|
|
91
|
+
|
|
92
|
+
raise ValueError("At least one of 'ratio' and 'min_score' must be specified.")
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class SelectTopK(Select):
|
|
96
|
+
r"""Selects the top-:math:`k` nodes with highest projection scores.
|
|
97
|
+
|
|
98
|
+
Example:
|
|
99
|
+
```python
|
|
100
|
+
import numpy as np
|
|
101
|
+
from k3_node.layers import SelectTopK
|
|
102
|
+
|
|
103
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
104
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
105
|
+
|
|
106
|
+
select = SelectTopK(in_channels=8, ratio=0.5)
|
|
107
|
+
out = select(x, batch)
|
|
108
|
+
print(tuple(out.node_index.shape)) # (6,): indices of the kept nodes
|
|
109
|
+
print(tuple(out.weight.shape)) # (6,): their scores
|
|
110
|
+
```
|
|
111
|
+
"""
|
|
112
|
+
def __init__(
|
|
113
|
+
self,
|
|
114
|
+
in_channels: int,
|
|
115
|
+
ratio: Union[int, float] = 0.5,
|
|
116
|
+
min_score: Optional[float] = None,
|
|
117
|
+
act: Union[str, Callable] = "tanh",
|
|
118
|
+
**kwargs,
|
|
119
|
+
):
|
|
120
|
+
super().__init__(**kwargs)
|
|
121
|
+
|
|
122
|
+
if ratio is None and min_score is None:
|
|
123
|
+
raise ValueError(
|
|
124
|
+
f"At least one of 'ratio' and 'min_score' must be specified in '{self.__class__.__name__}'"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
self.in_channels = in_channels
|
|
128
|
+
self.ratio = ratio
|
|
129
|
+
self.min_score = min_score
|
|
130
|
+
self.act_fn = act if callable(act) else layers.Activation(act)
|
|
131
|
+
|
|
132
|
+
self.weight = self.add_weight(
|
|
133
|
+
shape=(1, in_channels),
|
|
134
|
+
initializer=initializers.RandomUniform(
|
|
135
|
+
minval=-1.0 / (in_channels**0.5), maxval=1.0 / (in_channels**0.5)
|
|
136
|
+
),
|
|
137
|
+
trainable=True,
|
|
138
|
+
name="weight",
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
def reset_parameters(self):
|
|
142
|
+
limit = 1.0 / (self.in_channels**0.5)
|
|
143
|
+
init = initializers.RandomUniform(minval=-limit, maxval=limit)
|
|
144
|
+
self.weight.assign(init(self.weight.shape, dtype=self.weight.dtype))
|
|
145
|
+
|
|
146
|
+
def build(self, input_shape=None):
|
|
147
|
+
if hasattr(self.act_fn, "built") and not self.act_fn.built:
|
|
148
|
+
self.act_fn.build(input_shape)
|
|
149
|
+
self.built = True
|
|
150
|
+
|
|
151
|
+
def call(self, x, batch=None) -> SelectOutput:
|
|
152
|
+
num_nodes = ops.shape(x)[0]
|
|
153
|
+
if batch is None:
|
|
154
|
+
batch = ops.zeros((num_nodes,), dtype="int32")
|
|
155
|
+
else:
|
|
156
|
+
batch = ops.cast(batch, dtype="int32")
|
|
157
|
+
|
|
158
|
+
if len(ops.shape(x)) == 1:
|
|
159
|
+
x = ops.expand_dims(x, axis=-1)
|
|
160
|
+
|
|
161
|
+
score = ops.sum(x * self.weight, axis=-1)
|
|
162
|
+
|
|
163
|
+
if is_tracing(x) or is_tracing(batch):
|
|
164
|
+
node_index = ops.arange(num_nodes, dtype="int32")
|
|
165
|
+
return SelectOutput(
|
|
166
|
+
node_index=node_index,
|
|
167
|
+
num_nodes=num_nodes,
|
|
168
|
+
cluster_index=node_index,
|
|
169
|
+
num_clusters=num_nodes,
|
|
170
|
+
weight=score,
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
if self.min_score is None:
|
|
174
|
+
norm_w = ops.sqrt(ops.sum(ops.power(self.weight, 2), axis=-1))
|
|
175
|
+
score = self.act_fn(score / norm_w)
|
|
176
|
+
else:
|
|
177
|
+
# Graph-wise softmax
|
|
178
|
+
num_graphs = ops.max(batch) + 1 if ops.shape(batch)[0] > 0 else 0
|
|
179
|
+
try:
|
|
180
|
+
num_graphs = int(num_graphs)
|
|
181
|
+
except (TypeError, ValueError):
|
|
182
|
+
pass
|
|
183
|
+
score_max = segment_max(score, batch, num_segments=num_graphs)
|
|
184
|
+
score_max_exp = ops.take(score_max, batch, axis=0)
|
|
185
|
+
exp_score = ops.exp(score - score_max_exp)
|
|
186
|
+
exp_sum = segment_sum(exp_score, batch, num_segments=num_graphs)
|
|
187
|
+
exp_sum_exp = ops.take(exp_sum, batch, axis=0)
|
|
188
|
+
score = exp_score / (exp_sum_exp + 1e-12)
|
|
189
|
+
|
|
190
|
+
node_index = topk(score, self.ratio, batch, self.min_score)
|
|
191
|
+
num_selected = ops.shape(node_index)[0]
|
|
192
|
+
|
|
193
|
+
return SelectOutput(
|
|
194
|
+
node_index=node_index,
|
|
195
|
+
num_nodes=num_nodes,
|
|
196
|
+
cluster_index=ops.arange(num_selected, dtype="int32"),
|
|
197
|
+
num_clusters=num_selected,
|
|
198
|
+
weight=ops.take(score, node_index, axis=0),
|
|
199
|
+
)
|
|
200
|
+
|
|
201
|
+
def __repr__(self) -> str:
|
|
202
|
+
if self.min_score is None:
|
|
203
|
+
arg = f"ratio={self.ratio}"
|
|
204
|
+
else:
|
|
205
|
+
arg = f"min_score={self.min_score}"
|
|
206
|
+
return f"{self.__class__.__name__}({self.in_channels}, {arg})"
|