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
k3_node/io/npz.py
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
from typing import Any, Dict
|
|
2
|
+
import numpy as np
|
|
3
|
+
import scipy.sparse as sp
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.data import Data
|
|
7
|
+
from k3_node.layers.conv.utils import remove_self_loops
|
|
8
|
+
from k3_node.transforms.utils import to_undirected as to_undirected_fn
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def read_npz(path: str, to_undirected: bool = True) -> Data:
|
|
12
|
+
"""Reads a `.npz` graph archive (e.g. Amazon, Coauthor) and returns a `Data` object."""
|
|
13
|
+
with np.load(path) as f:
|
|
14
|
+
return parse_npz(f, to_undirected=to_undirected)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def parse_npz(f: Dict[str, Any], to_undirected: bool = True) -> Data:
|
|
18
|
+
x_sp = sp.csr_matrix(
|
|
19
|
+
(f["attr_data"], f["attr_indices"], f["attr_indptr"]),
|
|
20
|
+
shape=tuple(f["attr_shape"]),
|
|
21
|
+
)
|
|
22
|
+
x = np.array(x_sp.todense(), dtype=np.float32)
|
|
23
|
+
x[x > 0] = 1.0
|
|
24
|
+
|
|
25
|
+
adj = sp.csr_matrix(
|
|
26
|
+
(f["adj_data"], f["adj_indices"], f["adj_indptr"]),
|
|
27
|
+
shape=tuple(f["adj_shape"]),
|
|
28
|
+
).tocoo()
|
|
29
|
+
|
|
30
|
+
row = np.array(adj.row, dtype=np.int64)
|
|
31
|
+
col = np.array(adj.col, dtype=np.int64)
|
|
32
|
+
edge_index = np.stack([row, col], axis=0)
|
|
33
|
+
|
|
34
|
+
edge_index, _ = remove_self_loops(edge_index)
|
|
35
|
+
if to_undirected:
|
|
36
|
+
edge_index = to_undirected_fn(edge_index, num_nodes=x.shape[0])
|
|
37
|
+
|
|
38
|
+
y = np.array(f["labels"], dtype=np.int64)
|
|
39
|
+
|
|
40
|
+
return Data(
|
|
41
|
+
x=ops.convert_to_tensor(x, dtype="float32"),
|
|
42
|
+
edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
|
|
43
|
+
y=ops.convert_to_tensor(y, dtype="int64"),
|
|
44
|
+
)
|
|
45
|
+
|
k3_node/io/off.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
from typing import List
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def parse_off(src: List[str]) -> Data:
|
|
9
|
+
r"""Parses the lines of an OFF (Object File Format) mesh into ``Data(pos, face)``;
|
|
10
|
+
quadrilaterals are split into two triangles."""
|
|
11
|
+
if src[0] == 'OFF':
|
|
12
|
+
src = src[1:]
|
|
13
|
+
else: # some files lack the line break after "OFF"
|
|
14
|
+
src[0] = src[0][3:]
|
|
15
|
+
num_nodes, num_faces = (int(item) for item in src[0].split()[:2])
|
|
16
|
+
pos = np.array([[float(v) for v in line.split()[:3]] for line in src[1:1 + num_nodes]], dtype=np.float32)
|
|
17
|
+
faces = [[int(v) for v in line.strip().split()] for line in src[1 + num_nodes:1 + num_nodes + num_faces]]
|
|
18
|
+
tri = [f[1:4] for f in faces if f[0] == 3]
|
|
19
|
+
for f in faces:
|
|
20
|
+
if f[0] == 4:
|
|
21
|
+
tri += [[f[1], f[2], f[3]], [f[1], f[3], f[4]]]
|
|
22
|
+
face = np.array(tri, dtype=np.int64).reshape(-1, 3).T
|
|
23
|
+
return Data(pos=pos, face=face)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def read_off(path: str) -> Data:
|
|
27
|
+
r"""Reads an OFF (Object File Format) mesh file into ``Data(pos, face)``."""
|
|
28
|
+
with open(path) as f:
|
|
29
|
+
return parse_off(f.read().split('\n')[:-1])
|
k3_node/io/planetoid.py
ADDED
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
import os.path as osp
|
|
2
|
+
import pickle
|
|
3
|
+
import warnings
|
|
4
|
+
from typing import Dict, List, Optional
|
|
5
|
+
import numpy as np
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.data import Data
|
|
9
|
+
from k3_node.io.txt_array import read_txt_array
|
|
10
|
+
from k3_node.layers.conv.utils import remove_self_loops
|
|
11
|
+
from k3_node.utils.graph import coalesce
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def index_to_mask(index, size: int):
|
|
15
|
+
"""Converts 1D index array into a boolean mask."""
|
|
16
|
+
mask = np.zeros(size, dtype=bool)
|
|
17
|
+
idx_np = ops.convert_to_numpy(index)
|
|
18
|
+
mask[idx_np] = True
|
|
19
|
+
return ops.convert_to_tensor(mask, dtype="bool")
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def edge_index_from_dict(graph_dict: Dict[int, List[int]], num_nodes: Optional[int] = None):
|
|
23
|
+
rows: List[int] = []
|
|
24
|
+
cols: List[int] = []
|
|
25
|
+
for key, val_list in graph_dict.items():
|
|
26
|
+
for val in val_list:
|
|
27
|
+
rows.append(key)
|
|
28
|
+
cols.append(val)
|
|
29
|
+
if len(rows) == 0:
|
|
30
|
+
return ops.zeros((2, 0), dtype="int64")
|
|
31
|
+
|
|
32
|
+
edge_index = np.array([rows, cols], dtype=np.int64)
|
|
33
|
+
edge_index, _ = remove_self_loops(edge_index)
|
|
34
|
+
edge_index, _ = coalesce(edge_index, num_nodes=num_nodes, sort_by_row=False)
|
|
35
|
+
return ops.convert_to_tensor(edge_index, dtype="int64")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def read_file(folder: str, prefix: str, name: str):
|
|
39
|
+
path = osp.join(folder, f"ind.{prefix.lower()}.{name}")
|
|
40
|
+
if name == "test.index":
|
|
41
|
+
return read_txt_array(path, dtype="int64")
|
|
42
|
+
|
|
43
|
+
with open(path, "rb") as f:
|
|
44
|
+
warnings.filterwarnings("ignore", ".*`scipy.sparse.csr` name.*")
|
|
45
|
+
out = pickle.load(f, encoding="latin1")
|
|
46
|
+
|
|
47
|
+
if name == "graph":
|
|
48
|
+
return out
|
|
49
|
+
|
|
50
|
+
if hasattr(out, "todense"):
|
|
51
|
+
out = out.todense()
|
|
52
|
+
return np.array(out, dtype=np.float32)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def read_planetoid_data(folder: str, prefix: str) -> Data:
|
|
56
|
+
"""Reads planetoid citation graph files and returns a `Data` object."""
|
|
57
|
+
names = ["x", "tx", "allx", "y", "ty", "ally", "graph", "test.index"]
|
|
58
|
+
items = [read_file(folder, prefix, name) for name in names]
|
|
59
|
+
x, tx, allx, y, ty, ally, graph, test_index = items
|
|
60
|
+
|
|
61
|
+
test_index_np = ops.convert_to_numpy(test_index).astype(np.int64)
|
|
62
|
+
sorted_test_index = np.sort(test_index_np)
|
|
63
|
+
|
|
64
|
+
train_index = np.arange(y.shape[0], dtype=np.int64)
|
|
65
|
+
val_index = np.arange(y.shape[0], y.shape[0] + 500, dtype=np.int64)
|
|
66
|
+
|
|
67
|
+
if prefix.lower() == "citeseer":
|
|
68
|
+
len_test_indices = int(np.max(test_index_np) - np.min(test_index_np)) + 1
|
|
69
|
+
tx_ext = np.zeros((len_test_indices, tx.shape[1]), dtype=tx.dtype)
|
|
70
|
+
tx_ext[sorted_test_index - np.min(test_index_np), :] = tx
|
|
71
|
+
ty_ext = np.zeros((len_test_indices, ty.shape[1]), dtype=ty.dtype)
|
|
72
|
+
ty_ext[sorted_test_index - np.min(test_index_np), :] = ty
|
|
73
|
+
tx, ty = tx_ext, ty_ext
|
|
74
|
+
|
|
75
|
+
x = np.concatenate([allx, tx], axis=0)
|
|
76
|
+
x[test_index_np] = x[sorted_test_index]
|
|
77
|
+
|
|
78
|
+
y_cat = np.concatenate([ally, ty], axis=0)
|
|
79
|
+
y = np.argmax(y_cat, axis=1).astype(np.int64)
|
|
80
|
+
y[test_index_np] = y[sorted_test_index]
|
|
81
|
+
|
|
82
|
+
num_nodes = y.shape[0]
|
|
83
|
+
train_mask = index_to_mask(train_index, size=num_nodes)
|
|
84
|
+
val_mask = index_to_mask(val_index, size=num_nodes)
|
|
85
|
+
test_mask = index_to_mask(test_index_np, size=num_nodes)
|
|
86
|
+
|
|
87
|
+
edge_index = edge_index_from_dict(graph, num_nodes=num_nodes)
|
|
88
|
+
|
|
89
|
+
data = Data(
|
|
90
|
+
x=ops.convert_to_tensor(x, dtype="float32"),
|
|
91
|
+
edge_index=edge_index,
|
|
92
|
+
y=ops.convert_to_tensor(y, dtype="int64"),
|
|
93
|
+
)
|
|
94
|
+
data.train_mask = train_mask
|
|
95
|
+
data.val_mask = val_mask
|
|
96
|
+
data.test_mask = test_mask
|
|
97
|
+
|
|
98
|
+
return data
|
k3_node/io/tu.py
ADDED
|
@@ -0,0 +1,137 @@
|
|
|
1
|
+
import glob
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Any, Dict, List, Optional, Tuple
|
|
4
|
+
import numpy as np
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.data import Data
|
|
8
|
+
from k3_node.io.txt_array import read_txt_array
|
|
9
|
+
from k3_node.layers.conv.utils import remove_self_loops
|
|
10
|
+
from k3_node.utils.graph import coalesce
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def np_one_hot(arr: np.ndarray) -> np.ndarray:
|
|
14
|
+
"""One-hot encodes a 1D integer array."""
|
|
15
|
+
arr = arr.astype(np.int64)
|
|
16
|
+
if arr.size == 0:
|
|
17
|
+
return np.zeros((0, 0), dtype=np.float32)
|
|
18
|
+
min_val = np.min(arr)
|
|
19
|
+
arr = arr - min_val
|
|
20
|
+
max_val = np.max(arr)
|
|
21
|
+
out = np.zeros((arr.shape[0], max_val + 1), dtype=np.float32)
|
|
22
|
+
out[np.arange(arr.shape[0]), arr] = 1.0
|
|
23
|
+
return out
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def read_tu_data(folder: str, prefix: str) -> Tuple[Data, Dict[str, Any], Dict[str, int]]:
|
|
27
|
+
"""Reads a TU benchmark dataset from disk and returns `(data, slices, sizes)`."""
|
|
28
|
+
files = sorted(glob.glob(osp.join(folder, f"{prefix}_*.txt")))
|
|
29
|
+
names = [osp.basename(f)[len(prefix) + 1 : -4] for f in files]
|
|
30
|
+
|
|
31
|
+
def read_file(name: str, dtype: str = "float32"):
|
|
32
|
+
path = osp.join(folder, f"{prefix}_{name}.txt")
|
|
33
|
+
return ops.convert_to_numpy(read_txt_array(path, sep=",", dtype=dtype))
|
|
34
|
+
|
|
35
|
+
# Adjacency: 1-indexed in raw TU files
|
|
36
|
+
edge_index_np = read_file("A", dtype="int64")
|
|
37
|
+
if edge_index_np.ndim == 1:
|
|
38
|
+
edge_index_np = edge_index_np.reshape(-1, 2)
|
|
39
|
+
edge_index_np = edge_index_np.T - 1 # 0-indexed (2, E)
|
|
40
|
+
|
|
41
|
+
# Graph indicator: which graph each node belongs to (1-indexed)
|
|
42
|
+
batch_np = read_file("graph_indicator", dtype="int64") - 1 # 0-indexed (N,)
|
|
43
|
+
|
|
44
|
+
node_attribute = np.empty((batch_np.shape[0], 0), dtype=np.float32)
|
|
45
|
+
if "node_attributes" in names:
|
|
46
|
+
node_attribute = read_file("node_attributes", dtype="float32")
|
|
47
|
+
if node_attribute.ndim == 1:
|
|
48
|
+
node_attribute = node_attribute[:, None]
|
|
49
|
+
|
|
50
|
+
node_label = np.empty((batch_np.shape[0], 0), dtype=np.float32)
|
|
51
|
+
if "node_labels" in names:
|
|
52
|
+
node_label_raw = read_file("node_labels", dtype="int64")
|
|
53
|
+
if node_label_raw.ndim == 1:
|
|
54
|
+
node_label_raw = node_label_raw[:, None]
|
|
55
|
+
encoded = [np_one_hot(node_label_raw[:, i]) for i in range(node_label_raw.shape[1])]
|
|
56
|
+
node_label = np.concatenate(encoded, axis=-1)
|
|
57
|
+
|
|
58
|
+
edge_attribute = np.empty((edge_index_np.shape[1], 0), dtype=np.float32)
|
|
59
|
+
if "edge_attributes" in names:
|
|
60
|
+
edge_attribute = read_file("edge_attributes", dtype="float32")
|
|
61
|
+
if edge_attribute.ndim == 1:
|
|
62
|
+
edge_attribute = edge_attribute[:, None]
|
|
63
|
+
|
|
64
|
+
edge_label = np.empty((edge_index_np.shape[1], 0), dtype=np.float32)
|
|
65
|
+
if "edge_labels" in names:
|
|
66
|
+
edge_label_raw = read_file("edge_labels", dtype="int64")
|
|
67
|
+
if edge_label_raw.ndim == 1:
|
|
68
|
+
edge_label_raw = edge_label_raw[:, None]
|
|
69
|
+
encoded = [np_one_hot(edge_label_raw[:, i]) for i in range(edge_label_raw.shape[1])]
|
|
70
|
+
edge_label = np.concatenate(encoded, axis=-1)
|
|
71
|
+
|
|
72
|
+
# Combine attributes and one-hot labels
|
|
73
|
+
x_list = [arr for arr in [node_attribute, node_label] if arr.shape[1] > 0]
|
|
74
|
+
x_np = np.concatenate(x_list, axis=-1) if len(x_list) > 0 else None
|
|
75
|
+
|
|
76
|
+
edge_attr_list = [arr for arr in [edge_attribute, edge_label] if arr.shape[1] > 0]
|
|
77
|
+
edge_attr_np = np.concatenate(edge_attr_list, axis=-1) if len(edge_attr_list) > 0 else None
|
|
78
|
+
|
|
79
|
+
y_np = None
|
|
80
|
+
if "graph_attributes" in names:
|
|
81
|
+
y_np = read_file("graph_attributes", dtype="float32")
|
|
82
|
+
elif "graph_labels" in names:
|
|
83
|
+
y_raw = read_file("graph_labels", dtype="int64")
|
|
84
|
+
_, y_inv = np.unique(y_raw, return_inverse=True)
|
|
85
|
+
y_np = y_inv.astype(np.int64)
|
|
86
|
+
|
|
87
|
+
num_nodes = x_np.shape[0] if x_np is not None else int(np.max(edge_index_np)) + 1
|
|
88
|
+
edge_index_np, edge_attr_np = remove_self_loops(edge_index_np, edge_attr_np)
|
|
89
|
+
edge_index_t, edge_attr_t = coalesce(edge_index_np, edge_attr_np, num_nodes=num_nodes)
|
|
90
|
+
edge_index_np = ops.convert_to_numpy(edge_index_t)
|
|
91
|
+
edge_attr_np = ops.convert_to_numpy(edge_attr_t) if edge_attr_t is not None else None
|
|
92
|
+
|
|
93
|
+
# Convert to tensors
|
|
94
|
+
edge_index = ops.convert_to_tensor(edge_index_np, dtype="int64")
|
|
95
|
+
x = ops.convert_to_tensor(x_np, dtype="float32") if x_np is not None else None
|
|
96
|
+
edge_attr = ops.convert_to_tensor(edge_attr_np, dtype="float32") if edge_attr_np is not None else None
|
|
97
|
+
y = ops.convert_to_tensor(y_np, dtype="float32" if "graph_attributes" in names else "int64") if y_np is not None else None
|
|
98
|
+
|
|
99
|
+
# Compute graph slices
|
|
100
|
+
num_graphs = int(np.max(batch_np)) + 1
|
|
101
|
+
node_counts = np.bincount(batch_np, minlength=num_graphs)
|
|
102
|
+
node_slice = np.pad(np.cumsum(node_counts), (1, 0))
|
|
103
|
+
|
|
104
|
+
row = edge_index_np[0]
|
|
105
|
+
edge_batch = batch_np[row]
|
|
106
|
+
edge_counts = np.bincount(edge_batch, minlength=num_graphs)
|
|
107
|
+
edge_slice = np.pad(np.cumsum(edge_counts), (1, 0))
|
|
108
|
+
|
|
109
|
+
# Shift edge indices so each graph starts at 0
|
|
110
|
+
shift = node_slice[edge_batch]
|
|
111
|
+
edge_index_shifted = edge_index_np - shift[None, :]
|
|
112
|
+
edge_index = ops.convert_to_tensor(edge_index_shifted, dtype="int64")
|
|
113
|
+
|
|
114
|
+
data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y)
|
|
115
|
+
slices: Dict[str, Any] = {
|
|
116
|
+
"edge_index": ops.convert_to_tensor(edge_slice, dtype="int64"),
|
|
117
|
+
}
|
|
118
|
+
if x is not None:
|
|
119
|
+
slices["x"] = ops.convert_to_tensor(node_slice, dtype="int64")
|
|
120
|
+
else:
|
|
121
|
+
data.num_nodes = int(batch_np.shape[0])
|
|
122
|
+
if edge_attr is not None:
|
|
123
|
+
slices["edge_attr"] = ops.convert_to_tensor(edge_slice, dtype="int64")
|
|
124
|
+
if y is not None:
|
|
125
|
+
if y_np.shape[0] == batch_np.shape[0]:
|
|
126
|
+
slices["y"] = ops.convert_to_tensor(node_slice, dtype="int64")
|
|
127
|
+
else:
|
|
128
|
+
slices["y"] = ops.convert_to_tensor(np.arange(num_graphs + 1, dtype=np.int64), dtype="int64")
|
|
129
|
+
|
|
130
|
+
sizes = {
|
|
131
|
+
"num_node_attributes": node_attribute.shape[-1],
|
|
132
|
+
"num_node_labels": node_label.shape[-1],
|
|
133
|
+
"num_edge_attributes": edge_attribute.shape[-1],
|
|
134
|
+
"num_edge_labels": edge_label.shape[-1],
|
|
135
|
+
}
|
|
136
|
+
|
|
137
|
+
return data, slices, sizes
|
k3_node/io/txt_array.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from typing import List, Optional, Union
|
|
2
|
+
import numpy as np
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def parse_txt_array(
|
|
7
|
+
src: List[str],
|
|
8
|
+
sep: Optional[str] = None,
|
|
9
|
+
start: int = 0,
|
|
10
|
+
end: Optional[int] = None,
|
|
11
|
+
dtype: Optional[str] = None,
|
|
12
|
+
):
|
|
13
|
+
"""Parses a list of string rows into a tensor."""
|
|
14
|
+
lines = [line.strip() for line in src if line.strip()]
|
|
15
|
+
if len(lines) == 0:
|
|
16
|
+
return ops.zeros((0,), dtype=dtype or "float32")
|
|
17
|
+
|
|
18
|
+
split_lines = []
|
|
19
|
+
is_float = False
|
|
20
|
+
for line in lines:
|
|
21
|
+
parts = line.split(sep)[start:end]
|
|
22
|
+
row = []
|
|
23
|
+
for x in parts:
|
|
24
|
+
x_str = x.strip()
|
|
25
|
+
if not x_str:
|
|
26
|
+
continue
|
|
27
|
+
if "." in x_str or "e" in x_str.lower():
|
|
28
|
+
is_float = True
|
|
29
|
+
row.append(float(x_str))
|
|
30
|
+
else:
|
|
31
|
+
try:
|
|
32
|
+
row.append(int(x_str))
|
|
33
|
+
except ValueError:
|
|
34
|
+
is_float = True
|
|
35
|
+
row.append(float(x_str))
|
|
36
|
+
split_lines.append(row)
|
|
37
|
+
|
|
38
|
+
if dtype is None:
|
|
39
|
+
dtype = "float32" if is_float else "int64"
|
|
40
|
+
|
|
41
|
+
arr = np.array(split_lines, dtype=dtype)
|
|
42
|
+
if arr.ndim > 1 and arr.shape[1] == 1:
|
|
43
|
+
arr = np.squeeze(arr, axis=1)
|
|
44
|
+
return ops.convert_to_tensor(arr, dtype=dtype)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def read_txt_array(
|
|
48
|
+
path: str,
|
|
49
|
+
sep: Optional[str] = None,
|
|
50
|
+
start: int = 0,
|
|
51
|
+
end: Optional[int] = None,
|
|
52
|
+
dtype: Optional[str] = None,
|
|
53
|
+
):
|
|
54
|
+
"""Reads a text array from a file and returns a tensor."""
|
|
55
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
56
|
+
src = f.read().split("\n")
|
|
57
|
+
return parse_txt_array(src, sep, start, end, dtype)
|
|
58
|
+
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
"""
|
|
2
|
+
`k3_node.layers` module provides access to various layers for building
|
|
3
|
+
graph neural networks.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from .conv import *
|
|
7
|
+
from .norm import *
|
|
8
|
+
from .aggr import *
|
|
9
|
+
from .attention import *
|
|
10
|
+
from .dense import *
|
|
11
|
+
from .pool import *
|
|
12
|
+
from .unpool import *
|
|
13
|
+
from .kge import *
|
|
14
|
+
from .functional import *
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
r"""Aggregation operators for graph neural networks."""
|
|
2
|
+
|
|
3
|
+
from .base import Aggregation, from_dense_batch, ptr2index, to_dense_adj, to_dense_batch
|
|
4
|
+
from .basic import (
|
|
5
|
+
MaxAggregation,
|
|
6
|
+
MeanAggregation,
|
|
7
|
+
MinAggregation,
|
|
8
|
+
MulAggregation,
|
|
9
|
+
PowerMeanAggregation,
|
|
10
|
+
SoftmaxAggregation,
|
|
11
|
+
StdAggregation,
|
|
12
|
+
SumAggregation,
|
|
13
|
+
VarAggregation,
|
|
14
|
+
)
|
|
15
|
+
from .quantile import MedianAggregation, QuantileAggregation
|
|
16
|
+
from .attention import AttentionalAggregation
|
|
17
|
+
from .set2set import Set2Set
|
|
18
|
+
from .scaler import DegreeScalerAggregation
|
|
19
|
+
from .sort import SortAggregation
|
|
20
|
+
from .multi import MultiAggregation
|
|
21
|
+
from .deep_sets import DeepSetsAggregation
|
|
22
|
+
from .mlp import MLPAggregation
|
|
23
|
+
from .lstm import LSTMAggregation
|
|
24
|
+
from .gru import GRUAggregation
|
|
25
|
+
from .set_transformer import SetTransformerAggregation
|
|
26
|
+
from .gmt import GraphMultisetTransformer
|
|
27
|
+
from .variance_preserving import VariancePreservingAggregation
|
|
28
|
+
from .patch_transformer import PatchTransformerAggregation
|
|
29
|
+
from .lcm import LCMAggregation
|
|
30
|
+
from .equilibrium import EquilibriumAggregation
|
|
31
|
+
from .fused import FusedAggregation
|
|
32
|
+
from .resolver import aggregation_resolver
|
|
33
|
+
|
|
34
|
+
__all__ = [
|
|
35
|
+
"Aggregation",
|
|
36
|
+
"ptr2index",
|
|
37
|
+
"to_dense_batch",
|
|
38
|
+
"from_dense_batch",
|
|
39
|
+
"to_dense_adj",
|
|
40
|
+
"SumAggregation",
|
|
41
|
+
"MeanAggregation",
|
|
42
|
+
"MaxAggregation",
|
|
43
|
+
"MinAggregation",
|
|
44
|
+
"MulAggregation",
|
|
45
|
+
"VarAggregation",
|
|
46
|
+
"StdAggregation",
|
|
47
|
+
"SoftmaxAggregation",
|
|
48
|
+
"PowerMeanAggregation",
|
|
49
|
+
"QuantileAggregation",
|
|
50
|
+
"MedianAggregation",
|
|
51
|
+
"AttentionalAggregation",
|
|
52
|
+
"Set2Set",
|
|
53
|
+
"DegreeScalerAggregation",
|
|
54
|
+
"SortAggregation",
|
|
55
|
+
"MultiAggregation",
|
|
56
|
+
"DeepSetsAggregation",
|
|
57
|
+
"MLPAggregation",
|
|
58
|
+
"LSTMAggregation",
|
|
59
|
+
"GRUAggregation",
|
|
60
|
+
"SetTransformerAggregation",
|
|
61
|
+
"GraphMultisetTransformer",
|
|
62
|
+
"VariancePreservingAggregation",
|
|
63
|
+
"PatchTransformerAggregation",
|
|
64
|
+
"LCMAggregation",
|
|
65
|
+
"EquilibriumAggregation",
|
|
66
|
+
"FusedAggregation",
|
|
67
|
+
"aggregation_resolver",
|
|
68
|
+
]
|
|
69
|
+
|
|
70
|
+
classes = __all__
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from .base import Aggregation
|
|
5
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class AttentionalAggregation(Aggregation):
|
|
9
|
+
r"""The soft attention aggregation layer from the `"Graph Matching Networks
|
|
10
|
+
for Learning the Similarity of Graph Structured Objects"
|
|
11
|
+
<https://arxiv.org/abs/1904.12787>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
import keras
|
|
17
|
+
from k3_node.layers import AttentionalAggregation
|
|
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 = AttentionalAggregation(gate_nn=keras.layers.Dense(1), nn=keras.layers.Dense(16))
|
|
23
|
+
out = aggr(x, index=index, dim_size=2) # attention-weighted sum per set
|
|
24
|
+
print(tuple(out.shape)) # (2, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
gate_nn,
|
|
31
|
+
nn: Optional[any] = None,
|
|
32
|
+
**kwargs,
|
|
33
|
+
):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
self.gate_nn = gate_nn
|
|
36
|
+
self.nn = nn
|
|
37
|
+
|
|
38
|
+
def reset_parameters(self):
|
|
39
|
+
if hasattr(self.gate_nn, "reset_parameters"):
|
|
40
|
+
self.gate_nn.reset_parameters()
|
|
41
|
+
if self.nn is not None and hasattr(self.nn, "reset_parameters"):
|
|
42
|
+
self.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
|
+
gate = self.gate_nn(x)
|
|
54
|
+
if self.nn is not None:
|
|
55
|
+
x = self.nn(x)
|
|
56
|
+
|
|
57
|
+
if ptr is not None and index is None:
|
|
58
|
+
from .base import ptr2index
|
|
59
|
+
index = ptr2index(ptr)
|
|
60
|
+
|
|
61
|
+
index = ops.cast(index, dtype="int32")
|
|
62
|
+
if dim_size is None: # a tensor while tracing; don't test its truth value
|
|
63
|
+
dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
64
|
+
|
|
65
|
+
# Graph-wise softmax over groups
|
|
66
|
+
max_val = segment_max(gate, index, num_segments=dim_size)
|
|
67
|
+
max_exp = ops.take(max_val, index, axis=0)
|
|
68
|
+
exp_gate = ops.exp(gate - max_exp)
|
|
69
|
+
sum_exp = segment_sum(exp_gate, index, num_segments=dim_size)
|
|
70
|
+
sum_exp_exp = ops.take(sum_exp, index, axis=0)
|
|
71
|
+
alpha = exp_gate / ops.maximum(sum_exp_exp, 1e-12)
|
|
72
|
+
|
|
73
|
+
return self.reduce(alpha * x, index, ptr, dim_size, dim, reduce="sum")
|
|
74
|
+
|
|
75
|
+
def __repr__(self) -> str:
|
|
76
|
+
return f"{self.__class__.__name__}(gate_nn={self.gate_nn}, nn={self.nn})"
|
|
77
|
+
|