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,137 @@
|
|
|
1
|
+
import math
|
|
2
|
+
from typing import List, Optional, Union
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from .base import Aggregation
|
|
7
|
+
from .utils import MultiheadAttentionBlock
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class PatchTransformerAggregation(Aggregation):
|
|
11
|
+
r"""Performs patch transformer aggregation in which the elements to
|
|
12
|
+
aggregate are processed by multi-head attention blocks across patches.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import PatchTransformerAggregation
|
|
18
|
+
|
|
19
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
20
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
21
|
+
|
|
22
|
+
aggr = PatchTransformerAggregation(in_channels=8, out_channels=16, patch_size=2, hidden_channels=8)
|
|
23
|
+
out = aggr(x, index=index, dim_size=2)
|
|
24
|
+
print(tuple(out.shape)) # (2, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
in_channels: int,
|
|
31
|
+
out_channels: int,
|
|
32
|
+
patch_size: int,
|
|
33
|
+
hidden_channels: int,
|
|
34
|
+
num_transformer_blocks: int = 1,
|
|
35
|
+
heads: int = 1,
|
|
36
|
+
dropout: float = 0.0,
|
|
37
|
+
aggr: Union[str, List[str]] = "mean",
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
super().__init__(**kwargs)
|
|
41
|
+
self.in_channels = in_channels
|
|
42
|
+
self.out_channels = out_channels
|
|
43
|
+
self.patch_size = patch_size
|
|
44
|
+
self.hidden_channels = hidden_channels
|
|
45
|
+
self.aggrs = [aggr] if isinstance(aggr, str) else list(aggr)
|
|
46
|
+
|
|
47
|
+
self.lin = layers.Dense(hidden_channels)
|
|
48
|
+
self.pad_projector = layers.Dense(hidden_channels)
|
|
49
|
+
self.blocks = [
|
|
50
|
+
MultiheadAttentionBlock(
|
|
51
|
+
channels=hidden_channels,
|
|
52
|
+
heads=heads,
|
|
53
|
+
layer_norm=True,
|
|
54
|
+
dropout=dropout,
|
|
55
|
+
)
|
|
56
|
+
for _ in range(num_transformer_blocks)
|
|
57
|
+
]
|
|
58
|
+
self.fc = layers.Dense(out_channels)
|
|
59
|
+
|
|
60
|
+
def reset_parameters(self):
|
|
61
|
+
self.lin.reset_parameters()
|
|
62
|
+
self.pad_projector.reset_parameters()
|
|
63
|
+
for block in self.blocks:
|
|
64
|
+
block.reset_parameters()
|
|
65
|
+
self.fc.reset_parameters()
|
|
66
|
+
|
|
67
|
+
def call(
|
|
68
|
+
self,
|
|
69
|
+
x,
|
|
70
|
+
index: Optional[any] = None,
|
|
71
|
+
ptr: Optional[any] = None,
|
|
72
|
+
dim_size: Optional[int] = None,
|
|
73
|
+
dim: int = -2,
|
|
74
|
+
max_num_elements: Optional[int] = None,
|
|
75
|
+
training: bool = False,
|
|
76
|
+
**kwargs,
|
|
77
|
+
):
|
|
78
|
+
if max_num_elements is None:
|
|
79
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
80
|
+
if is_tracing(x) or is_tracing(index):
|
|
81
|
+
if hasattr(x, "shape") and x.shape[0] is not None:
|
|
82
|
+
max_num_elements = int(x.shape[0])
|
|
83
|
+
else:
|
|
84
|
+
max_num_elements = 16
|
|
85
|
+
elif ptr is not None:
|
|
86
|
+
ptr_np = ops.convert_to_numpy(ptr)
|
|
87
|
+
count = ptr_np[1:] - ptr_np[:-1]
|
|
88
|
+
max_num_elements = int(np.max(count)) if len(count) > 0 else 1
|
|
89
|
+
else:
|
|
90
|
+
idx_np = ops.convert_to_numpy(index).astype(np.int64)
|
|
91
|
+
counts = np.bincount(idx_np)
|
|
92
|
+
max_num_elements = int(np.max(counts)) if len(counts) > 0 else 1
|
|
93
|
+
|
|
94
|
+
# Ensure max_num_elements is a multiple of patch_size
|
|
95
|
+
num_patches = max(math.ceil(max_num_elements / self.patch_size), 1)
|
|
96
|
+
target_elements = num_patches * self.patch_size
|
|
97
|
+
|
|
98
|
+
x_dense, _ = self.to_dense_batch(
|
|
99
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
100
|
+
max_num_elements=target_elements,
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
B = ops.shape(x_dense)[0]
|
|
104
|
+
x_proj = self.lin(x_dense) # [B, target_elements, hidden_channels]
|
|
105
|
+
|
|
106
|
+
# Reshape to patches: [B, num_patches, patch_size * hidden_channels]
|
|
107
|
+
x_patches = ops.reshape(x_proj, (B, num_patches, self.patch_size * self.hidden_channels))
|
|
108
|
+
x_patches = self.pad_projector(x_patches) # [B, num_patches, hidden_channels]
|
|
109
|
+
|
|
110
|
+
# Process through transformer blocks
|
|
111
|
+
for block in self.blocks:
|
|
112
|
+
x_patches = block(x_patches, x_patches, training=training)
|
|
113
|
+
|
|
114
|
+
outs = []
|
|
115
|
+
for aggr_mode in self.aggrs:
|
|
116
|
+
if aggr_mode == "mean":
|
|
117
|
+
outs.append(ops.mean(x_patches, axis=1))
|
|
118
|
+
elif aggr_mode == "sum":
|
|
119
|
+
outs.append(ops.sum(x_patches, axis=1))
|
|
120
|
+
elif aggr_mode == "max":
|
|
121
|
+
outs.append(ops.max(x_patches, axis=1))
|
|
122
|
+
elif aggr_mode == "min":
|
|
123
|
+
outs.append(ops.min(x_patches, axis=1))
|
|
124
|
+
elif aggr_mode == "var":
|
|
125
|
+
outs.append(ops.var(x_patches, axis=1))
|
|
126
|
+
elif aggr_mode == "std":
|
|
127
|
+
outs.append(ops.std(x_patches, axis=1))
|
|
128
|
+
|
|
129
|
+
combined = ops.concatenate(outs, axis=-1) if len(outs) > 1 else outs[0]
|
|
130
|
+
return self.fc(combined)
|
|
131
|
+
|
|
132
|
+
def __repr__(self) -> str:
|
|
133
|
+
return (
|
|
134
|
+
f"{self.__class__.__name__}({self.in_channels}, "
|
|
135
|
+
f"{self.out_channels}, patch_size={self.patch_size})"
|
|
136
|
+
)
|
|
137
|
+
|
|
@@ -0,0 +1,125 @@
|
|
|
1
|
+
from typing import List, Optional, Union
|
|
2
|
+
from keras import ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from .base import Aggregation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class QuantileAggregation(Aggregation):
|
|
9
|
+
r"""An aggregation operator that returns the feature-wise :math:`q`-th
|
|
10
|
+
quantile of a set :math:`\mathcal{X}`.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import QuantileAggregation
|
|
16
|
+
|
|
17
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
18
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
19
|
+
|
|
20
|
+
aggr = QuantileAggregation(q=0.75)
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 8)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
interpolations = {"linear", "lower", "higher", "nearest", "midpoint"}
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
q: Union[float, List[float]],
|
|
30
|
+
interpolation: str = "linear",
|
|
31
|
+
fill_value: float = 0.0,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
|
|
36
|
+
qs = [q] if not isinstance(q, (list, tuple)) else list(q)
|
|
37
|
+
if len(qs) == 0:
|
|
38
|
+
raise ValueError("Provide at least one quantile value for `q`.")
|
|
39
|
+
if not all(0.0 <= quantile <= 1.0 for quantile in qs):
|
|
40
|
+
raise ValueError("`q` must be in the range [0, 1].")
|
|
41
|
+
if interpolation not in self.interpolations:
|
|
42
|
+
raise ValueError(f"Invalid interpolation method got ('{interpolation}')")
|
|
43
|
+
|
|
44
|
+
self.q = qs
|
|
45
|
+
self.interpolation = interpolation
|
|
46
|
+
self.fill_value = fill_value
|
|
47
|
+
|
|
48
|
+
def call(
|
|
49
|
+
self,
|
|
50
|
+
x,
|
|
51
|
+
index: Optional[any] = None,
|
|
52
|
+
ptr: Optional[any] = None,
|
|
53
|
+
dim_size: Optional[int] = None,
|
|
54
|
+
dim: int = -2,
|
|
55
|
+
**kwargs,
|
|
56
|
+
):
|
|
57
|
+
self.assert_index_present(index)
|
|
58
|
+
from k3_node.ops.segment import segment_sum
|
|
59
|
+
|
|
60
|
+
# Sort every set's values in a dense [sets, max_size, features] tensor; the padding sorts
|
|
61
|
+
# last (+inf) and is then zeroed. Static shapes and differentiable, like PyG's version.
|
|
62
|
+
dense, _ = self.to_dense_batch(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
63
|
+
fill_value=float("inf"))
|
|
64
|
+
dense = ops.sort(dense, axis=1)
|
|
65
|
+
dense = ops.where(ops.isinf(dense), ops.zeros_like(dense), dense)
|
|
66
|
+
|
|
67
|
+
index_i = ops.cast(index, "int32")
|
|
68
|
+
count = segment_sum(ops.ones_like(index_i), index_i, num_segments=ops.shape(dense)[0])
|
|
69
|
+
count_f = ops.cast(count, dense.dtype)
|
|
70
|
+
last = ops.maximum(count - 1, 0)
|
|
71
|
+
|
|
72
|
+
def gather(position): # the value at `position` of every set, per feature
|
|
73
|
+
position = ops.minimum(ops.maximum(ops.cast(position, "int32"), 0), last)
|
|
74
|
+
position = ops.broadcast_to(ops.reshape(position, (-1, 1, 1)),
|
|
75
|
+
(ops.shape(dense)[0], 1, ops.shape(dense)[2]))
|
|
76
|
+
return ops.take_along_axis(dense, position, axis=1)[:, 0]
|
|
77
|
+
|
|
78
|
+
outs = []
|
|
79
|
+
for q_val in self.q:
|
|
80
|
+
q_point = q_val * (count_f - 1.0)
|
|
81
|
+
if self.interpolation == "lower":
|
|
82
|
+
quantile = gather(ops.floor(q_point))
|
|
83
|
+
elif self.interpolation == "higher":
|
|
84
|
+
quantile = gather(ops.ceil(q_point))
|
|
85
|
+
elif self.interpolation == "nearest":
|
|
86
|
+
quantile = gather(ops.round(q_point))
|
|
87
|
+
else:
|
|
88
|
+
low, high = gather(ops.floor(q_point)), gather(ops.ceil(q_point))
|
|
89
|
+
if self.interpolation == "linear":
|
|
90
|
+
frac = ops.expand_dims(q_point - ops.floor(q_point), -1)
|
|
91
|
+
quantile = low + (high - low) * frac
|
|
92
|
+
else: # midpoint
|
|
93
|
+
quantile = 0.5 * low + 0.5 * high
|
|
94
|
+
empty = ops.expand_dims(count == 0, -1)
|
|
95
|
+
outs.append(ops.where(empty, ops.cast(self.fill_value, quantile.dtype), quantile))
|
|
96
|
+
return outs[0] if len(outs) == 1 else ops.concatenate(outs, axis=-1)
|
|
97
|
+
|
|
98
|
+
def __repr__(self) -> str:
|
|
99
|
+
q_str = self.q[0] if len(self.q) == 1 else self.q
|
|
100
|
+
return f"{self.__class__.__name__}(q={q_str})"
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class MedianAggregation(QuantileAggregation):
|
|
104
|
+
r"""An aggregation operator that returns the feature-wise median of a set.
|
|
105
|
+
|
|
106
|
+
Example:
|
|
107
|
+
```python
|
|
108
|
+
import numpy as np
|
|
109
|
+
from k3_node.layers import MedianAggregation
|
|
110
|
+
|
|
111
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
112
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
113
|
+
|
|
114
|
+
aggr = MedianAggregation()
|
|
115
|
+
out = aggr(x, index=index, dim_size=2)
|
|
116
|
+
print(tuple(out.shape)) # (2, 8)
|
|
117
|
+
```
|
|
118
|
+
"""
|
|
119
|
+
|
|
120
|
+
def __init__(self, fill_value: float = 0.0, **kwargs):
|
|
121
|
+
super().__init__(0.5, "lower", fill_value=fill_value, **kwargs)
|
|
122
|
+
|
|
123
|
+
def __repr__(self) -> str:
|
|
124
|
+
return f"{self.__class__.__name__}()"
|
|
125
|
+
|
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
from typing import Union
|
|
2
|
+
|
|
3
|
+
from .base import Aggregation
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def aggregation_resolver(
|
|
7
|
+
query: Union[str, Aggregation],
|
|
8
|
+
*args,
|
|
9
|
+
**kwargs,
|
|
10
|
+
) -> Aggregation:
|
|
11
|
+
r"""Resolves an aggregation string or instance to an `Aggregation` object.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import aggregation_resolver
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
20
|
+
|
|
21
|
+
aggr = aggregation_resolver("mean") # build an Aggregation from its name
|
|
22
|
+
print(type(aggr).__name__, tuple(aggr(x, index=index, dim_size=2).shape)) # MeanAggregation (2, 8)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
if isinstance(query, Aggregation):
|
|
26
|
+
return query
|
|
27
|
+
|
|
28
|
+
if not isinstance(query, str):
|
|
29
|
+
raise ValueError(f"Expected string or Aggregation instance, got {type(query)}")
|
|
30
|
+
|
|
31
|
+
query_norm = query.lower().strip()
|
|
32
|
+
|
|
33
|
+
from .basic import (
|
|
34
|
+
MaxAggregation,
|
|
35
|
+
MeanAggregation,
|
|
36
|
+
MinAggregation,
|
|
37
|
+
MulAggregation,
|
|
38
|
+
PowerMeanAggregation,
|
|
39
|
+
SoftmaxAggregation,
|
|
40
|
+
StdAggregation,
|
|
41
|
+
SumAggregation,
|
|
42
|
+
VarAggregation,
|
|
43
|
+
)
|
|
44
|
+
from .quantile import MedianAggregation, QuantileAggregation
|
|
45
|
+
from .variance_preserving import VariancePreservingAggregation
|
|
46
|
+
|
|
47
|
+
AGGR_DICT = {
|
|
48
|
+
"sum": SumAggregation,
|
|
49
|
+
"add": SumAggregation,
|
|
50
|
+
"mean": MeanAggregation,
|
|
51
|
+
"max": MaxAggregation,
|
|
52
|
+
"min": MinAggregation,
|
|
53
|
+
"mul": MulAggregation,
|
|
54
|
+
"var": VarAggregation,
|
|
55
|
+
"std": StdAggregation,
|
|
56
|
+
"softmax": SoftmaxAggregation,
|
|
57
|
+
"powermean": PowerMeanAggregation,
|
|
58
|
+
"median": MedianAggregation,
|
|
59
|
+
"quantile": QuantileAggregation,
|
|
60
|
+
"variance_preserving": VariancePreservingAggregation,
|
|
61
|
+
"vpa": VariancePreservingAggregation,
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
if query_norm in AGGR_DICT:
|
|
65
|
+
return AGGR_DICT[query_norm](*args, **kwargs)
|
|
66
|
+
|
|
67
|
+
raise ValueError(f"Could not resolve aggregation '{query}'")
|
|
68
|
+
|
|
@@ -0,0 +1,133 @@
|
|
|
1
|
+
from typing import Any, Dict, List, Optional, Union
|
|
2
|
+
from keras import initializers, ops
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from .base import Aggregation
|
|
6
|
+
from k3_node.ops.segment import segment_sum
|
|
7
|
+
from k3_node.ops.creation import full
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class DegreeScalerAggregation(Aggregation):
|
|
11
|
+
r"""Combines one or more aggregators and transforms its output with one or
|
|
12
|
+
more scalers as introduced in the `"Principal Neighbourhood Aggregation for
|
|
13
|
+
Graph Nets" <https://arxiv.org/abs/2004.05718>`_ paper.
|
|
14
|
+
|
|
15
|
+
Example:
|
|
16
|
+
```python
|
|
17
|
+
import numpy as np
|
|
18
|
+
from k3_node.layers import DegreeScalerAggregation
|
|
19
|
+
|
|
20
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
21
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
22
|
+
|
|
23
|
+
deg = np.array([0, 3, 5, 2]) # in-degree histogram of the training graphs
|
|
24
|
+
aggr = DegreeScalerAggregation(aggr=["mean", "max"], scaler=["identity", "amplification"], deg=deg)
|
|
25
|
+
out = aggr(x, index=index, dim_size=2)
|
|
26
|
+
print(tuple(out.shape)) # (2, 32): 2 aggregators x 2 scalers x 8 features
|
|
27
|
+
```
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
def __init__(
|
|
31
|
+
self,
|
|
32
|
+
aggr: Union[str, List[str], Aggregation],
|
|
33
|
+
scaler: Union[str, List[str]],
|
|
34
|
+
deg,
|
|
35
|
+
train_norm: bool = False,
|
|
36
|
+
aggr_kwargs: Optional[List[Dict[str, Any]]] = None,
|
|
37
|
+
**kwargs,
|
|
38
|
+
):
|
|
39
|
+
super().__init__(**kwargs)
|
|
40
|
+
|
|
41
|
+
from .resolver import aggregation_resolver
|
|
42
|
+
from .multi import MultiAggregation
|
|
43
|
+
|
|
44
|
+
if isinstance(aggr, (str, Aggregation)):
|
|
45
|
+
self.aggr = aggregation_resolver(aggr, **(aggr_kwargs or {}))
|
|
46
|
+
elif isinstance(aggr, (tuple, list)):
|
|
47
|
+
self.aggr = MultiAggregation(aggr, aggr_kwargs)
|
|
48
|
+
else:
|
|
49
|
+
raise ValueError(
|
|
50
|
+
f"Only strings, list, tuples and instances of "
|
|
51
|
+
f"`Aggregation` are valid aggregation schemes (got '{type(aggr)}')"
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
self.scaler = [scaler] if isinstance(scaler, str) else list(scaler)
|
|
55
|
+
|
|
56
|
+
deg_np = ops.convert_to_numpy(deg).astype(np.float32)
|
|
57
|
+
N = float(np.sum(deg_np))
|
|
58
|
+
bin_degree = np.arange(len(deg_np), dtype=np.float32)
|
|
59
|
+
|
|
60
|
+
self.init_avg_deg_lin = float(np.sum(bin_degree * deg_np)) / max(N, 1.0)
|
|
61
|
+
self.init_avg_deg_log = float(np.sum(np.log(bin_degree + 1.0) * deg_np)) / max(N, 1.0)
|
|
62
|
+
self.train_norm = train_norm
|
|
63
|
+
|
|
64
|
+
if train_norm:
|
|
65
|
+
self.avg_deg_lin = self.add_weight(
|
|
66
|
+
shape=(1,),
|
|
67
|
+
initializer=initializers.Constant(self.init_avg_deg_lin),
|
|
68
|
+
trainable=True,
|
|
69
|
+
name="avg_deg_lin",
|
|
70
|
+
)
|
|
71
|
+
self.avg_deg_log = self.add_weight(
|
|
72
|
+
shape=(1,),
|
|
73
|
+
initializer=initializers.Constant(self.init_avg_deg_log),
|
|
74
|
+
trainable=True,
|
|
75
|
+
name="avg_deg_log",
|
|
76
|
+
)
|
|
77
|
+
else:
|
|
78
|
+
self.avg_deg_lin = self.init_avg_deg_lin
|
|
79
|
+
self.avg_deg_log = self.init_avg_deg_log
|
|
80
|
+
|
|
81
|
+
def reset_parameters(self):
|
|
82
|
+
if hasattr(self.aggr, "reset_parameters"):
|
|
83
|
+
self.aggr.reset_parameters()
|
|
84
|
+
if self.train_norm:
|
|
85
|
+
self.avg_deg_lin.assign(full((1,), self.init_avg_deg_lin, dtype=self.avg_deg_lin.dtype))
|
|
86
|
+
self.avg_deg_log.assign(full((1,), self.init_avg_deg_log, dtype=self.avg_deg_log.dtype))
|
|
87
|
+
|
|
88
|
+
def call(
|
|
89
|
+
self,
|
|
90
|
+
x,
|
|
91
|
+
index: Optional[any] = None,
|
|
92
|
+
ptr: Optional[any] = None,
|
|
93
|
+
dim_size: Optional[int] = None,
|
|
94
|
+
dim: int = -2,
|
|
95
|
+
**kwargs,
|
|
96
|
+
):
|
|
97
|
+
self.assert_index_present(index)
|
|
98
|
+
|
|
99
|
+
out = self.aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim)
|
|
100
|
+
|
|
101
|
+
index = ops.cast(index, dtype="int32")
|
|
102
|
+
if dim_size is None: # a tensor while tracing; don't test its truth value
|
|
103
|
+
dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
104
|
+
|
|
105
|
+
# Compute degree per index
|
|
106
|
+
ones = ops.ones((ops.shape(index)[0],), dtype=out.dtype)
|
|
107
|
+
deg = segment_sum(ones, index, num_segments=dim_size)
|
|
108
|
+
deg = ops.reshape(deg, (dim_size,) + (1,) * (len(ops.shape(out)) - 1))
|
|
109
|
+
|
|
110
|
+
avg_deg_log = self.avg_deg_log
|
|
111
|
+
avg_deg_lin = self.avg_deg_lin
|
|
112
|
+
|
|
113
|
+
outs = []
|
|
114
|
+
for scaler in self.scaler:
|
|
115
|
+
if scaler == "identity":
|
|
116
|
+
out_scaler = out
|
|
117
|
+
elif scaler == "amplification":
|
|
118
|
+
out_scaler = out * (ops.log(deg + 1.0) / avg_deg_log)
|
|
119
|
+
elif scaler == "attenuation":
|
|
120
|
+
out_scaler = out * (avg_deg_log / ops.log(ops.maximum(deg, 1.0) + 1.0))
|
|
121
|
+
elif scaler == "linear":
|
|
122
|
+
out_scaler = out * (deg / avg_deg_lin)
|
|
123
|
+
elif scaler == "inverse_linear":
|
|
124
|
+
out_scaler = out * (avg_deg_lin / ops.maximum(deg, 1.0))
|
|
125
|
+
else:
|
|
126
|
+
raise ValueError(f"Unknown scaler '{scaler}'")
|
|
127
|
+
outs.append(out_scaler)
|
|
128
|
+
|
|
129
|
+
return ops.concatenate(outs, axis=-1) if len(outs) > 1 else outs[0]
|
|
130
|
+
|
|
131
|
+
def __repr__(self) -> str:
|
|
132
|
+
return f"{self.__class__.__name__}(aggr={self.aggr}, scaler={self.scaler})"
|
|
133
|
+
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Set2Set(Aggregation):
|
|
9
|
+
r"""The Set2Set aggregation operator based on iterative content-based
|
|
10
|
+
attention, as described in the `"Order Matters: Sequence to sequence for
|
|
11
|
+
Sets" <https://arxiv.org/abs/1511.06391>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import Set2Set
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
20
|
+
|
|
21
|
+
aggr = Set2Set(in_channels=8, processing_steps=2)
|
|
22
|
+
out = aggr(x, index=index, dim_size=2)
|
|
23
|
+
print(tuple(out.shape)) # (2, 16)
|
|
24
|
+
```
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(self, in_channels: int, processing_steps: int, **kwargs):
|
|
28
|
+
super().__init__(**kwargs)
|
|
29
|
+
self.in_channels = in_channels
|
|
30
|
+
self.out_channels = 2 * in_channels
|
|
31
|
+
self.processing_steps = processing_steps
|
|
32
|
+
self.lstm_cell = layers.LSTMCell(in_channels)
|
|
33
|
+
|
|
34
|
+
def build(self, input_shape=None):
|
|
35
|
+
self.lstm_cell.build((None, self.out_channels))
|
|
36
|
+
super().build(input_shape)
|
|
37
|
+
|
|
38
|
+
def reset_parameters(self):
|
|
39
|
+
self.lstm_cell.reset_parameters()
|
|
40
|
+
|
|
41
|
+
def call(
|
|
42
|
+
self,
|
|
43
|
+
x,
|
|
44
|
+
index: Optional[any] = None,
|
|
45
|
+
ptr: Optional[any] = None,
|
|
46
|
+
dim_size: Optional[int] = None,
|
|
47
|
+
dim: int = -2,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
if ptr is not None and index is None:
|
|
51
|
+
from .base import ptr2index
|
|
52
|
+
index = ptr2index(ptr)
|
|
53
|
+
|
|
54
|
+
self.assert_index_present(index)
|
|
55
|
+
self.assert_two_dimensional_input(x, dim)
|
|
56
|
+
|
|
57
|
+
index = ops.cast(index, dtype="int32")
|
|
58
|
+
if dim_size is None: # a tensor while tracing; don't test its truth value
|
|
59
|
+
dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
60
|
+
|
|
61
|
+
# Initial hidden states: [dim_size, in_channels]
|
|
62
|
+
h = [
|
|
63
|
+
ops.zeros((dim_size, self.in_channels), dtype=x.dtype),
|
|
64
|
+
ops.zeros((dim_size, self.in_channels), dtype=x.dtype),
|
|
65
|
+
]
|
|
66
|
+
q_star = ops.zeros((dim_size, self.out_channels), dtype=x.dtype)
|
|
67
|
+
|
|
68
|
+
for _ in range(self.processing_steps):
|
|
69
|
+
q, h = self.lstm_cell(q_star, h)
|
|
70
|
+
q_taken = ops.take(q, index, axis=0)
|
|
71
|
+
e = ops.sum(x * q_taken, axis=-1, keepdims=True)
|
|
72
|
+
|
|
73
|
+
max_e = segment_max(e, index, num_segments=dim_size)
|
|
74
|
+
max_e_exp = ops.take(max_e, index, axis=0)
|
|
75
|
+
exp_e = ops.exp(e - max_e_exp)
|
|
76
|
+
sum_exp = segment_sum(exp_e, index, num_segments=dim_size)
|
|
77
|
+
sum_exp_exp = ops.take(sum_exp, index, axis=0)
|
|
78
|
+
a = exp_e / ops.maximum(sum_exp_exp, 1e-12)
|
|
79
|
+
|
|
80
|
+
r = segment_sum(a * x, index, num_segments=dim_size)
|
|
81
|
+
q_star = ops.concatenate([q, r], axis=-1)
|
|
82
|
+
|
|
83
|
+
return q_star
|
|
84
|
+
|
|
85
|
+
def __repr__(self) -> str:
|
|
86
|
+
return f"{self.__class__.__name__}({self.in_channels}, {self.out_channels})"
|
|
87
|
+
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
from .utils import PoolingByMultiheadAttention, SetAttentionBlock
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class SetTransformerAggregation(Aggregation):
|
|
9
|
+
r"""Performs "Set Transformer" aggregation in which the elements to
|
|
10
|
+
aggregate are processed by multi-head attention blocks, as described in
|
|
11
|
+
the `"Graph Neural Networks with Adaptive Readouts"
|
|
12
|
+
<https://arxiv.org/abs/2211.04952>`_ paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import SetTransformerAggregation
|
|
18
|
+
|
|
19
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
20
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
21
|
+
|
|
22
|
+
aggr = SetTransformerAggregation(channels=8, num_seed_points=2)
|
|
23
|
+
out = aggr(x, index=index, dim_size=2)
|
|
24
|
+
print(tuple(out.shape)) # (2, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
channels: int,
|
|
31
|
+
num_seed_points: int = 1,
|
|
32
|
+
num_encoder_blocks: int = 1,
|
|
33
|
+
num_decoder_blocks: int = 1,
|
|
34
|
+
heads: int = 1,
|
|
35
|
+
concat: bool = True,
|
|
36
|
+
layer_norm: bool = False,
|
|
37
|
+
dropout: float = 0.0,
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
super().__init__(**kwargs)
|
|
41
|
+
self.channels = channels
|
|
42
|
+
self.num_seed_points = num_seed_points
|
|
43
|
+
self.heads = heads
|
|
44
|
+
self.concat = concat
|
|
45
|
+
self.layer_norm = layer_norm
|
|
46
|
+
self.dropout = dropout
|
|
47
|
+
|
|
48
|
+
self.encoders = [
|
|
49
|
+
SetAttentionBlock(channels, heads, layer_norm, dropout)
|
|
50
|
+
for _ in range(num_encoder_blocks)
|
|
51
|
+
]
|
|
52
|
+
self.pma = PoolingByMultiheadAttention(
|
|
53
|
+
channels, num_seed_points, heads, layer_norm, dropout
|
|
54
|
+
)
|
|
55
|
+
self.decoders = [
|
|
56
|
+
SetAttentionBlock(channels, heads, layer_norm, dropout)
|
|
57
|
+
for _ in range(num_decoder_blocks)
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
def reset_parameters(self):
|
|
61
|
+
for encoder in self.encoders:
|
|
62
|
+
encoder.reset_parameters()
|
|
63
|
+
self.pma.reset_parameters()
|
|
64
|
+
for decoder in self.decoders:
|
|
65
|
+
decoder.reset_parameters()
|
|
66
|
+
|
|
67
|
+
def call(
|
|
68
|
+
self,
|
|
69
|
+
x,
|
|
70
|
+
index: Optional[any] = None,
|
|
71
|
+
ptr: Optional[any] = None,
|
|
72
|
+
dim_size: Optional[int] = None,
|
|
73
|
+
dim: int = -2,
|
|
74
|
+
max_num_elements: Optional[int] = None,
|
|
75
|
+
training: bool = False,
|
|
76
|
+
**kwargs,
|
|
77
|
+
):
|
|
78
|
+
x_dense, mask = self.to_dense_batch(
|
|
79
|
+
x, index=index, ptr=ptr, dim_size=dim_size, dim=dim,
|
|
80
|
+
max_num_elements=max_num_elements,
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
for encoder in self.encoders:
|
|
84
|
+
x_dense = encoder(x_dense, mask=mask, training=training)
|
|
85
|
+
|
|
86
|
+
x_dense = self.pma(x_dense, mask=mask, training=training)
|
|
87
|
+
|
|
88
|
+
for decoder in self.decoders:
|
|
89
|
+
x_dense = decoder(x_dense, training=training)
|
|
90
|
+
|
|
91
|
+
# Handle NaNs if any
|
|
92
|
+
x_dense = ops.where(ops.isnan(x_dense), 0.0, x_dense)
|
|
93
|
+
|
|
94
|
+
if self.concat:
|
|
95
|
+
B = ops.shape(x_dense)[0]
|
|
96
|
+
return ops.reshape(x_dense, (B, self.num_seed_points * self.channels))
|
|
97
|
+
else:
|
|
98
|
+
return ops.mean(x_dense, axis=1)
|
|
99
|
+
|
|
100
|
+
def __repr__(self) -> str:
|
|
101
|
+
return (
|
|
102
|
+
f"{self.__class__.__name__}({self.channels}, "
|
|
103
|
+
f"num_seed_points={self.num_seed_points}, "
|
|
104
|
+
f"heads={self.heads}, layer_norm={self.layer_norm}, "
|
|
105
|
+
f"dropout={self.dropout})"
|
|
106
|
+
)
|
|
107
|
+
|