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,412 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import initializers, ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
6
|
+
from k3_node.ops.creation import full
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class SumAggregation(Aggregation):
|
|
10
|
+
r"""An aggregation operator that sums up features across a set of elements.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import SumAggregation
|
|
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 = SumAggregation()
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 8)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def call(
|
|
27
|
+
self,
|
|
28
|
+
x,
|
|
29
|
+
index: Optional[any] = None,
|
|
30
|
+
ptr: Optional[any] = None,
|
|
31
|
+
dim_size: Optional[int] = None,
|
|
32
|
+
dim: int = -2,
|
|
33
|
+
**kwargs,
|
|
34
|
+
):
|
|
35
|
+
return self.reduce(x, index, ptr, dim_size, dim, reduce="sum")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class MeanAggregation(Aggregation):
|
|
39
|
+
r"""An aggregation operator that averages features across a set of elements.
|
|
40
|
+
|
|
41
|
+
Example:
|
|
42
|
+
```python
|
|
43
|
+
import numpy as np
|
|
44
|
+
from k3_node.layers import MeanAggregation
|
|
45
|
+
|
|
46
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
47
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
48
|
+
|
|
49
|
+
aggr = MeanAggregation()
|
|
50
|
+
out = aggr(x, index=index, dim_size=2)
|
|
51
|
+
print(tuple(out.shape)) # (2, 8)
|
|
52
|
+
```
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
def call(
|
|
56
|
+
self,
|
|
57
|
+
x,
|
|
58
|
+
index: Optional[any] = None,
|
|
59
|
+
ptr: Optional[any] = None,
|
|
60
|
+
dim_size: Optional[int] = None,
|
|
61
|
+
dim: int = -2,
|
|
62
|
+
**kwargs,
|
|
63
|
+
):
|
|
64
|
+
return self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class MaxAggregation(Aggregation):
|
|
68
|
+
r"""An aggregation operator that takes the feature-wise maximum across a set of elements.
|
|
69
|
+
|
|
70
|
+
Example:
|
|
71
|
+
```python
|
|
72
|
+
import numpy as np
|
|
73
|
+
from k3_node.layers import MaxAggregation
|
|
74
|
+
|
|
75
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
76
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
77
|
+
|
|
78
|
+
aggr = MaxAggregation()
|
|
79
|
+
out = aggr(x, index=index, dim_size=2)
|
|
80
|
+
print(tuple(out.shape)) # (2, 8)
|
|
81
|
+
```
|
|
82
|
+
"""
|
|
83
|
+
|
|
84
|
+
def call(
|
|
85
|
+
self,
|
|
86
|
+
x,
|
|
87
|
+
index: Optional[any] = None,
|
|
88
|
+
ptr: Optional[any] = None,
|
|
89
|
+
dim_size: Optional[int] = None,
|
|
90
|
+
dim: int = -2,
|
|
91
|
+
**kwargs,
|
|
92
|
+
):
|
|
93
|
+
return self.reduce(x, index, ptr, dim_size, dim, reduce="max")
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
class MinAggregation(Aggregation):
|
|
97
|
+
r"""An aggregation operator that takes the feature-wise minimum across a set of elements.
|
|
98
|
+
|
|
99
|
+
Example:
|
|
100
|
+
```python
|
|
101
|
+
import numpy as np
|
|
102
|
+
from k3_node.layers import MinAggregation
|
|
103
|
+
|
|
104
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
105
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
106
|
+
|
|
107
|
+
aggr = MinAggregation()
|
|
108
|
+
out = aggr(x, index=index, dim_size=2)
|
|
109
|
+
print(tuple(out.shape)) # (2, 8)
|
|
110
|
+
```
|
|
111
|
+
"""
|
|
112
|
+
|
|
113
|
+
def call(
|
|
114
|
+
self,
|
|
115
|
+
x,
|
|
116
|
+
index: Optional[any] = None,
|
|
117
|
+
ptr: Optional[any] = None,
|
|
118
|
+
dim_size: Optional[int] = None,
|
|
119
|
+
dim: int = -2,
|
|
120
|
+
**kwargs,
|
|
121
|
+
):
|
|
122
|
+
return self.reduce(x, index, ptr, dim_size, dim, reduce="min")
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class MulAggregation(Aggregation):
|
|
126
|
+
r"""An aggregation operator that multiplies features across a set of elements.
|
|
127
|
+
|
|
128
|
+
Example:
|
|
129
|
+
```python
|
|
130
|
+
import numpy as np
|
|
131
|
+
from k3_node.layers import MulAggregation
|
|
132
|
+
|
|
133
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
134
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
135
|
+
|
|
136
|
+
aggr = MulAggregation()
|
|
137
|
+
out = aggr(x, index=index, dim_size=2)
|
|
138
|
+
print(tuple(out.shape)) # (2, 8)
|
|
139
|
+
```
|
|
140
|
+
"""
|
|
141
|
+
|
|
142
|
+
def call(
|
|
143
|
+
self,
|
|
144
|
+
x,
|
|
145
|
+
index: Optional[any] = None,
|
|
146
|
+
ptr: Optional[any] = None,
|
|
147
|
+
dim_size: Optional[int] = None,
|
|
148
|
+
dim: int = -2,
|
|
149
|
+
**kwargs,
|
|
150
|
+
):
|
|
151
|
+
self.assert_index_present(index)
|
|
152
|
+
return self.reduce(x, index, ptr, dim_size, dim, reduce="mul")
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
class VarAggregation(Aggregation):
|
|
156
|
+
r"""An aggregation operator that takes the feature-wise variance across a set of elements.
|
|
157
|
+
|
|
158
|
+
Example:
|
|
159
|
+
```python
|
|
160
|
+
import numpy as np
|
|
161
|
+
from k3_node.layers import VarAggregation
|
|
162
|
+
|
|
163
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
164
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
165
|
+
|
|
166
|
+
aggr = VarAggregation()
|
|
167
|
+
out = aggr(x, index=index, dim_size=2)
|
|
168
|
+
print(tuple(out.shape)) # (2, 8)
|
|
169
|
+
```
|
|
170
|
+
"""
|
|
171
|
+
|
|
172
|
+
def __init__(self, semi_grad: bool = False, **kwargs):
|
|
173
|
+
super().__init__(**kwargs)
|
|
174
|
+
self.semi_grad = semi_grad
|
|
175
|
+
|
|
176
|
+
def call(
|
|
177
|
+
self,
|
|
178
|
+
x,
|
|
179
|
+
index: Optional[any] = None,
|
|
180
|
+
ptr: Optional[any] = None,
|
|
181
|
+
dim_size: Optional[int] = None,
|
|
182
|
+
dim: int = -2,
|
|
183
|
+
**kwargs,
|
|
184
|
+
):
|
|
185
|
+
mean = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
|
|
186
|
+
x_sq = ops.power(x, 2)
|
|
187
|
+
if self.semi_grad:
|
|
188
|
+
x_sq = ops.stop_gradient(x_sq)
|
|
189
|
+
mean2 = self.reduce(x_sq, index, ptr, dim_size, dim, reduce="mean")
|
|
190
|
+
return mean2 - ops.power(mean, 2)
|
|
191
|
+
|
|
192
|
+
def __repr__(self) -> str:
|
|
193
|
+
return f"{self.__class__.__name__}(semi_grad={self.semi_grad})"
|
|
194
|
+
|
|
195
|
+
|
|
196
|
+
class StdAggregation(Aggregation):
|
|
197
|
+
r"""An aggregation operator that takes the feature-wise standard deviation across a set of elements.
|
|
198
|
+
|
|
199
|
+
Example:
|
|
200
|
+
```python
|
|
201
|
+
import numpy as np
|
|
202
|
+
from k3_node.layers import StdAggregation
|
|
203
|
+
|
|
204
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
205
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
206
|
+
|
|
207
|
+
aggr = StdAggregation()
|
|
208
|
+
out = aggr(x, index=index, dim_size=2)
|
|
209
|
+
print(tuple(out.shape)) # (2, 8)
|
|
210
|
+
```
|
|
211
|
+
"""
|
|
212
|
+
|
|
213
|
+
def __init__(self, semi_grad: bool = False, **kwargs):
|
|
214
|
+
super().__init__(**kwargs)
|
|
215
|
+
self.semi_grad = semi_grad
|
|
216
|
+
self.var_aggr = VarAggregation(semi_grad=semi_grad)
|
|
217
|
+
|
|
218
|
+
def call(
|
|
219
|
+
self,
|
|
220
|
+
x,
|
|
221
|
+
index: Optional[any] = None,
|
|
222
|
+
ptr: Optional[any] = None,
|
|
223
|
+
dim_size: Optional[int] = None,
|
|
224
|
+
dim: int = -2,
|
|
225
|
+
**kwargs,
|
|
226
|
+
):
|
|
227
|
+
var = self.var_aggr(x, index, ptr, dim_size, dim)
|
|
228
|
+
out = ops.sqrt(ops.maximum(var, 1e-5))
|
|
229
|
+
out = ops.where(out <= (1e-5**0.5), 0.0, out)
|
|
230
|
+
return out
|
|
231
|
+
|
|
232
|
+
def __repr__(self) -> str:
|
|
233
|
+
return f"{self.__class__.__name__}(semi_grad={self.semi_grad})"
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
class SoftmaxAggregation(Aggregation):
|
|
237
|
+
r"""The softmax aggregation operator based on a temperature term.
|
|
238
|
+
|
|
239
|
+
Example:
|
|
240
|
+
```python
|
|
241
|
+
import numpy as np
|
|
242
|
+
from k3_node.layers import SoftmaxAggregation
|
|
243
|
+
|
|
244
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
245
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
246
|
+
|
|
247
|
+
aggr = SoftmaxAggregation(learn=True)
|
|
248
|
+
out = aggr(x, index=index, dim_size=2)
|
|
249
|
+
print(tuple(out.shape)) # (2, 8)
|
|
250
|
+
```
|
|
251
|
+
"""
|
|
252
|
+
|
|
253
|
+
def __init__(
|
|
254
|
+
self,
|
|
255
|
+
t: float = 1.0,
|
|
256
|
+
learn: bool = False,
|
|
257
|
+
semi_grad: bool = False,
|
|
258
|
+
channels: int = 1,
|
|
259
|
+
**kwargs,
|
|
260
|
+
):
|
|
261
|
+
super().__init__(**kwargs)
|
|
262
|
+
|
|
263
|
+
if learn and semi_grad:
|
|
264
|
+
raise ValueError(
|
|
265
|
+
f"Cannot enable 'semi_grad' in '{self.__class__.__name__}' in "
|
|
266
|
+
f"case the temperature term 't' is learnable"
|
|
267
|
+
)
|
|
268
|
+
|
|
269
|
+
if not learn and channels != 1:
|
|
270
|
+
raise ValueError(
|
|
271
|
+
f"Cannot set 'channels' greater than '1' in case "
|
|
272
|
+
f"'{self.__class__.__name__}' is not trainable"
|
|
273
|
+
)
|
|
274
|
+
|
|
275
|
+
self._init_t = t
|
|
276
|
+
self.learn = learn
|
|
277
|
+
self.semi_grad = semi_grad
|
|
278
|
+
self.channels = channels
|
|
279
|
+
|
|
280
|
+
if learn:
|
|
281
|
+
self.t = self.add_weight(
|
|
282
|
+
shape=(channels,),
|
|
283
|
+
initializer=initializers.Constant(t),
|
|
284
|
+
trainable=True,
|
|
285
|
+
name="t",
|
|
286
|
+
)
|
|
287
|
+
else:
|
|
288
|
+
self.t = t
|
|
289
|
+
|
|
290
|
+
def reset_parameters(self):
|
|
291
|
+
if self.learn:
|
|
292
|
+
self.t.assign(full((self.channels,), self._init_t, dtype=self.t.dtype))
|
|
293
|
+
|
|
294
|
+
def call(
|
|
295
|
+
self,
|
|
296
|
+
x,
|
|
297
|
+
index: Optional[any] = None,
|
|
298
|
+
ptr: Optional[any] = None,
|
|
299
|
+
dim_size: Optional[int] = None,
|
|
300
|
+
dim: int = -2,
|
|
301
|
+
**kwargs,
|
|
302
|
+
):
|
|
303
|
+
t = self.t
|
|
304
|
+
if self.channels != 1:
|
|
305
|
+
self.assert_two_dimensional_input(x, dim)
|
|
306
|
+
t = ops.reshape(t, (1, self.channels))
|
|
307
|
+
|
|
308
|
+
alpha = x
|
|
309
|
+
if self.learn or t != 1.0:
|
|
310
|
+
alpha = x * t
|
|
311
|
+
|
|
312
|
+
if not self.learn and self.semi_grad:
|
|
313
|
+
alpha = ops.stop_gradient(alpha)
|
|
314
|
+
|
|
315
|
+
# Graph-wise softmax over segments
|
|
316
|
+
index = ops.cast(index, dtype="int32")
|
|
317
|
+
if dim_size is None: # a tensor while tracing; don't test its truth value
|
|
318
|
+
dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
319
|
+
|
|
320
|
+
max_val = segment_max(alpha, index, num_segments=dim_size)
|
|
321
|
+
max_exp = ops.take(max_val, index, axis=0)
|
|
322
|
+
exp_alpha = ops.exp(alpha - max_exp)
|
|
323
|
+
sum_exp = segment_sum(exp_alpha, index, num_segments=dim_size)
|
|
324
|
+
sum_exp_taken = ops.take(sum_exp, index, axis=0)
|
|
325
|
+
alpha_sm = exp_alpha / ops.maximum(sum_exp_taken, 1e-12)
|
|
326
|
+
|
|
327
|
+
return self.reduce(x * alpha_sm, index, ptr, dim_size, dim, reduce="sum")
|
|
328
|
+
|
|
329
|
+
def __repr__(self) -> str:
|
|
330
|
+
return f"{self.__class__.__name__}(learn={self.learn})"
|
|
331
|
+
|
|
332
|
+
|
|
333
|
+
class PowerMeanAggregation(Aggregation):
|
|
334
|
+
r"""The powermean aggregation operator based on a power term.
|
|
335
|
+
|
|
336
|
+
Example:
|
|
337
|
+
```python
|
|
338
|
+
import numpy as np
|
|
339
|
+
from k3_node.layers import PowerMeanAggregation
|
|
340
|
+
|
|
341
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
342
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
343
|
+
|
|
344
|
+
aggr = PowerMeanAggregation(learn=True)
|
|
345
|
+
out = aggr(x, index=index, dim_size=2)
|
|
346
|
+
print(tuple(out.shape)) # (2, 8)
|
|
347
|
+
```
|
|
348
|
+
"""
|
|
349
|
+
|
|
350
|
+
def __init__(
|
|
351
|
+
self,
|
|
352
|
+
p: float = 1.0,
|
|
353
|
+
learn: bool = False,
|
|
354
|
+
channels: int = 1,
|
|
355
|
+
clamp_min: Optional[float] = 1e-4,
|
|
356
|
+
clamp_max: Optional[float] = 100.0,
|
|
357
|
+
**kwargs,
|
|
358
|
+
):
|
|
359
|
+
super().__init__(**kwargs)
|
|
360
|
+
|
|
361
|
+
if not learn and channels != 1:
|
|
362
|
+
raise ValueError(
|
|
363
|
+
f"Cannot set 'channels' greater than '1' in case "
|
|
364
|
+
f"'{self.__class__.__name__}' is not trainable"
|
|
365
|
+
)
|
|
366
|
+
|
|
367
|
+
self._init_p = p
|
|
368
|
+
self.learn = learn
|
|
369
|
+
self.channels = channels
|
|
370
|
+
self.min_value = clamp_min if clamp_min is not None else 1e-4
|
|
371
|
+
self.max_value = clamp_max if clamp_max is not None else 100.0
|
|
372
|
+
|
|
373
|
+
if learn:
|
|
374
|
+
self.p = self.add_weight(
|
|
375
|
+
shape=(channels,),
|
|
376
|
+
initializer=initializers.Constant(p),
|
|
377
|
+
trainable=True,
|
|
378
|
+
name="p",
|
|
379
|
+
)
|
|
380
|
+
else:
|
|
381
|
+
self.p = p
|
|
382
|
+
|
|
383
|
+
def reset_parameters(self):
|
|
384
|
+
if self.learn:
|
|
385
|
+
self.p.assign(full((self.channels,), self._init_p, dtype=self.p.dtype))
|
|
386
|
+
|
|
387
|
+
def call(
|
|
388
|
+
self,
|
|
389
|
+
x,
|
|
390
|
+
index: Optional[any] = None,
|
|
391
|
+
ptr: Optional[any] = None,
|
|
392
|
+
dim_size: Optional[int] = None,
|
|
393
|
+
dim: int = -2,
|
|
394
|
+
**kwargs,
|
|
395
|
+
):
|
|
396
|
+
p = self.p
|
|
397
|
+
if self.channels != 1:
|
|
398
|
+
self.assert_two_dimensional_input(x, dim)
|
|
399
|
+
p = ops.reshape(p, (-1, self.channels))
|
|
400
|
+
|
|
401
|
+
if self.learn or p != 1.0:
|
|
402
|
+
x = ops.power(ops.clip(x, self.min_value, self.max_value), p)
|
|
403
|
+
|
|
404
|
+
out = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
|
|
405
|
+
|
|
406
|
+
if self.learn or p != 1.0:
|
|
407
|
+
out = ops.power(ops.clip(out, self.min_value, self.max_value), 1.0 / p)
|
|
408
|
+
|
|
409
|
+
return out
|
|
410
|
+
|
|
411
|
+
def __repr__(self) -> str:
|
|
412
|
+
return f"{self.__class__.__name__}(learn={self.learn})"
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from .base import Aggregation
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class DeepSetsAggregation(Aggregation):
|
|
6
|
+
r"""Performs Deep Sets aggregation in which the elements to aggregate are
|
|
7
|
+
first transformed by a Multi-Layer Perceptron (MLP)
|
|
8
|
+
:math:`\phi_{\mathbf{\Theta}}`, summed, and then transformed by another MLP
|
|
9
|
+
:math:`\rho_{\mathbf{\Theta}}`.
|
|
10
|
+
|
|
11
|
+
Example:
|
|
12
|
+
```python
|
|
13
|
+
import numpy as np
|
|
14
|
+
import keras
|
|
15
|
+
from k3_node.layers import DeepSetsAggregation
|
|
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 = DeepSetsAggregation(local_nn=keras.layers.Dense(16), global_nn=keras.layers.Dense(16))
|
|
21
|
+
out = aggr(x, index=index, dim_size=2)
|
|
22
|
+
print(tuple(out.shape)) # (2, 16)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
local_nn: Optional[any] = None,
|
|
29
|
+
global_nn: Optional[any] = None,
|
|
30
|
+
local_mlp: Optional[any] = None,
|
|
31
|
+
global_mlp: Optional[any] = None,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
self.local_nn = local_nn if local_nn is not None else local_mlp
|
|
36
|
+
self.global_nn = global_nn if global_nn is not None else global_mlp
|
|
37
|
+
|
|
38
|
+
def reset_parameters(self):
|
|
39
|
+
if self.local_nn is not None and hasattr(self.local_nn, "reset_parameters"):
|
|
40
|
+
self.local_nn.reset_parameters()
|
|
41
|
+
if self.global_nn is not None and hasattr(self.global_nn, "reset_parameters"):
|
|
42
|
+
self.global_nn.reset_parameters()
|
|
43
|
+
|
|
44
|
+
def call(
|
|
45
|
+
self,
|
|
46
|
+
x,
|
|
47
|
+
index: Optional[any] = None,
|
|
48
|
+
ptr: Optional[any] = None,
|
|
49
|
+
dim_size: Optional[int] = None,
|
|
50
|
+
dim: int = -2,
|
|
51
|
+
**kwargs,
|
|
52
|
+
):
|
|
53
|
+
if self.local_nn is not None:
|
|
54
|
+
x = self.local_nn(x)
|
|
55
|
+
|
|
56
|
+
x = self.reduce(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, reduce="sum")
|
|
57
|
+
|
|
58
|
+
if self.global_nn is not None:
|
|
59
|
+
x = self.global_nn(x)
|
|
60
|
+
|
|
61
|
+
return x
|
|
62
|
+
|
|
63
|
+
def __repr__(self) -> str:
|
|
64
|
+
return f"{self.__class__.__name__}(local_nn={self.local_nn}, global_nn={self.global_nn})"
|
|
65
|
+
|
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.aggr import Aggregation
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class DeepSetsAggregation(Aggregation):
|
|
8
|
+
def __init__(
|
|
9
|
+
self,
|
|
10
|
+
local_mlp=None,
|
|
11
|
+
global_mlp=None,
|
|
12
|
+
):
|
|
13
|
+
super().__init__()
|
|
14
|
+
|
|
15
|
+
self.local_mlp = local_mlp
|
|
16
|
+
self.global_mlp = global_mlp
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def call(self, x, index=None, axis=-2):
|
|
20
|
+
|
|
21
|
+
if self.local_mlp is not None:
|
|
22
|
+
x = self.local_mlp(x)
|
|
23
|
+
|
|
24
|
+
x = self.reduce(x, index=index, axis=axis, reduce_fn=ops.segment_sum)
|
|
25
|
+
|
|
26
|
+
if self.global_mlp is not None:
|
|
27
|
+
x = self.global_mlp(x) # Assuming batch handling within MLP
|
|
28
|
+
|
|
29
|
+
return x
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
from typing import List, Optional
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class ResNetPotential(layers.Layer):
|
|
8
|
+
def __init__(self, in_channels: int, out_channels: int, num_layers: List[int], **kwargs):
|
|
9
|
+
super().__init__(**kwargs)
|
|
10
|
+
sizes = [in_channels] + num_layers + [out_channels]
|
|
11
|
+
self.layers_list = [
|
|
12
|
+
layers.Dense(out_size, activation="tanh")
|
|
13
|
+
for in_size, out_size in zip(sizes[:-2], sizes[1:-1])
|
|
14
|
+
]
|
|
15
|
+
self.final_layer = layers.Dense(sizes[-1])
|
|
16
|
+
self.res_trans = [
|
|
17
|
+
layers.Dense(layer_size)
|
|
18
|
+
for layer_size in num_layers + [out_channels]
|
|
19
|
+
]
|
|
20
|
+
|
|
21
|
+
def reset_parameters(self):
|
|
22
|
+
for l in self.layers_list:
|
|
23
|
+
l.reset_parameters()
|
|
24
|
+
self.final_layer.reset_parameters()
|
|
25
|
+
for r in self.res_trans:
|
|
26
|
+
r.reset_parameters()
|
|
27
|
+
|
|
28
|
+
def call(self, inp):
|
|
29
|
+
h = inp
|
|
30
|
+
for layer, res in zip(self.layers_list + [self.final_layer], self.res_trans):
|
|
31
|
+
h_next = layer(h)
|
|
32
|
+
h = res(inp) + h_next
|
|
33
|
+
return h
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class EquilibriumAggregation(Aggregation):
|
|
37
|
+
r"""The equilibrium aggregation layer from the `"Equilibrium Aggregation:
|
|
38
|
+
Encoding Sets via Optimization" <https://arxiv.org/abs/2202.12795>`_ paper.
|
|
39
|
+
|
|
40
|
+
Example:
|
|
41
|
+
```python
|
|
42
|
+
import numpy as np
|
|
43
|
+
from k3_node.layers import EquilibriumAggregation
|
|
44
|
+
|
|
45
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
46
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
47
|
+
|
|
48
|
+
aggr = EquilibriumAggregation(in_channels=8, out_channels=16, num_layers=[8])
|
|
49
|
+
out = aggr(x, index=index, dim_size=2)
|
|
50
|
+
print(tuple(out.shape)) # (2, 16)
|
|
51
|
+
```
|
|
52
|
+
"""
|
|
53
|
+
|
|
54
|
+
def __init__(
|
|
55
|
+
self,
|
|
56
|
+
in_channels: int,
|
|
57
|
+
out_channels: int,
|
|
58
|
+
num_layers: List[int],
|
|
59
|
+
grad_iter: int = 5,
|
|
60
|
+
lamb: float = 0.1,
|
|
61
|
+
**kwargs,
|
|
62
|
+
):
|
|
63
|
+
super().__init__(**kwargs)
|
|
64
|
+
self.in_channels = in_channels
|
|
65
|
+
self.out_channels = out_channels
|
|
66
|
+
self.num_layers = num_layers
|
|
67
|
+
self.grad_iter = grad_iter
|
|
68
|
+
self.initial_lamb = lamb
|
|
69
|
+
|
|
70
|
+
self.potential = ResNetPotential(in_channels + out_channels, out_channels, num_layers)
|
|
71
|
+
self.proj = layers.Dense(out_channels)
|
|
72
|
+
|
|
73
|
+
def reset_parameters(self):
|
|
74
|
+
self.potential.reset_parameters()
|
|
75
|
+
self.proj.reset_parameters()
|
|
76
|
+
|
|
77
|
+
def call(
|
|
78
|
+
self,
|
|
79
|
+
x,
|
|
80
|
+
index: Optional[any] = None,
|
|
81
|
+
ptr: Optional[any] = None,
|
|
82
|
+
dim_size: Optional[int] = None,
|
|
83
|
+
dim: int = -2,
|
|
84
|
+
**kwargs,
|
|
85
|
+
):
|
|
86
|
+
self.assert_index_present(index)
|
|
87
|
+
index = ops.cast(index, dtype="int32")
|
|
88
|
+
if dim_size is None:
|
|
89
|
+
dim_size = ops.max(index) + 1
|
|
90
|
+
|
|
91
|
+
# Initial mean aggregation as starting state
|
|
92
|
+
x_mean = self.reduce(x, index, ptr, dim_size, dim, reduce="mean")
|
|
93
|
+
y = self.proj(x_mean)
|
|
94
|
+
|
|
95
|
+
# Unrolled iterative equilibrium updates
|
|
96
|
+
for _ in range(self.grad_iter):
|
|
97
|
+
y_expanded = ops.take(y, index, axis=0)
|
|
98
|
+
inp = ops.concatenate([x, y_expanded], axis=-1)
|
|
99
|
+
pot = self.potential(inp)
|
|
100
|
+
pot_mean = self.reduce(pot, index, ptr, dim_size, dim, reduce="mean")
|
|
101
|
+
y = y + 0.1 * pot_mean - self.initial_lamb * y
|
|
102
|
+
|
|
103
|
+
return y
|
|
104
|
+
|
|
105
|
+
def __repr__(self) -> str:
|
|
106
|
+
return f"{self.__class__.__name__}()"
|
|
107
|
+
|
|
@@ -0,0 +1,43 @@
|
|
|
1
|
+
from typing import List, Union
|
|
2
|
+
from .base import Aggregation
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class FusedAggregation(Aggregation):
|
|
6
|
+
r"""Helper class to fuse computation of multiple aggregations together.
|
|
7
|
+
|
|
8
|
+
Example:
|
|
9
|
+
```python
|
|
10
|
+
import numpy as np
|
|
11
|
+
from k3_node.layers import FusedAggregation
|
|
12
|
+
|
|
13
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
14
|
+
index = np.repeat([0, 1], 5) # aggregate nodes 0-4 into set 0 and nodes 5-9 into set 1
|
|
15
|
+
|
|
16
|
+
aggr = FusedAggregation(aggrs=["sum", "mean", "max"])
|
|
17
|
+
outs = aggr(x, index=index, dim_size=2) # one result per aggregation, computed together
|
|
18
|
+
print(len(outs), tuple(outs[0].shape)) # 3 (2, 8)
|
|
19
|
+
```
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
def __init__(self, aggrs: List[Union[Aggregation, str]], **kwargs):
|
|
23
|
+
super().__init__(**kwargs)
|
|
24
|
+
from .resolver import aggregation_resolver
|
|
25
|
+
|
|
26
|
+
self.aggrs = [aggregation_resolver(aggr) for aggr in aggrs]
|
|
27
|
+
|
|
28
|
+
def reset_parameters(self):
|
|
29
|
+
for aggr in self.aggrs:
|
|
30
|
+
if hasattr(aggr, "reset_parameters"):
|
|
31
|
+
aggr.reset_parameters()
|
|
32
|
+
|
|
33
|
+
def call(
|
|
34
|
+
self,
|
|
35
|
+
x,
|
|
36
|
+
index=None,
|
|
37
|
+
ptr=None,
|
|
38
|
+
dim_size=None,
|
|
39
|
+
dim=-2,
|
|
40
|
+
**kwargs,
|
|
41
|
+
):
|
|
42
|
+
return [aggr(x, index=index, ptr=ptr, dim_size=dim_size, dim=dim, **kwargs) for aggr in self.aggrs]
|
|
43
|
+
|