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,96 @@
|
|
|
1
|
+
import math
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.kge.base import KGEModel, margin_ranking_loss, normalize
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class TransE(KGEModel):
|
|
10
|
+
r"""The TransE model from the `"Translating Embeddings for Modeling
|
|
11
|
+
Multi-Relational Data" <https://proceedings.neurips.cc/paper/2013/file/
|
|
12
|
+
1cecc7a77928ca8133fa24680a88d2f9-Paper.pdf>`_ paper.
|
|
13
|
+
|
|
14
|
+
:class:`TransE` models relations as a translation from head to tail
|
|
15
|
+
entities such that
|
|
16
|
+
|
|
17
|
+
.. math::
|
|
18
|
+
\mathbf{e}_h + \mathbf{e}_r \approx \mathbf{e}_t,
|
|
19
|
+
|
|
20
|
+
resulting in the scoring function:
|
|
21
|
+
|
|
22
|
+
.. math::
|
|
23
|
+
d(h, r, t) = - {\| \mathbf{e}_h + \mathbf{e}_r - \mathbf{e}_t \|}_p
|
|
24
|
+
|
|
25
|
+
Args:
|
|
26
|
+
num_nodes (int): The number of nodes/entities in the graph.
|
|
27
|
+
num_relations (int): The number of relations in the graph.
|
|
28
|
+
hidden_channels (int): The hidden embedding size.
|
|
29
|
+
margin (float, optional): The margin of the ranking loss.
|
|
30
|
+
(default: :obj:`1.0`)
|
|
31
|
+
p_norm (float, optional): The order embedding and distance
|
|
32
|
+
normalization. (default: :obj:`1.0`)
|
|
33
|
+
sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
|
|
34
|
+
|
|
35
|
+
Example:
|
|
36
|
+
```python
|
|
37
|
+
import numpy as np
|
|
38
|
+
from k3_node.layers import TransE
|
|
39
|
+
|
|
40
|
+
head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
|
|
41
|
+
rel = np.random.randint(0, 5, size=(10,))
|
|
42
|
+
tail = np.random.randint(0, 20, size=(10,))
|
|
43
|
+
|
|
44
|
+
model = TransE(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
45
|
+
score = model(head, rel, tail) # plausibility score of every triple
|
|
46
|
+
print(tuple(score.shape)) # (10,)
|
|
47
|
+
loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
|
|
48
|
+
print(tuple(loss.shape)) # (): a scalar
|
|
49
|
+
```
|
|
50
|
+
"""
|
|
51
|
+
def __init__(
|
|
52
|
+
self,
|
|
53
|
+
num_nodes: int,
|
|
54
|
+
num_relations: int,
|
|
55
|
+
hidden_channels: int,
|
|
56
|
+
margin: float = 1.0,
|
|
57
|
+
p_norm: float = 1.0,
|
|
58
|
+
sparse: bool = False,
|
|
59
|
+
**kwargs,
|
|
60
|
+
):
|
|
61
|
+
super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
|
|
62
|
+
|
|
63
|
+
self.p_norm = p_norm
|
|
64
|
+
self.margin = margin
|
|
65
|
+
|
|
66
|
+
self.reset_parameters()
|
|
67
|
+
|
|
68
|
+
def reset_parameters(self):
|
|
69
|
+
bound = 6.0 / math.sqrt(self.hidden_channels)
|
|
70
|
+
# A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
|
|
71
|
+
uniform = lambda shape: keras.initializers.RandomUniform(-bound, bound)(shape)
|
|
72
|
+
self.node_emb.embeddings.assign(uniform(ops.shape(self.node_emb.embeddings)))
|
|
73
|
+
self.rel_emb.embeddings.assign(uniform(ops.shape(self.rel_emb.embeddings)))
|
|
74
|
+
self.rel_emb.embeddings.assign(normalize(self.rel_emb.embeddings, p=self.p_norm, axis=-1))
|
|
75
|
+
|
|
76
|
+
def call(self, head_index, rel_type, tail_index):
|
|
77
|
+
head_index = ops.cast(head_index, "int32")
|
|
78
|
+
rel_type = ops.cast(rel_type, "int32")
|
|
79
|
+
tail_index = ops.cast(tail_index, "int32")
|
|
80
|
+
|
|
81
|
+
head = self.node_emb(head_index)
|
|
82
|
+
rel = self.rel_emb(rel_type)
|
|
83
|
+
tail = self.node_emb(tail_index)
|
|
84
|
+
|
|
85
|
+
head = normalize(head, p=self.p_norm, axis=-1)
|
|
86
|
+
tail = normalize(tail, p=self.p_norm, axis=-1)
|
|
87
|
+
|
|
88
|
+
# Calculate *negative* TransE norm:
|
|
89
|
+
diff = (head + rel) - tail
|
|
90
|
+
return -ops.power(ops.sum(ops.power(ops.abs(diff), self.p_norm), axis=-1), 1.0 / self.p_norm)
|
|
91
|
+
|
|
92
|
+
def loss(self, head_index, rel_type, tail_index):
|
|
93
|
+
pos_score = self(head_index, rel_type, tail_index)
|
|
94
|
+
neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
|
|
95
|
+
|
|
96
|
+
return margin_ranking_loss(pos_score, neg_score, margin=self.margin)
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from .batch_norm import BatchNorm, HeteroBatchNorm
|
|
2
|
+
from .diff_group_norm import DiffGroupNorm
|
|
3
|
+
from .graph_norm import GraphNorm
|
|
4
|
+
from .graph_size_norm import GraphSizeNorm
|
|
5
|
+
from .instance_norm import InstanceNorm
|
|
6
|
+
from .layer_norm import HeteroLayerNorm, LayerNorm
|
|
7
|
+
from .mean_subtraction_norm import MeanSubtractionNorm
|
|
8
|
+
from .msg_norm import MessageNorm
|
|
9
|
+
from .pair_norm import PairNorm
|
|
10
|
+
|
|
11
|
+
__all__ = [
|
|
12
|
+
"BatchNorm",
|
|
13
|
+
"HeteroBatchNorm",
|
|
14
|
+
"InstanceNorm",
|
|
15
|
+
"LayerNorm",
|
|
16
|
+
"HeteroLayerNorm",
|
|
17
|
+
"GraphNorm",
|
|
18
|
+
"GraphSizeNorm",
|
|
19
|
+
"PairNorm",
|
|
20
|
+
"MeanSubtractionNorm",
|
|
21
|
+
"MessageNorm",
|
|
22
|
+
"DiffGroupNorm",
|
|
23
|
+
]
|
|
@@ -0,0 +1,328 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
from k3_node.ops.segment import segment_sum
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class BatchNorm(layers.Layer):
|
|
7
|
+
r"""Applies batch normalization over a batch of features as described in
|
|
8
|
+
the `"Batch Normalization: Accelerating Deep Network Training by
|
|
9
|
+
Reducing Internal Covariate Shift" <https://arxiv.org/abs/1502.03167>`_
|
|
10
|
+
paper.
|
|
11
|
+
|
|
12
|
+
.. math::
|
|
13
|
+
\mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
|
|
14
|
+
\textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
|
|
15
|
+
\odot \gamma + \beta
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
in_channels (int): Size of each input sample.
|
|
19
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
20
|
+
stability. (default: :obj:`1e-5`)
|
|
21
|
+
momentum (float, optional): The value used for the running mean and
|
|
22
|
+
running variance computation. (default: :obj:`0.1`)
|
|
23
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
24
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
25
|
+
(default: :obj:`True`)
|
|
26
|
+
track_running_stats (bool, optional): If set to :obj:`True`, this
|
|
27
|
+
module tracks the running mean and variance, and when set to
|
|
28
|
+
:obj:`False`, this module does not track such statistics and always
|
|
29
|
+
uses batch statistics in both training and eval modes.
|
|
30
|
+
(default: :obj:`True`)
|
|
31
|
+
allow_single_element (bool, optional): If set to :obj:`True`, batches
|
|
32
|
+
with only a single element will work as during in evaluation.
|
|
33
|
+
That is the running mean and variance will be used.
|
|
34
|
+
Requires :obj:`track_running_stats=True`. (default: :obj:`False`)
|
|
35
|
+
|
|
36
|
+
Example:
|
|
37
|
+
```python
|
|
38
|
+
import numpy as np
|
|
39
|
+
from k3_node.layers import BatchNorm
|
|
40
|
+
|
|
41
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
42
|
+
|
|
43
|
+
layer = BatchNorm(in_channels=8)
|
|
44
|
+
out = layer(x, training=True) # uses batch statistics while training
|
|
45
|
+
print(tuple(out.shape)) # (10, 8)
|
|
46
|
+
```
|
|
47
|
+
"""
|
|
48
|
+
def __init__(
|
|
49
|
+
self,
|
|
50
|
+
in_channels: int,
|
|
51
|
+
eps: float = 1e-5,
|
|
52
|
+
momentum: Optional[float] = 0.1,
|
|
53
|
+
affine: bool = True,
|
|
54
|
+
track_running_stats: bool = True,
|
|
55
|
+
allow_single_element: bool = False,
|
|
56
|
+
**kwargs
|
|
57
|
+
):
|
|
58
|
+
super().__init__(**kwargs)
|
|
59
|
+
if allow_single_element and not track_running_stats:
|
|
60
|
+
raise ValueError("'allow_single_element' requires "
|
|
61
|
+
"'track_running_stats' to be set to `True`")
|
|
62
|
+
|
|
63
|
+
self.in_channels = in_channels
|
|
64
|
+
self.eps = eps
|
|
65
|
+
self.momentum = momentum
|
|
66
|
+
self.affine = affine
|
|
67
|
+
self.track_running_stats = track_running_stats
|
|
68
|
+
self.allow_single_element = allow_single_element
|
|
69
|
+
|
|
70
|
+
if affine:
|
|
71
|
+
self.weight = self.add_weight(
|
|
72
|
+
shape=(in_channels,),
|
|
73
|
+
initializer="ones",
|
|
74
|
+
trainable=True,
|
|
75
|
+
name="weight",
|
|
76
|
+
)
|
|
77
|
+
self.bias = self.add_weight(
|
|
78
|
+
shape=(in_channels,),
|
|
79
|
+
initializer="zeros",
|
|
80
|
+
trainable=True,
|
|
81
|
+
name="bias",
|
|
82
|
+
)
|
|
83
|
+
else:
|
|
84
|
+
self.weight = None
|
|
85
|
+
self.bias = None
|
|
86
|
+
|
|
87
|
+
if track_running_stats:
|
|
88
|
+
self.running_mean = self.add_weight(
|
|
89
|
+
shape=(in_channels,),
|
|
90
|
+
initializer="zeros",
|
|
91
|
+
trainable=False,
|
|
92
|
+
name="running_mean",
|
|
93
|
+
)
|
|
94
|
+
self.running_var = self.add_weight(
|
|
95
|
+
shape=(in_channels,),
|
|
96
|
+
initializer="ones",
|
|
97
|
+
trainable=False,
|
|
98
|
+
name="running_var",
|
|
99
|
+
)
|
|
100
|
+
self.num_batches_tracked = self.add_weight(
|
|
101
|
+
shape=(),
|
|
102
|
+
initializer="zeros",
|
|
103
|
+
dtype="int32",
|
|
104
|
+
trainable=False,
|
|
105
|
+
name="num_batches_tracked",
|
|
106
|
+
)
|
|
107
|
+
else:
|
|
108
|
+
self.running_mean = None
|
|
109
|
+
self.running_var = None
|
|
110
|
+
self.num_batches_tracked = None
|
|
111
|
+
|
|
112
|
+
def reset_running_stats(self):
|
|
113
|
+
if self.track_running_stats:
|
|
114
|
+
self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
|
|
115
|
+
self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
|
|
116
|
+
self.num_batches_tracked.assign(ops.cast(0, "int32"))
|
|
117
|
+
|
|
118
|
+
def reset_parameters(self):
|
|
119
|
+
self.reset_running_stats()
|
|
120
|
+
if self.affine:
|
|
121
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
122
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
123
|
+
|
|
124
|
+
def call(self, x, training=None):
|
|
125
|
+
num_samples_static = x.shape[0]
|
|
126
|
+
# Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
|
|
127
|
+
is_training = bool(training) if training is not None else False
|
|
128
|
+
|
|
129
|
+
if is_training:
|
|
130
|
+
if num_samples_static is not None and num_samples_static <= 1:
|
|
131
|
+
if not self.allow_single_element:
|
|
132
|
+
raise ValueError(f"Expected more than 1 value per channel when training, got input size {ops.shape(x)}")
|
|
133
|
+
# Evaluation behavior with running stats
|
|
134
|
+
mean = self.running_mean
|
|
135
|
+
var = self.running_var
|
|
136
|
+
else:
|
|
137
|
+
mean = ops.mean(x, axis=0)
|
|
138
|
+
var = ops.var(x, axis=0)
|
|
139
|
+
|
|
140
|
+
if self.track_running_stats:
|
|
141
|
+
n = ops.cast(ops.shape(x)[0], dtype=x.dtype)
|
|
142
|
+
unbiased_var = var * n / ops.maximum(n - 1.0, 1.0)
|
|
143
|
+
if self.momentum is None:
|
|
144
|
+
count = ops.cast(self.num_batches_tracked + 1, dtype=x.dtype)
|
|
145
|
+
m = 1.0 / count
|
|
146
|
+
else:
|
|
147
|
+
m = self.momentum
|
|
148
|
+
|
|
149
|
+
new_running_mean = (1.0 - m) * self.running_mean + m * mean
|
|
150
|
+
new_running_var = (1.0 - m) * self.running_var + m * unbiased_var
|
|
151
|
+
self.running_mean.assign(new_running_mean)
|
|
152
|
+
self.running_var.assign(new_running_var)
|
|
153
|
+
self.num_batches_tracked.assign(self.num_batches_tracked + 1)
|
|
154
|
+
else:
|
|
155
|
+
if self.track_running_stats:
|
|
156
|
+
mean = self.running_mean
|
|
157
|
+
var = self.running_var
|
|
158
|
+
else:
|
|
159
|
+
mean = ops.mean(x, axis=0)
|
|
160
|
+
var = ops.var(x, axis=0)
|
|
161
|
+
|
|
162
|
+
out = (x - mean) / ops.sqrt(var + self.eps)
|
|
163
|
+
|
|
164
|
+
if self.affine:
|
|
165
|
+
out = out * self.weight + self.bias
|
|
166
|
+
|
|
167
|
+
return out
|
|
168
|
+
|
|
169
|
+
def compute_output_shape(self, input_shape):
|
|
170
|
+
return input_shape
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
class HeteroBatchNorm(layers.Layer):
|
|
174
|
+
r"""Applies batch normalization over a batch of heterogeneous features as
|
|
175
|
+
described in the `"Batch Normalization: Accelerating Deep Network Training
|
|
176
|
+
by Reducing Internal Covariate Shift" <https://arxiv.org/abs/1502.03167>`_
|
|
177
|
+
paper.
|
|
178
|
+
Compared to :class:`BatchNorm`, :class:`HeteroBatchNorm` applies
|
|
179
|
+
normalization individually for each node or edge type.
|
|
180
|
+
|
|
181
|
+
Args:
|
|
182
|
+
in_channels (int): Size of each input sample.
|
|
183
|
+
num_types (int): The number of types.
|
|
184
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
185
|
+
stability. (default: :obj:`1e-5`)
|
|
186
|
+
momentum (float, optional): The value used for the running mean and
|
|
187
|
+
running variance computation. (default: :obj:`0.1`)
|
|
188
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
189
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
190
|
+
(default: :obj:`True`)
|
|
191
|
+
track_running_stats (bool, optional): If set to :obj:`True`, this
|
|
192
|
+
module tracks the running mean and variance, and when set to
|
|
193
|
+
:obj:`False`, this module does not track such statistics and always
|
|
194
|
+
uses batch statistics in both training and eval modes.
|
|
195
|
+
(default: :obj:`True`)
|
|
196
|
+
|
|
197
|
+
Example:
|
|
198
|
+
```python
|
|
199
|
+
import numpy as np
|
|
200
|
+
from k3_node.layers import HeteroBatchNorm
|
|
201
|
+
|
|
202
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
203
|
+
node_type = np.random.randint(0, 3, size=(10,)) # type of each node
|
|
204
|
+
|
|
205
|
+
layer = HeteroBatchNorm(in_channels=8, num_types=3)
|
|
206
|
+
out = layer(x, node_type, training=True) # separate statistics per node type
|
|
207
|
+
print(tuple(out.shape)) # (10, 8)
|
|
208
|
+
```
|
|
209
|
+
"""
|
|
210
|
+
def __init__(
|
|
211
|
+
self,
|
|
212
|
+
in_channels: int,
|
|
213
|
+
num_types: int,
|
|
214
|
+
eps: float = 1e-5,
|
|
215
|
+
momentum: Optional[float] = 0.1,
|
|
216
|
+
affine: bool = True,
|
|
217
|
+
track_running_stats: bool = True,
|
|
218
|
+
**kwargs
|
|
219
|
+
):
|
|
220
|
+
super().__init__(**kwargs)
|
|
221
|
+
self.in_channels = in_channels
|
|
222
|
+
self.num_types = num_types
|
|
223
|
+
self.eps = eps
|
|
224
|
+
self.momentum = momentum
|
|
225
|
+
self.affine = affine
|
|
226
|
+
self.track_running_stats = track_running_stats
|
|
227
|
+
|
|
228
|
+
if affine:
|
|
229
|
+
self.weight = self.add_weight(
|
|
230
|
+
shape=(num_types, in_channels),
|
|
231
|
+
initializer="ones",
|
|
232
|
+
trainable=True,
|
|
233
|
+
name="weight",
|
|
234
|
+
)
|
|
235
|
+
self.bias = self.add_weight(
|
|
236
|
+
shape=(num_types, in_channels),
|
|
237
|
+
initializer="zeros",
|
|
238
|
+
trainable=True,
|
|
239
|
+
name="bias",
|
|
240
|
+
)
|
|
241
|
+
else:
|
|
242
|
+
self.weight = None
|
|
243
|
+
self.bias = None
|
|
244
|
+
|
|
245
|
+
if track_running_stats:
|
|
246
|
+
self.running_mean = self.add_weight(
|
|
247
|
+
shape=(num_types, in_channels),
|
|
248
|
+
initializer="zeros",
|
|
249
|
+
trainable=False,
|
|
250
|
+
name="running_mean",
|
|
251
|
+
)
|
|
252
|
+
self.running_var = self.add_weight(
|
|
253
|
+
shape=(num_types, in_channels),
|
|
254
|
+
initializer="ones",
|
|
255
|
+
trainable=False,
|
|
256
|
+
name="running_var",
|
|
257
|
+
)
|
|
258
|
+
self.num_batches_tracked = self.add_weight(
|
|
259
|
+
shape=(),
|
|
260
|
+
initializer="zeros",
|
|
261
|
+
dtype="int32",
|
|
262
|
+
trainable=False,
|
|
263
|
+
name="num_batches_tracked",
|
|
264
|
+
)
|
|
265
|
+
else:
|
|
266
|
+
self.running_mean = None
|
|
267
|
+
self.running_var = None
|
|
268
|
+
self.num_batches_tracked = None
|
|
269
|
+
|
|
270
|
+
def reset_running_stats(self):
|
|
271
|
+
if self.track_running_stats:
|
|
272
|
+
self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
|
|
273
|
+
self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
|
|
274
|
+
self.num_batches_tracked.assign(ops.cast(0, "int32"))
|
|
275
|
+
|
|
276
|
+
def reset_parameters(self):
|
|
277
|
+
self.reset_running_stats()
|
|
278
|
+
if self.affine:
|
|
279
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
280
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
281
|
+
|
|
282
|
+
def call(self, x, type_vec=None, training=None):
|
|
283
|
+
if type_vec is None and isinstance(x, (tuple, list)):
|
|
284
|
+
x, type_vec = x
|
|
285
|
+
|
|
286
|
+
# Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
|
|
287
|
+
is_training = bool(training) if training is not None else False
|
|
288
|
+
type_vec = ops.cast(type_vec, "int32")
|
|
289
|
+
|
|
290
|
+
if not is_training and self.track_running_stats:
|
|
291
|
+
mean = self.running_mean
|
|
292
|
+
var = self.running_var
|
|
293
|
+
else:
|
|
294
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
295
|
+
counts = ops.maximum(segment_sum(ones, type_vec, num_segments=self.num_types), 1.0)
|
|
296
|
+
mean = segment_sum(x, type_vec, num_segments=self.num_types) / counts
|
|
297
|
+
x_c = x - ops.take(mean, type_vec, axis=0)
|
|
298
|
+
var = segment_sum(ops.power(x_c, 2), type_vec, num_segments=self.num_types) / counts
|
|
299
|
+
|
|
300
|
+
if is_training and self.track_running_stats:
|
|
301
|
+
if self.momentum is None:
|
|
302
|
+
count = ops.cast(self.num_batches_tracked + 1, dtype=x.dtype)
|
|
303
|
+
exp_avg_factor = 1.0 / count
|
|
304
|
+
else:
|
|
305
|
+
exp_avg_factor = self.momentum
|
|
306
|
+
|
|
307
|
+
new_running_mean = (1.0 - exp_avg_factor) * self.running_mean + exp_avg_factor * mean
|
|
308
|
+
new_running_var = (1.0 - exp_avg_factor) * self.running_var + exp_avg_factor * var
|
|
309
|
+
self.running_mean.assign(new_running_mean)
|
|
310
|
+
self.running_var.assign(new_running_var)
|
|
311
|
+
self.num_batches_tracked.assign(self.num_batches_tracked + 1)
|
|
312
|
+
|
|
313
|
+
std = ops.sqrt(var + self.eps)
|
|
314
|
+
mean_taken = ops.take(mean, type_vec, axis=0)
|
|
315
|
+
std_taken = ops.take(std, type_vec, axis=0)
|
|
316
|
+
out = (x - mean_taken) / std_taken
|
|
317
|
+
|
|
318
|
+
if self.affine:
|
|
319
|
+
w_taken = ops.take(self.weight, type_vec, axis=0)
|
|
320
|
+
b_taken = ops.take(self.bias, type_vec, axis=0)
|
|
321
|
+
out = out * w_taken + b_taken
|
|
322
|
+
|
|
323
|
+
return out
|
|
324
|
+
|
|
325
|
+
def compute_output_shape(self, input_shape):
|
|
326
|
+
if isinstance(input_shape, (tuple, list)):
|
|
327
|
+
return input_shape[0]
|
|
328
|
+
return input_shape
|
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from scipy.spatial.distance import cdist
|
|
3
|
+
from keras import initializers, layers, ops
|
|
4
|
+
|
|
5
|
+
from .batch_norm import BatchNorm
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class DiffGroupNorm(layers.Layer):
|
|
9
|
+
r"""The differentiable group normalization layer from the `"Towards Deeper
|
|
10
|
+
Graph Neural Networks with Differentiable Group Normalization"
|
|
11
|
+
<https://arxiv.org/abs/2006.06972>`_ paper, which normalizes node features
|
|
12
|
+
group-wise via a learnable soft cluster assignment.
|
|
13
|
+
|
|
14
|
+
.. math::
|
|
15
|
+
|
|
16
|
+
\mathbf{S} = \text{softmax} (\mathbf{X} \mathbf{W})
|
|
17
|
+
|
|
18
|
+
where :math:`\mathbf{W} \in \mathbb{R}^{F \times G}` denotes a trainable
|
|
19
|
+
weight matrix mapping each node into one of :math:`G` clusters.
|
|
20
|
+
Normalization is then performed group-wise via:
|
|
21
|
+
|
|
22
|
+
.. math::
|
|
23
|
+
|
|
24
|
+
\mathbf{X}^{\prime} = \mathbf{X} + \lambda \sum_{i = 1}^G
|
|
25
|
+
\text{BatchNorm}(\mathbf{S}[:, i] \odot \mathbf{X})
|
|
26
|
+
|
|
27
|
+
Args:
|
|
28
|
+
in_channels (int): Size of each input sample :math:`F`.
|
|
29
|
+
groups (int): The number of groups :math:`G`.
|
|
30
|
+
lamda (float, optional): The balancing factor :math:`\lambda` between
|
|
31
|
+
input embeddings and normalized embeddings. (default: :obj:`0.01`)
|
|
32
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
33
|
+
stability. (default: :obj:`1e-5`)
|
|
34
|
+
momentum (float, optional): The value used for the running mean and
|
|
35
|
+
running variance computation. (default: :obj:`0.1`)
|
|
36
|
+
affine (bool, optional): If set to :obj:`True`, this module has
|
|
37
|
+
learnable affine parameters :math:`\gamma` and :math:`\beta`.
|
|
38
|
+
(default: :obj:`True`)
|
|
39
|
+
track_running_stats (bool, optional): If set to :obj:`True`, this
|
|
40
|
+
module tracks the running mean and variance, and when set to
|
|
41
|
+
:obj:`False`, this module does not track such statistics and always
|
|
42
|
+
uses batch statistics in both training and eval modes.
|
|
43
|
+
(default: :obj:`True`)
|
|
44
|
+
|
|
45
|
+
Example:
|
|
46
|
+
```python
|
|
47
|
+
import numpy as np
|
|
48
|
+
from k3_node.layers import DiffGroupNorm
|
|
49
|
+
|
|
50
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
51
|
+
|
|
52
|
+
layer = DiffGroupNorm(in_channels=8, groups=2)
|
|
53
|
+
out = layer(x, training=True) # uses batch statistics while training
|
|
54
|
+
print(tuple(out.shape)) # (10, 8)
|
|
55
|
+
```
|
|
56
|
+
"""
|
|
57
|
+
def __init__(
|
|
58
|
+
self,
|
|
59
|
+
in_channels: int,
|
|
60
|
+
groups: int,
|
|
61
|
+
lamda: float = 0.01,
|
|
62
|
+
eps: float = 1e-5,
|
|
63
|
+
momentum: float = 0.1,
|
|
64
|
+
affine: bool = True,
|
|
65
|
+
track_running_stats: bool = True,
|
|
66
|
+
**kwargs
|
|
67
|
+
):
|
|
68
|
+
super().__init__(**kwargs)
|
|
69
|
+
self.in_channels = in_channels
|
|
70
|
+
self.groups = groups
|
|
71
|
+
self.lamda = lamda
|
|
72
|
+
self.eps = eps
|
|
73
|
+
self.momentum = momentum
|
|
74
|
+
self.affine = affine
|
|
75
|
+
self.track_running_stats = track_running_stats
|
|
76
|
+
|
|
77
|
+
self.lin_weight = self.add_weight(
|
|
78
|
+
shape=(in_channels, groups),
|
|
79
|
+
initializer="glorot_uniform",
|
|
80
|
+
trainable=True,
|
|
81
|
+
name="lin_weight",
|
|
82
|
+
)
|
|
83
|
+
self.norm = BatchNorm(
|
|
84
|
+
groups * in_channels,
|
|
85
|
+
eps=eps,
|
|
86
|
+
momentum=momentum,
|
|
87
|
+
affine=affine,
|
|
88
|
+
track_running_stats=track_running_stats,
|
|
89
|
+
name="norm",
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
def reset_parameters(self):
|
|
93
|
+
self.lin_weight.assign(
|
|
94
|
+
initializers.GlorotUniform()(shape=self.lin_weight.shape, dtype=self.lin_weight.dtype)
|
|
95
|
+
)
|
|
96
|
+
self.norm.reset_parameters()
|
|
97
|
+
|
|
98
|
+
def call(self, x, training=None):
|
|
99
|
+
F, G = self.in_channels, self.groups
|
|
100
|
+
|
|
101
|
+
s = ops.softmax(ops.matmul(x, self.lin_weight), axis=-1) # [N, G]
|
|
102
|
+
out = ops.expand_dims(s, axis=-1) * ops.expand_dims(x, axis=-2) # [N, G, F]
|
|
103
|
+
out_flat = ops.reshape(out, (-1, G * F))
|
|
104
|
+
out_norm = self.norm(out_flat, training=training)
|
|
105
|
+
out = ops.sum(ops.reshape(out_norm, (-1, G, F)), axis=-2) # [N, F]
|
|
106
|
+
|
|
107
|
+
return x + self.lamda * out
|
|
108
|
+
|
|
109
|
+
@staticmethod
|
|
110
|
+
def group_distance_ratio(x, y, eps: float = 1e-5) -> float:
|
|
111
|
+
r"""Measures the ratio of inter-group distance over intra-group
|
|
112
|
+
distance.
|
|
113
|
+
"""
|
|
114
|
+
x_np = ops.convert_to_numpy(x)
|
|
115
|
+
y_np = ops.convert_to_numpy(y).astype(np.int64)
|
|
116
|
+
|
|
117
|
+
num_classes = int(y_np.max()) + 1
|
|
118
|
+
|
|
119
|
+
numerator = 0.0
|
|
120
|
+
for i in range(num_classes):
|
|
121
|
+
mask = (y_np == i)
|
|
122
|
+
if not np.any(mask) or np.all(mask):
|
|
123
|
+
continue
|
|
124
|
+
dist = cdist(x_np[mask], x_np[~mask])
|
|
125
|
+
numerator += (1.0 / dist.size) * float(dist.sum())
|
|
126
|
+
numerator *= 1.0 / ((num_classes - 1) ** 2)
|
|
127
|
+
|
|
128
|
+
denominator = 0.0
|
|
129
|
+
for i in range(num_classes):
|
|
130
|
+
mask = (y_np == i)
|
|
131
|
+
if not np.any(mask):
|
|
132
|
+
continue
|
|
133
|
+
dist = cdist(x_np[mask], x_np[mask])
|
|
134
|
+
denominator += (1.0 / dist.size) * float(dist.sum())
|
|
135
|
+
denominator *= 1.0 / num_classes
|
|
136
|
+
|
|
137
|
+
return float(numerator / (denominator + eps))
|
|
138
|
+
|
|
139
|
+
def compute_output_shape(self, input_shape):
|
|
140
|
+
return input_shape
|
|
141
|
+
|
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
3
|
+
from k3_node.ops.segment import segment_sum
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class GraphNorm(layers.Layer):
|
|
7
|
+
r"""Applies graph normalization over individual graphs as described in the
|
|
8
|
+
`"GraphNorm: A Principled Approach to Accelerating Graph Neural Network
|
|
9
|
+
Training" <https://arxiv.org/abs/2009.03294>`_ paper.
|
|
10
|
+
|
|
11
|
+
.. math::
|
|
12
|
+
\mathbf{x}^{\prime}_i = \frac{\mathbf{x} - \alpha \odot
|
|
13
|
+
\textrm{E}[\mathbf{x}]}
|
|
14
|
+
{\sqrt{\textrm{Var}[\mathbf{x} - \alpha \odot \textrm{E}[\mathbf{x}]]
|
|
15
|
+
+ \epsilon}} \odot \gamma + \beta
|
|
16
|
+
|
|
17
|
+
where :math:`\alpha` denotes parameters that learn how much information
|
|
18
|
+
to keep in the mean.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
in_channels (int): Size of each input sample.
|
|
22
|
+
eps (float, optional): A value added to the denominator for numerical
|
|
23
|
+
stability. (default: :obj:`1e-5`)
|
|
24
|
+
|
|
25
|
+
Example:
|
|
26
|
+
```python
|
|
27
|
+
import numpy as np
|
|
28
|
+
from k3_node.layers import GraphNorm
|
|
29
|
+
|
|
30
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
31
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
32
|
+
|
|
33
|
+
layer = GraphNorm(in_channels=8)
|
|
34
|
+
out = layer(x, batch) # normalizes each graph separately
|
|
35
|
+
print(tuple(out.shape)) # (10, 8)
|
|
36
|
+
```
|
|
37
|
+
"""
|
|
38
|
+
def __init__(self, in_channels: int, eps: float = 1e-5, **kwargs):
|
|
39
|
+
super().__init__(**kwargs)
|
|
40
|
+
self.in_channels = in_channels
|
|
41
|
+
self.eps = eps
|
|
42
|
+
|
|
43
|
+
self.weight = self.add_weight(
|
|
44
|
+
shape=(in_channels,),
|
|
45
|
+
initializer="ones",
|
|
46
|
+
trainable=True,
|
|
47
|
+
name="weight",
|
|
48
|
+
)
|
|
49
|
+
self.bias = self.add_weight(
|
|
50
|
+
shape=(in_channels,),
|
|
51
|
+
initializer="zeros",
|
|
52
|
+
trainable=True,
|
|
53
|
+
name="bias",
|
|
54
|
+
)
|
|
55
|
+
self.mean_scale = self.add_weight(
|
|
56
|
+
shape=(in_channels,),
|
|
57
|
+
initializer="ones",
|
|
58
|
+
trainable=True,
|
|
59
|
+
name="mean_scale",
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
def reset_parameters(self):
|
|
63
|
+
self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
|
|
64
|
+
self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
|
|
65
|
+
self.mean_scale.assign(ops.ones(self.mean_scale.shape, dtype=self.mean_scale.dtype))
|
|
66
|
+
|
|
67
|
+
def call(self, x, batch=None, batch_size=None):
|
|
68
|
+
if batch is None and isinstance(x, (tuple, list)):
|
|
69
|
+
if len(x) == 2:
|
|
70
|
+
x, batch = x
|
|
71
|
+
elif len(x) == 3:
|
|
72
|
+
x, batch, batch_size = x
|
|
73
|
+
|
|
74
|
+
if batch is None:
|
|
75
|
+
mean = ops.mean(x, axis=0, keepdims=True)
|
|
76
|
+
out = x - mean * self.mean_scale
|
|
77
|
+
var = ops.mean(ops.power(out, 2), axis=0, keepdims=True)
|
|
78
|
+
std = ops.sqrt(var + self.eps)
|
|
79
|
+
return self.weight * out / std + self.bias
|
|
80
|
+
|
|
81
|
+
if batch_size is None:
|
|
82
|
+
if not is_tracing(batch):
|
|
83
|
+
try:
|
|
84
|
+
batch_size = int(ops.max(batch)) + 1
|
|
85
|
+
except Exception:
|
|
86
|
+
batch_size = None
|
|
87
|
+
elif not isinstance(batch_size, int):
|
|
88
|
+
try:
|
|
89
|
+
batch_size = int(batch_size)
|
|
90
|
+
except Exception:
|
|
91
|
+
pass
|
|
92
|
+
|
|
93
|
+
batch = ops.cast(batch, "int32")
|
|
94
|
+
ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
|
|
95
|
+
counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
|
|
96
|
+
mean = segment_sum(x, batch, num_segments=batch_size) / counts
|
|
97
|
+
out = x - ops.take(mean, batch, axis=0) * self.mean_scale
|
|
98
|
+
var = segment_sum(ops.power(out, 2), batch, num_segments=batch_size) / counts
|
|
99
|
+
std = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
|
|
100
|
+
return self.weight * out / std + self.bias
|
|
101
|
+
|
|
102
|
+
def compute_output_shape(self, input_shape):
|
|
103
|
+
if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
|
|
104
|
+
return input_shape[0]
|
|
105
|
+
return input_shape
|