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,57 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
from k3_node.ops.segment import segment_sum
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class GraphSizeNorm(layers.Layer):
|
|
6
|
+
r"""Applies Graph Size Normalization over each individual graph in a batch
|
|
7
|
+
of node features:
|
|
8
|
+
|
|
9
|
+
.. math::
|
|
10
|
+
\mathbf{x}^{\prime}_i = \frac{\mathbf{x}_i}{\sqrt{|\mathcal{V}|}}
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import GraphSizeNorm
|
|
16
|
+
|
|
17
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
18
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
19
|
+
|
|
20
|
+
layer = GraphSizeNorm()
|
|
21
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
22
|
+
print(tuple(out.shape)) # (10, 8)
|
|
23
|
+
```
|
|
24
|
+
"""
|
|
25
|
+
def __init__(self, **kwargs):
|
|
26
|
+
super().__init__(**kwargs)
|
|
27
|
+
|
|
28
|
+
def call(self, x, batch=None, batch_size=None):
|
|
29
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
30
|
+
if len(x) == 2:
|
|
31
|
+
x, batch = x
|
|
32
|
+
elif len(x) == 3:
|
|
33
|
+
x, batch, batch_size = x
|
|
34
|
+
|
|
35
|
+
if batch is None:
|
|
36
|
+
num_nodes = ops.cast(ops.shape(x)[0], dtype=x.dtype)
|
|
37
|
+
return x * ops.power(num_nodes, -0.5)
|
|
38
|
+
|
|
39
|
+
if batch_size is not None and not isinstance(batch_size, int):
|
|
40
|
+
try:
|
|
41
|
+
batch_size = int(batch_size)
|
|
42
|
+
except Exception:
|
|
43
|
+
pass
|
|
44
|
+
elif batch_size is None:
|
|
45
|
+
batch_size = ops.cast(ops.max(batch), "int32") + 1
|
|
46
|
+
|
|
47
|
+
batch = ops.cast(batch, "int32")
|
|
48
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
49
|
+
deg = segment_sum(ones, batch, num_segments=batch_size)
|
|
50
|
+
inv_sqrt_deg = ops.power(deg, -0.5)
|
|
51
|
+
scale = ops.take(inv_sqrt_deg, batch, axis=0)
|
|
52
|
+
return x * scale
|
|
53
|
+
|
|
54
|
+
def compute_output_shape(self, input_shape):
|
|
55
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
56
|
+
return input_shape[0]
|
|
57
|
+
return input_shape
|
|
@@ -0,0 +1,163 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
from k3_node.ops.segment import segment_sum
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class InstanceNorm(layers.Layer):
|
|
6
|
+
r"""Applies instance normalization over each individual example in a batch
|
|
7
|
+
of node features as described in the `"Instance Normalization: The Missing
|
|
8
|
+
Ingredient for Fast Stylization" <https://arxiv.org/abs/1607.06450>`_
|
|
9
|
+
paper.
|
|
10
|
+
|
|
11
|
+
.. math::
|
|
12
|
+
\mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
|
|
13
|
+
\textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
|
|
14
|
+
\odot \gamma + \beta
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
in_channels (int): Size of each input sample.
|
|
18
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
19
|
+
stability. (default: :obj:`1e-5`)
|
|
20
|
+
momentum (float, optional): The value used for the running mean and
|
|
21
|
+
running variance computation. (default: :obj:`0.1`)
|
|
22
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
23
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
24
|
+
(default: :obj:`False`)
|
|
25
|
+
track_running_stats (bool, optional): If set to :obj:`True`, this
|
|
26
|
+
module tracks the running mean and variance, and when set to
|
|
27
|
+
:obj:`False`, this module does not track such statistics and always
|
|
28
|
+
uses instance statistics in both training and eval modes.
|
|
29
|
+
(default: :obj:`False`)
|
|
30
|
+
|
|
31
|
+
Example:
|
|
32
|
+
```python
|
|
33
|
+
import numpy as np
|
|
34
|
+
from k3_node.layers import InstanceNorm
|
|
35
|
+
|
|
36
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
37
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
38
|
+
|
|
39
|
+
layer = InstanceNorm(in_channels=8)
|
|
40
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
41
|
+
print(tuple(out.shape)) # (10, 8)
|
|
42
|
+
```
|
|
43
|
+
"""
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
in_channels: int,
|
|
47
|
+
eps: float = 1e-5,
|
|
48
|
+
momentum: float = 0.1,
|
|
49
|
+
affine: bool = False,
|
|
50
|
+
track_running_stats: bool = False,
|
|
51
|
+
**kwargs
|
|
52
|
+
):
|
|
53
|
+
super().__init__(**kwargs)
|
|
54
|
+
self.in_channels = in_channels
|
|
55
|
+
self.eps = eps
|
|
56
|
+
self.momentum = momentum
|
|
57
|
+
self.affine = affine
|
|
58
|
+
self.track_running_stats = track_running_stats
|
|
59
|
+
|
|
60
|
+
if affine:
|
|
61
|
+
self.weight = self.add_weight(
|
|
62
|
+
shape=(in_channels,),
|
|
63
|
+
initializer="ones",
|
|
64
|
+
trainable=True,
|
|
65
|
+
name="weight",
|
|
66
|
+
)
|
|
67
|
+
self.bias = self.add_weight(
|
|
68
|
+
shape=(in_channels,),
|
|
69
|
+
initializer="zeros",
|
|
70
|
+
trainable=True,
|
|
71
|
+
name="bias",
|
|
72
|
+
)
|
|
73
|
+
else:
|
|
74
|
+
self.weight = None
|
|
75
|
+
self.bias = None
|
|
76
|
+
|
|
77
|
+
if track_running_stats:
|
|
78
|
+
self.running_mean = self.add_weight(
|
|
79
|
+
shape=(in_channels,),
|
|
80
|
+
initializer="zeros",
|
|
81
|
+
trainable=False,
|
|
82
|
+
name="running_mean",
|
|
83
|
+
)
|
|
84
|
+
self.running_var = self.add_weight(
|
|
85
|
+
shape=(in_channels,),
|
|
86
|
+
initializer="ones",
|
|
87
|
+
trainable=False,
|
|
88
|
+
name="running_var",
|
|
89
|
+
)
|
|
90
|
+
else:
|
|
91
|
+
self.running_mean = None
|
|
92
|
+
self.running_var = None
|
|
93
|
+
|
|
94
|
+
def reset_running_stats(self):
|
|
95
|
+
if self.track_running_stats:
|
|
96
|
+
self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
|
|
97
|
+
self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
|
|
98
|
+
|
|
99
|
+
def reset_parameters(self):
|
|
100
|
+
self.reset_running_stats()
|
|
101
|
+
if self.affine:
|
|
102
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
103
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
104
|
+
|
|
105
|
+
def call(self, x, batch=None, batch_size=None, training=None):
|
|
106
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
107
|
+
if len(x) == 2:
|
|
108
|
+
x, batch = x
|
|
109
|
+
elif len(x) == 3:
|
|
110
|
+
x, batch, batch_size = x
|
|
111
|
+
|
|
112
|
+
# Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
|
|
113
|
+
is_training = bool(training) if training is not None else False
|
|
114
|
+
|
|
115
|
+
if batch is None:
|
|
116
|
+
batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
|
|
117
|
+
batch_size = 1
|
|
118
|
+
elif batch_size is not None and not isinstance(batch_size, int):
|
|
119
|
+
try:
|
|
120
|
+
batch_size = int(batch_size)
|
|
121
|
+
except Exception:
|
|
122
|
+
pass
|
|
123
|
+
elif batch_size is None:
|
|
124
|
+
batch_size = ops.cast(ops.max(batch), "int32") + 1
|
|
125
|
+
|
|
126
|
+
batch = ops.cast(batch, "int32")
|
|
127
|
+
|
|
128
|
+
if is_training or not self.track_running_stats:
|
|
129
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
130
|
+
counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
|
|
131
|
+
unbiased_counts = ops.maximum(counts - 1.0, 1.0)
|
|
132
|
+
|
|
133
|
+
mean = segment_sum(x, batch, num_segments=batch_size) / counts
|
|
134
|
+
x_c = x - ops.take(mean, batch, axis=0)
|
|
135
|
+
sq_diff = segment_sum(ops.power(x_c, 2), batch, num_segments=batch_size)
|
|
136
|
+
var = sq_diff / counts
|
|
137
|
+
unbiased_var = sq_diff / unbiased_counts
|
|
138
|
+
|
|
139
|
+
if is_training and self.track_running_stats:
|
|
140
|
+
m = self.momentum
|
|
141
|
+
cur_mean = ops.mean(mean, axis=0)
|
|
142
|
+
cur_var = ops.mean(unbiased_var, axis=0)
|
|
143
|
+
new_running_mean = (1.0 - m) * self.running_mean + m * cur_mean
|
|
144
|
+
new_running_var = (1.0 - m) * self.running_var + m * cur_var
|
|
145
|
+
self.running_mean.assign(new_running_mean)
|
|
146
|
+
self.running_var.assign(new_running_var)
|
|
147
|
+
|
|
148
|
+
std_x = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
|
|
149
|
+
out = x_c / std_x
|
|
150
|
+
else:
|
|
151
|
+
x_c = x - self.running_mean
|
|
152
|
+
std = ops.sqrt(self.running_var + self.eps)
|
|
153
|
+
out = x_c / std
|
|
154
|
+
|
|
155
|
+
if self.affine:
|
|
156
|
+
out = out * self.weight + self.bias
|
|
157
|
+
|
|
158
|
+
return out
|
|
159
|
+
|
|
160
|
+
def compute_output_shape(self, input_shape):
|
|
161
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
162
|
+
return input_shape[0]
|
|
163
|
+
return input_shape
|
|
@@ -0,0 +1,245 @@
|
|
|
1
|
+
from typing import List, Optional, Union
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
from k3_node.ops.segment import segment_sum
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class LayerNorm(layers.Layer):
|
|
7
|
+
r"""Applies layer normalization over each individual example in a batch
|
|
8
|
+
of features as described in the `"Layer Normalization"
|
|
9
|
+
<https://arxiv.org/abs/1607.06450>`_ paper.
|
|
10
|
+
|
|
11
|
+
.. math::
|
|
12
|
+
\mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
|
|
13
|
+
\textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
|
|
14
|
+
\odot \gamma + \beta
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
in_channels (int): Size of each input sample.
|
|
18
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
19
|
+
stability. (default: :obj:`1e-5`)
|
|
20
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
21
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
22
|
+
(default: :obj:`True`)
|
|
23
|
+
mode (str, optional): The normalization mode to use for layer
|
|
24
|
+
normalization (:obj:`"graph"` or :obj:`"node"`). If :obj:`"graph"`
|
|
25
|
+
is used, each graph will be considered as an element to be
|
|
26
|
+
normalized. If `"node"` is used, each node will be considered as
|
|
27
|
+
an element to be normalized. (default: :obj:`"graph"`)
|
|
28
|
+
|
|
29
|
+
Example:
|
|
30
|
+
```python
|
|
31
|
+
import numpy as np
|
|
32
|
+
from k3_node.layers import LayerNorm
|
|
33
|
+
|
|
34
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
35
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
36
|
+
|
|
37
|
+
layer = LayerNorm(in_channels=8, mode="graph")
|
|
38
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
39
|
+
print(tuple(out.shape)) # (10, 8)
|
|
40
|
+
```
|
|
41
|
+
"""
|
|
42
|
+
def __init__(
|
|
43
|
+
self,
|
|
44
|
+
in_channels: int,
|
|
45
|
+
eps: float = 1e-5,
|
|
46
|
+
affine: bool = True,
|
|
47
|
+
mode: str = 'graph',
|
|
48
|
+
**kwargs
|
|
49
|
+
):
|
|
50
|
+
super().__init__(**kwargs)
|
|
51
|
+
if mode not in ('graph', 'node'):
|
|
52
|
+
raise ValueError(f"Unknown normalization mode: {mode}")
|
|
53
|
+
|
|
54
|
+
self.in_channels = in_channels
|
|
55
|
+
self.eps = eps
|
|
56
|
+
self.affine = affine
|
|
57
|
+
self.mode = mode
|
|
58
|
+
|
|
59
|
+
if affine:
|
|
60
|
+
self.weight = self.add_weight(
|
|
61
|
+
shape=(in_channels,),
|
|
62
|
+
initializer="ones",
|
|
63
|
+
trainable=True,
|
|
64
|
+
name="weight",
|
|
65
|
+
)
|
|
66
|
+
self.bias = self.add_weight(
|
|
67
|
+
shape=(in_channels,),
|
|
68
|
+
initializer="zeros",
|
|
69
|
+
trainable=True,
|
|
70
|
+
name="bias",
|
|
71
|
+
)
|
|
72
|
+
else:
|
|
73
|
+
self.weight = None
|
|
74
|
+
self.bias = None
|
|
75
|
+
|
|
76
|
+
def reset_parameters(self):
|
|
77
|
+
if self.affine:
|
|
78
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
79
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
80
|
+
|
|
81
|
+
def call(self, x, batch=None, batch_size=None):
|
|
82
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
83
|
+
if len(x) == 2:
|
|
84
|
+
x, batch = x
|
|
85
|
+
elif len(x) == 3:
|
|
86
|
+
x, batch, batch_size = x
|
|
87
|
+
|
|
88
|
+
if self.mode == 'graph':
|
|
89
|
+
if batch is None:
|
|
90
|
+
mean = ops.mean(x)
|
|
91
|
+
var = ops.mean(ops.power(x - mean, 2))
|
|
92
|
+
out = (x - mean) / ops.sqrt(var + self.eps)
|
|
93
|
+
else:
|
|
94
|
+
if batch_size is not None and not isinstance(batch_size, int):
|
|
95
|
+
try:
|
|
96
|
+
batch_size = int(batch_size)
|
|
97
|
+
except Exception:
|
|
98
|
+
pass
|
|
99
|
+
elif batch_size is None:
|
|
100
|
+
batch_size = ops.cast(ops.max(batch), "int32") + 1
|
|
101
|
+
|
|
102
|
+
batch = ops.cast(batch, "int32")
|
|
103
|
+
in_channels = ops.cast(ops.shape(x)[-1], dtype=x.dtype)
|
|
104
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
105
|
+
node_counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
|
|
106
|
+
total_count = node_counts * in_channels
|
|
107
|
+
|
|
108
|
+
sum_x = ops.sum(segment_sum(x, batch, num_segments=batch_size), axis=-1, keepdims=True)
|
|
109
|
+
mean = sum_x / total_count
|
|
110
|
+
x_centered = x - ops.take(mean, batch, axis=0)
|
|
111
|
+
|
|
112
|
+
sum_sq = ops.sum(segment_sum(ops.power(x_centered, 2), batch, num_segments=batch_size), axis=-1, keepdims=True)
|
|
113
|
+
var = sum_sq / total_count
|
|
114
|
+
std_x = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
|
|
115
|
+
out = x_centered / std_x
|
|
116
|
+
|
|
117
|
+
if self.affine:
|
|
118
|
+
out = out * self.weight + self.bias
|
|
119
|
+
return out
|
|
120
|
+
|
|
121
|
+
elif self.mode == 'node':
|
|
122
|
+
mean = ops.mean(x, axis=-1, keepdims=True)
|
|
123
|
+
var = ops.var(x, axis=-1, keepdims=True)
|
|
124
|
+
out = (x - mean) / ops.sqrt(var + self.eps)
|
|
125
|
+
if self.affine:
|
|
126
|
+
out = out * self.weight + self.bias
|
|
127
|
+
return out
|
|
128
|
+
|
|
129
|
+
def compute_output_shape(self, input_shape):
|
|
130
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
131
|
+
return input_shape[0]
|
|
132
|
+
return input_shape
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
class HeteroLayerNorm(layers.Layer):
|
|
136
|
+
r"""Applies layer normalization over each individual example in a batch
|
|
137
|
+
of heterogeneous features as described in the `"Layer Normalization"
|
|
138
|
+
<https://arxiv.org/abs/1607.06450>`_ paper.
|
|
139
|
+
Compared to :class:`LayerNorm`, :class:`HeteroLayerNorm` applies
|
|
140
|
+
normalization individually for each node or edge type.
|
|
141
|
+
|
|
142
|
+
Args:
|
|
143
|
+
in_channels (int): Size of each input sample.
|
|
144
|
+
num_types (int): The number of types.
|
|
145
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
146
|
+
stability. (default: :obj:`1e-5`)
|
|
147
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
148
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
149
|
+
(default: :obj:`True`)
|
|
150
|
+
mode (str, optional): The normalization mode to use for layer
|
|
151
|
+
normalization (:obj:`"node"`). (default: :obj:`"node"`)
|
|
152
|
+
|
|
153
|
+
Example:
|
|
154
|
+
```python
|
|
155
|
+
import numpy as np
|
|
156
|
+
from k3_node.layers import HeteroLayerNorm
|
|
157
|
+
|
|
158
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
159
|
+
node_type = np.random.randint(0, 3, size=(10,)) # type of each node
|
|
160
|
+
|
|
161
|
+
layer = HeteroLayerNorm(in_channels=8, num_types=3)
|
|
162
|
+
out = layer(x, node_type) # separate statistics per node type
|
|
163
|
+
print(tuple(out.shape)) # (10, 8)
|
|
164
|
+
```
|
|
165
|
+
"""
|
|
166
|
+
def __init__(
|
|
167
|
+
self,
|
|
168
|
+
in_channels: int,
|
|
169
|
+
num_types: int,
|
|
170
|
+
eps: float = 1e-5,
|
|
171
|
+
affine: bool = True,
|
|
172
|
+
mode: str = 'node',
|
|
173
|
+
**kwargs
|
|
174
|
+
):
|
|
175
|
+
super().__init__(**kwargs)
|
|
176
|
+
if mode != 'node':
|
|
177
|
+
raise ValueError(f"HeteroLayerNorm only supports mode='node' (got '{mode}')")
|
|
178
|
+
|
|
179
|
+
self.in_channels = in_channels
|
|
180
|
+
self.num_types = num_types
|
|
181
|
+
self.eps = eps
|
|
182
|
+
self.affine = affine
|
|
183
|
+
self.mode = mode
|
|
184
|
+
|
|
185
|
+
if affine:
|
|
186
|
+
self.weight = self.add_weight(
|
|
187
|
+
shape=(num_types, in_channels),
|
|
188
|
+
initializer="ones",
|
|
189
|
+
trainable=True,
|
|
190
|
+
name="weight",
|
|
191
|
+
)
|
|
192
|
+
self.bias = self.add_weight(
|
|
193
|
+
shape=(num_types, in_channels),
|
|
194
|
+
initializer="zeros",
|
|
195
|
+
trainable=True,
|
|
196
|
+
name="bias",
|
|
197
|
+
)
|
|
198
|
+
else:
|
|
199
|
+
self.weight = None
|
|
200
|
+
self.bias = None
|
|
201
|
+
|
|
202
|
+
def reset_parameters(self):
|
|
203
|
+
if self.affine:
|
|
204
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
205
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
206
|
+
|
|
207
|
+
def call(
|
|
208
|
+
self,
|
|
209
|
+
x,
|
|
210
|
+
type_vec=None,
|
|
211
|
+
type_ptr: Optional[Union[list, tuple]] = None,
|
|
212
|
+
):
|
|
213
|
+
if type_vec is None and isinstance(x, (tuple, list)):
|
|
214
|
+
if len(x) == 2:
|
|
215
|
+
x, type_vec = x
|
|
216
|
+
elif len(x) == 3:
|
|
217
|
+
x, type_vec, type_ptr = x
|
|
218
|
+
|
|
219
|
+
if type_vec is None and type_ptr is None:
|
|
220
|
+
raise ValueError("Either 'type_vec' or 'type_ptr' must be given")
|
|
221
|
+
|
|
222
|
+
mean = ops.mean(x, axis=-1, keepdims=True)
|
|
223
|
+
var = ops.var(x, axis=-1, keepdims=True)
|
|
224
|
+
out = (x - mean) / ops.sqrt(var + self.eps)
|
|
225
|
+
|
|
226
|
+
if self.affine:
|
|
227
|
+
if type_ptr is not None:
|
|
228
|
+
parts = []
|
|
229
|
+
for i in range(len(type_ptr) - 1):
|
|
230
|
+
s, e = type_ptr[i], type_ptr[i + 1]
|
|
231
|
+
part = out[s:e] * self.weight[i] + self.bias[i]
|
|
232
|
+
parts.append(part)
|
|
233
|
+
out = ops.concatenate(parts, axis=0)
|
|
234
|
+
else:
|
|
235
|
+
type_vec = ops.cast(type_vec, "int32")
|
|
236
|
+
w = ops.take(self.weight, type_vec, axis=0)
|
|
237
|
+
b = ops.take(self.bias, type_vec, axis=0)
|
|
238
|
+
out = out * w + b
|
|
239
|
+
|
|
240
|
+
return out
|
|
241
|
+
|
|
242
|
+
def compute_output_shape(self, input_shape):
|
|
243
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
244
|
+
return input_shape[0]
|
|
245
|
+
return input_shape
|
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
from k3_node.ops.segment import segment_sum
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class MeanSubtractionNorm(layers.Layer):
|
|
6
|
+
r"""Applies layer normalization by subtracting the mean from the inputs
|
|
7
|
+
as described in the `"Revisiting 'Over-smoothing' in Deep GCNs"
|
|
8
|
+
<https://arxiv.org/abs/2003.13663>`_ paper.
|
|
9
|
+
|
|
10
|
+
.. math::
|
|
11
|
+
\mathbf{x}_i = \mathbf{x}_i - \frac{1}{|\mathcal{V}|}
|
|
12
|
+
\sum_{j \in \mathcal{V}} \mathbf{x}_j
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import MeanSubtractionNorm
|
|
18
|
+
|
|
19
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
20
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
21
|
+
|
|
22
|
+
layer = MeanSubtractionNorm()
|
|
23
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
24
|
+
print(tuple(out.shape)) # (10, 8)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(self, **kwargs):
|
|
28
|
+
super().__init__(**kwargs)
|
|
29
|
+
|
|
30
|
+
def call(self, x, batch=None, dim_size=None):
|
|
31
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
32
|
+
if len(x) == 2:
|
|
33
|
+
x, batch = x
|
|
34
|
+
elif len(x) == 3:
|
|
35
|
+
x, batch, dim_size = x
|
|
36
|
+
|
|
37
|
+
if batch is None:
|
|
38
|
+
return x - ops.mean(x, axis=0, keepdims=True)
|
|
39
|
+
|
|
40
|
+
if dim_size is not None and not isinstance(dim_size, int):
|
|
41
|
+
try:
|
|
42
|
+
dim_size = int(dim_size)
|
|
43
|
+
except Exception:
|
|
44
|
+
pass
|
|
45
|
+
elif dim_size is None:
|
|
46
|
+
dim_size = ops.cast(ops.max(batch), "int32") + 1
|
|
47
|
+
|
|
48
|
+
batch = ops.cast(batch, "int32")
|
|
49
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
50
|
+
counts = ops.maximum(segment_sum(ones, batch, num_segments=dim_size), 1.0)
|
|
51
|
+
mean = segment_sum(x, batch, num_segments=dim_size) / counts
|
|
52
|
+
return x - ops.take(mean, batch, axis=0)
|
|
53
|
+
|
|
54
|
+
def compute_output_shape(self, input_shape):
|
|
55
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
56
|
+
return input_shape[0]
|
|
57
|
+
return input_shape
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class MessageNorm(layers.Layer):
|
|
5
|
+
r"""Applies message normalization over the aggregated messages as described
|
|
6
|
+
in the `"DeeperGCNs: All You Need to Train Deeper GCNs"
|
|
7
|
+
<https://arxiv.org/abs/2006.07739>`_ paper.
|
|
8
|
+
|
|
9
|
+
.. math::
|
|
10
|
+
|
|
11
|
+
\mathbf{x}_i^{\prime} = \mathbf{x}_{i} + s \cdot
|
|
12
|
+
{\| \mathbf{x}_i \|}_2 \cdot
|
|
13
|
+
\frac{\mathbf{m}_{i}}{{\|\mathbf{m}_i\|}_2}
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
learn_scale (bool, optional): If set to :obj:`True`, will learn the
|
|
17
|
+
scaling factor :math:`s` of message normalization.
|
|
18
|
+
(default: :obj:`False`)
|
|
19
|
+
|
|
20
|
+
Example:
|
|
21
|
+
```python
|
|
22
|
+
import numpy as np
|
|
23
|
+
from k3_node.layers import MessageNorm
|
|
24
|
+
|
|
25
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
26
|
+
|
|
27
|
+
msg = np.random.rand(10, 8).astype("float32") # aggregated messages for each node
|
|
28
|
+
layer = MessageNorm(learn_scale=True)
|
|
29
|
+
out = layer(x, msg) # rescales messages to the norm of x
|
|
30
|
+
print(tuple(out.shape)) # (10, 8)
|
|
31
|
+
```
|
|
32
|
+
"""
|
|
33
|
+
def __init__(self, learn_scale: bool = False, **kwargs):
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
self.learn_scale = learn_scale
|
|
36
|
+
self.scale = self.add_weight(
|
|
37
|
+
shape=(1,),
|
|
38
|
+
initializer="ones",
|
|
39
|
+
trainable=learn_scale,
|
|
40
|
+
name="scale",
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
def reset_parameters(self):
|
|
44
|
+
self.scale.assign(ops.ones(self.scale.shape, dtype=self.scale.dtype))
|
|
45
|
+
|
|
46
|
+
def call(self, x, msg=None, p=2.0):
|
|
47
|
+
if msg is None and isinstance(x, (tuple, list)):
|
|
48
|
+
x, msg = x
|
|
49
|
+
|
|
50
|
+
msg_norm = ops.maximum(ops.norm(msg, ord=p, axis=-1, keepdims=True), 1e-12)
|
|
51
|
+
msg_normalized = msg / msg_norm
|
|
52
|
+
x_norm = ops.norm(x, ord=p, axis=-1, keepdims=True)
|
|
53
|
+
return msg_normalized * x_norm * self.scale
|
|
54
|
+
|
|
55
|
+
def compute_output_shape(self, input_shape):
|
|
56
|
+
if isinstance(input_shape, (tuple, list)):
|
|
57
|
+
return input_shape[0]
|
|
58
|
+
return input_shape
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
from k3_node.ops.segment import segment_sum
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class PairNorm(layers.Layer):
|
|
6
|
+
r"""Applies pair normalization over node features as described in the
|
|
7
|
+
`"PairNorm: Tackling Oversmoothing in GNNs"
|
|
8
|
+
<https://arxiv.org/abs/1909.12223>`_ paper.
|
|
9
|
+
|
|
10
|
+
.. math::
|
|
11
|
+
\mathbf{x}_i^c &= \mathbf{x}_i - \frac{1}{n}
|
|
12
|
+
\sum_{i=1}^n \mathbf{x}_i \\
|
|
13
|
+
|
|
14
|
+
\mathbf{x}_i^{\prime} &= s \cdot
|
|
15
|
+
\frac{\mathbf{x}_i^c}{\sqrt{\frac{1}{n} \sum_{i=1}^n
|
|
16
|
+
{\| \mathbf{x}_i^c \|}^2_2}}
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
scale (float, optional): Scaling factor :math:`s` of normalization.
|
|
20
|
+
(default: :obj:`1.0`)
|
|
21
|
+
scale_individually (bool, optional): If set to :obj:`True`, will
|
|
22
|
+
compute the scaling step as :math:`\mathbf{x}^{\prime}_i = s \cdot
|
|
23
|
+
\frac{\mathbf{x}_i^c}{{\| \mathbf{x}_i^c \|}_2}`.
|
|
24
|
+
(default: :obj:`False`)
|
|
25
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
26
|
+
stability. (default: :obj:`1e-5`)
|
|
27
|
+
|
|
28
|
+
Example:
|
|
29
|
+
```python
|
|
30
|
+
import numpy as np
|
|
31
|
+
from k3_node.layers import PairNorm
|
|
32
|
+
|
|
33
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
34
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
35
|
+
|
|
36
|
+
layer = PairNorm()
|
|
37
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
38
|
+
print(tuple(out.shape)) # (10, 8)
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
def __init__(self, scale: float = 1.0, scale_individually: bool = False,
|
|
42
|
+
eps: float = 1e-5, **kwargs):
|
|
43
|
+
super().__init__(**kwargs)
|
|
44
|
+
self.scale = scale
|
|
45
|
+
self.scale_individually = scale_individually
|
|
46
|
+
self.eps = eps
|
|
47
|
+
|
|
48
|
+
def call(self, x, batch=None, batch_size=None):
|
|
49
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
50
|
+
if len(x) == 2:
|
|
51
|
+
x, batch = x
|
|
52
|
+
elif len(x) == 3:
|
|
53
|
+
x, batch, batch_size = x
|
|
54
|
+
|
|
55
|
+
scale = self.scale
|
|
56
|
+
|
|
57
|
+
if batch is None:
|
|
58
|
+
x = x - ops.mean(x, axis=0, keepdims=True)
|
|
59
|
+
|
|
60
|
+
if not self.scale_individually:
|
|
61
|
+
mean_sq = ops.mean(ops.sum(ops.power(x, 2), axis=-1))
|
|
62
|
+
return scale * x / ops.sqrt(self.eps + mean_sq)
|
|
63
|
+
else:
|
|
64
|
+
norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True))
|
|
65
|
+
return scale * x / (self.eps + norm)
|
|
66
|
+
|
|
67
|
+
if batch_size is not None and not isinstance(batch_size, int):
|
|
68
|
+
try:
|
|
69
|
+
batch_size = int(batch_size)
|
|
70
|
+
except Exception:
|
|
71
|
+
pass
|
|
72
|
+
elif batch_size is None:
|
|
73
|
+
batch_size = ops.cast(ops.max(batch), "int32") + 1
|
|
74
|
+
|
|
75
|
+
batch = ops.cast(batch, "int32")
|
|
76
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
77
|
+
counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
|
|
78
|
+
mean = segment_sum(x, batch, num_segments=batch_size) / counts
|
|
79
|
+
x = x - ops.take(mean, batch, axis=0)
|
|
80
|
+
|
|
81
|
+
if not self.scale_individually:
|
|
82
|
+
sq_sum = ops.sum(ops.power(x, 2), axis=-1, keepdims=True)
|
|
83
|
+
mean_sq = segment_sum(sq_sum, batch, num_segments=batch_size) / counts
|
|
84
|
+
denom = ops.sqrt(self.eps + ops.take(mean_sq, batch, axis=0))
|
|
85
|
+
return scale * x / denom
|
|
86
|
+
else:
|
|
87
|
+
norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True))
|
|
88
|
+
return scale * x / (self.eps + norm)
|
|
89
|
+
|
|
90
|
+
def compute_output_shape(self, input_shape):
|
|
91
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
92
|
+
return input_shape[0]
|
|
93
|
+
return input_shape
|
|
94
|
+
|