k3-node 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- k3_node/__init__.py +122 -0
- k3_node/applications/__init__.py +17 -0
- k3_node/applications/bio/__init__.py +21 -0
- k3_node/applications/chemistry/__init__.py +155 -0
- k3_node/applications/materials/__init__.py +127 -0
- k3_node/applications/materials/basis.py +449 -0
- k3_node/applications/materials/chgnet.py +360 -0
- k3_node/applications/materials/core.py +351 -0
- k3_node/applications/materials/grace.py +246 -0
- k3_node/applications/materials/io.py +230 -0
- k3_node/applications/materials/m3gnet.py +462 -0
- k3_node/applications/materials/megnet.py +395 -0
- k3_node/applications/materials/qet.py +220 -0
- k3_node/applications/materials/readout.py +235 -0
- k3_node/applications/materials/so3net.py +234 -0
- k3_node/applications/materials/tensornet.py +381 -0
- k3_node/applications/materials/test_materials.py +167 -0
- k3_node/applications/materials/wrappers.py +95 -0
- k3_node/data/__init__.py +47 -0
- k3_node/data/batch.py +102 -0
- k3_node/data/collate.py +282 -0
- k3_node/data/data.py +532 -0
- k3_node/data/database.py +154 -0
- k3_node/data/dataset.py +182 -0
- k3_node/data/download.py +49 -0
- k3_node/data/extract.py +45 -0
- k3_node/data/feature_store.py +70 -0
- k3_node/data/graph_store.py +92 -0
- k3_node/data/hetero_data.py +374 -0
- k3_node/data/hypergraph_data.py +59 -0
- k3_node/data/in_memory_dataset.py +177 -0
- k3_node/data/makedirs.py +7 -0
- k3_node/data/on_disk_dataset.py +77 -0
- k3_node/data/separate.py +115 -0
- k3_node/data/storage.py +593 -0
- k3_node/data/temporal.py +154 -0
- k3_node/data/test_batch.py +67 -0
- k3_node/data/test_data.py +68 -0
- k3_node/data/test_dataset_and_stores.py +111 -0
- k3_node/data/test_hetero_data.py +33 -0
- k3_node/data/test_temporal_and_hyper.py +32 -0
- k3_node/data/view.py +43 -0
- k3_node/datasets/__init__.py +88 -0
- k3_node/datasets/actor.py +101 -0
- k3_node/datasets/airports.py +84 -0
- k3_node/datasets/amazon.py +66 -0
- k3_node/datasets/ba2motif_dataset.py +73 -0
- k3_node/datasets/ba_shapes.py +81 -0
- k3_node/datasets/bitcoin_otc.py +77 -0
- k3_node/datasets/citation_full.py +81 -0
- k3_node/datasets/coauthor.py +66 -0
- k3_node/datasets/dblp.py +106 -0
- k3_node/datasets/digits.py +63 -0
- k3_node/datasets/email_eu_core.py +60 -0
- k3_node/datasets/entities.py +158 -0
- k3_node/datasets/explainer_dataset.py +101 -0
- k3_node/datasets/facebook.py +51 -0
- k3_node/datasets/fake.py +256 -0
- k3_node/datasets/freebase.py +90 -0
- k3_node/datasets/geometric_shapes.py +69 -0
- k3_node/datasets/github.py +51 -0
- k3_node/datasets/graph_generator/__init__.py +6 -0
- k3_node/datasets/graph_generator/ba_graph.py +20 -0
- k3_node/datasets/graph_generator/base.py +29 -0
- k3_node/datasets/graph_generator/er_graph.py +21 -0
- k3_node/datasets/icews.py +58 -0
- k3_node/datasets/imdb.py +96 -0
- k3_node/datasets/jodie.py +56 -0
- k3_node/datasets/karate.py +56 -0
- k3_node/datasets/lastfm_asia.py +51 -0
- k3_node/datasets/mesh_correspondence.py +50 -0
- k3_node/datasets/molecule_net.py +148 -0
- k3_node/datasets/motif_generator/__init__.py +7 -0
- k3_node/datasets/motif_generator/base.py +29 -0
- k3_node/datasets/motif_generator/custom.py +17 -0
- k3_node/datasets/motif_generator/cycle.py +25 -0
- k3_node/datasets/motif_generator/house.py +27 -0
- k3_node/datasets/movielens.py +55 -0
- k3_node/datasets/planetoid.py +137 -0
- k3_node/datasets/polblogs.py +63 -0
- k3_node/datasets/ppi.py +189 -0
- k3_node/datasets/qm7.py +65 -0
- k3_node/datasets/qm9.py +132 -0
- k3_node/datasets/reddit.py +121 -0
- k3_node/datasets/sbm_dataset.py +165 -0
- k3_node/datasets/seal.py +74 -0
- k3_node/datasets/shape_scenes.py +92 -0
- k3_node/datasets/test_datasets.py +322 -0
- k3_node/datasets/tu_dataset.py +131 -0
- k3_node/datasets/twitch.py +66 -0
- k3_node/datasets/webkb.py +102 -0
- k3_node/datasets/wikics.py +85 -0
- k3_node/datasets/word_net.py +184 -0
- k3_node/etl/__init__.py +37 -0
- k3_node/etl/encoders.py +248 -0
- k3_node/etl/graph_builders.py +270 -0
- k3_node/etl/relational_to_graph.py +201 -0
- k3_node/etl/table_to_graph.py +244 -0
- k3_node/etl/test_etl.py +318 -0
- k3_node/export/__init__.py +15 -0
- k3_node/export/cross_backend.py +172 -0
- k3_node/export/onnx_exporter.py +190 -0
- k3_node/export/runtime.py +254 -0
- k3_node/export/tensorrt_exporter.py +201 -0
- k3_node/export/test_export.py +337 -0
- k3_node/export/tflite_exporter.py +112 -0
- k3_node/hub/__init__.py +29 -0
- k3_node/hub/dataset_hub.py +242 -0
- k3_node/hub/hub_mixin.py +599 -0
- k3_node/hub/model_card.py +133 -0
- k3_node/hub/test_hub.py +419 -0
- k3_node/io/__init__.py +22 -0
- k3_node/io/fs.py +117 -0
- k3_node/io/npz.py +45 -0
- k3_node/io/off.py +29 -0
- k3_node/io/planetoid.py +98 -0
- k3_node/io/tu.py +137 -0
- k3_node/io/txt_array.py +58 -0
- k3_node/layers/__init__.py +14 -0
- k3_node/layers/aggr/__init__.py +70 -0
- k3_node/layers/aggr/attention.py +77 -0
- k3_node/layers/aggr/base.py +403 -0
- k3_node/layers/aggr/basic.py +412 -0
- k3_node/layers/aggr/deep_sets.py +65 -0
- k3_node/layers/aggr/deepsets.py +29 -0
- k3_node/layers/aggr/equilibrium.py +107 -0
- k3_node/layers/aggr/fused.py +43 -0
- k3_node/layers/aggr/gmt.py +89 -0
- k3_node/layers/aggr/gru.py +58 -0
- k3_node/layers/aggr/lcm.py +143 -0
- k3_node/layers/aggr/lstm.py +58 -0
- k3_node/layers/aggr/mlp.py +75 -0
- k3_node/layers/aggr/multi.py +154 -0
- k3_node/layers/aggr/patch_transformer.py +137 -0
- k3_node/layers/aggr/quantile.py +125 -0
- k3_node/layers/aggr/resolver.py +68 -0
- k3_node/layers/aggr/scaler.py +133 -0
- k3_node/layers/aggr/set2set.py +87 -0
- k3_node/layers/aggr/set_transformer.py +107 -0
- k3_node/layers/aggr/sort.py +68 -0
- k3_node/layers/aggr/test_aggr.py +337 -0
- k3_node/layers/aggr/utils.py +210 -0
- k3_node/layers/aggr/variance_preserving.py +54 -0
- k3_node/layers/attention/__init__.py +5 -0
- k3_node/layers/attention/pair_attention.py +448 -0
- k3_node/layers/attention/performer.py +187 -0
- k3_node/layers/attention/polynormer.py +160 -0
- k3_node/layers/attention/qformer.py +143 -0
- k3_node/layers/attention/sgformer.py +106 -0
- k3_node/layers/attention/test_attention.py +68 -0
- k3_node/layers/attention/test_pair_attention.py +91 -0
- k3_node/layers/conv/__init__.py +149 -0
- k3_node/layers/conv/agnn_conv.py +120 -0
- k3_node/layers/conv/antisymmetric_conv.py +94 -0
- k3_node/layers/conv/appnp.py +105 -0
- k3_node/layers/conv/appnp_conv.py +157 -0
- k3_node/layers/conv/arma_conv.py +231 -0
- k3_node/layers/conv/cg_conv.py +92 -0
- k3_node/layers/conv/cheb_conv.py +137 -0
- k3_node/layers/conv/cluster_gcn_conv.py +102 -0
- k3_node/layers/conv/conv.py +100 -0
- k3_node/layers/conv/crystal_conv.py +140 -0
- k3_node/layers/conv/cugraph.py +84 -0
- k3_node/layers/conv/diffusion_conv.py +144 -0
- k3_node/layers/conv/dir_gnn_conv.py +93 -0
- k3_node/layers/conv/dna_conv.py +192 -0
- k3_node/layers/conv/edge_conv.py +107 -0
- k3_node/layers/conv/eg_conv.py +155 -0
- k3_node/layers/conv/fa_conv.py +107 -0
- k3_node/layers/conv/feast_conv.py +126 -0
- k3_node/layers/conv/film_conv.py +143 -0
- k3_node/layers/conv/gat_conv.py +244 -0
- k3_node/layers/conv/gated_graph_conv.py +136 -0
- k3_node/layers/conv/gatv2_conv.py +205 -0
- k3_node/layers/conv/gcn.py +144 -0
- k3_node/layers/conv/gcn2_conv.py +126 -0
- k3_node/layers/conv/gcn_conv.py +135 -0
- k3_node/layers/conv/gen_conv.py +163 -0
- k3_node/layers/conv/general_conv.py +218 -0
- k3_node/layers/conv/gin_conv.py +218 -0
- k3_node/layers/conv/gmm_conv.py +172 -0
- k3_node/layers/conv/gps_conv.py +153 -0
- k3_node/layers/conv/graph_attention.py +262 -0
- k3_node/layers/conv/graph_conv.py +84 -0
- k3_node/layers/conv/gravnet_conv.py +93 -0
- k3_node/layers/conv/han_conv.py +175 -0
- k3_node/layers/conv/heat_conv.py +131 -0
- k3_node/layers/conv/hetero_conv.py +128 -0
- k3_node/layers/conv/hgt_conv.py +218 -0
- k3_node/layers/conv/hypergraph_conv.py +182 -0
- k3_node/layers/conv/le_conv.py +81 -0
- k3_node/layers/conv/lg_conv.py +58 -0
- k3_node/layers/conv/meshcnn_conv.py +84 -0
- k3_node/layers/conv/message_passing.py +451 -0
- k3_node/layers/conv/mf_conv.py +95 -0
- k3_node/layers/conv/mixhop_conv.py +108 -0
- k3_node/layers/conv/nn_conv.py +110 -0
- k3_node/layers/conv/pan_conv.py +100 -0
- k3_node/layers/conv/pdn_conv.py +109 -0
- k3_node/layers/conv/pna_conv.py +177 -0
- k3_node/layers/conv/point_conv.py +101 -0
- k3_node/layers/conv/point_gnn_conv.py +90 -0
- k3_node/layers/conv/point_transformer_conv.py +132 -0
- k3_node/layers/conv/ppf_conv.py +135 -0
- k3_node/layers/conv/ppnp.py +89 -0
- k3_node/layers/conv/res_gated_graph_conv.py +126 -0
- k3_node/layers/conv/rgat_conv.py +251 -0
- k3_node/layers/conv/rgcn_conv.py +321 -0
- k3_node/layers/conv/sage_conv.py +154 -0
- k3_node/layers/conv/sg_conv.py +96 -0
- k3_node/layers/conv/signed_conv.py +100 -0
- k3_node/layers/conv/simple_conv.py +75 -0
- k3_node/layers/conv/spline_conv.py +182 -0
- k3_node/layers/conv/ssg_conv.py +101 -0
- k3_node/layers/conv/supergat_conv.py +195 -0
- k3_node/layers/conv/tag_conv.py +98 -0
- k3_node/layers/conv/test_backend_consistency.py +164 -0
- k3_node/layers/conv/test_conv.py +176 -0
- k3_node/layers/conv/test_conv_pyg.py +566 -0
- k3_node/layers/conv/transformer_conv.py +168 -0
- k3_node/layers/conv/utils.py +403 -0
- k3_node/layers/conv/wl_conv.py +151 -0
- k3_node/layers/conv/x_conv.py +187 -0
- k3_node/layers/dense/__init__.py +40 -0
- k3_node/layers/dense/dense_gat_conv.py +149 -0
- k3_node/layers/dense/dense_gcn_conv.py +117 -0
- k3_node/layers/dense/dense_gin_conv.py +88 -0
- k3_node/layers/dense/dense_graph_conv.py +95 -0
- k3_node/layers/dense/dense_sage_conv.py +85 -0
- k3_node/layers/dense/diff_pool.py +76 -0
- k3_node/layers/dense/dmon_pool.py +223 -0
- k3_node/layers/dense/linear.py +327 -0
- k3_node/layers/dense/mincut_pool.py +92 -0
- k3_node/layers/dense/test_dense.py +377 -0
- k3_node/layers/functional/__init__.py +13 -0
- k3_node/layers/functional/bro.py +49 -0
- k3_node/layers/functional/edge_dropout.py +55 -0
- k3_node/layers/functional/gini.py +44 -0
- k3_node/layers/functional/test_functional.py +34 -0
- k3_node/layers/kge/__init__.py +17 -0
- k3_node/layers/kge/base.py +255 -0
- k3_node/layers/kge/complex.py +98 -0
- k3_node/layers/kge/distmult.py +79 -0
- k3_node/layers/kge/loader.py +50 -0
- k3_node/layers/kge/rotate.py +103 -0
- k3_node/layers/kge/test_kge.py +76 -0
- k3_node/layers/kge/transe.py +96 -0
- k3_node/layers/norm/__init__.py +23 -0
- k3_node/layers/norm/batch_norm.py +328 -0
- k3_node/layers/norm/diff_group_norm.py +141 -0
- k3_node/layers/norm/graph_norm.py +105 -0
- k3_node/layers/norm/graph_size_norm.py +57 -0
- k3_node/layers/norm/instance_norm.py +163 -0
- k3_node/layers/norm/layer_norm.py +245 -0
- k3_node/layers/norm/mean_subtraction_norm.py +57 -0
- k3_node/layers/norm/msg_norm.py +58 -0
- k3_node/layers/norm/pair_norm.py +94 -0
- k3_node/layers/norm/test_norm.py +275 -0
- k3_node/layers/pool/__init__.py +83 -0
- k3_node/layers/pool/approx_knn.py +101 -0
- k3_node/layers/pool/asap.py +173 -0
- k3_node/layers/pool/avg_pool.py +165 -0
- k3_node/layers/pool/cluster_pool.py +168 -0
- k3_node/layers/pool/connect/__init__.py +10 -0
- k3_node/layers/pool/connect/base.py +103 -0
- k3_node/layers/pool/connect/filter_edges.py +113 -0
- k3_node/layers/pool/consecutive.py +30 -0
- k3_node/layers/pool/decimation.py +48 -0
- k3_node/layers/pool/edge_pool.py +189 -0
- k3_node/layers/pool/glob.py +139 -0
- k3_node/layers/pool/graclus.py +66 -0
- k3_node/layers/pool/knn.py +253 -0
- k3_node/layers/pool/max_pool.py +159 -0
- k3_node/layers/pool/mem_pool.py +145 -0
- k3_node/layers/pool/pan_pool.py +144 -0
- k3_node/layers/pool/point_cloud.py +212 -0
- k3_node/layers/pool/pool.py +119 -0
- k3_node/layers/pool/sag_pool.py +174 -0
- k3_node/layers/pool/select/__init__.py +10 -0
- k3_node/layers/pool/select/base.py +112 -0
- k3_node/layers/pool/select/topk.py +206 -0
- k3_node/layers/pool/test_pool.py +456 -0
- k3_node/layers/pool/topk_pool.py +103 -0
- k3_node/layers/pool/voxel_grid.py +70 -0
- k3_node/layers/unpool/__init__.py +9 -0
- k3_node/layers/unpool/knn_interpolate.py +57 -0
- k3_node/layers/unpool/test_unpool.py +31 -0
- k3_node/loader/__init__.py +62 -0
- k3_node/loader/base.py +69 -0
- k3_node/loader/cache.py +68 -0
- k3_node/loader/cluster.py +127 -0
- k3_node/loader/data_list_loader.py +45 -0
- k3_node/loader/dataloader.py +117 -0
- k3_node/loader/dense_data_loader.py +62 -0
- k3_node/loader/dynamic_batch_sampler.py +93 -0
- k3_node/loader/graph_saint.py +188 -0
- k3_node/loader/hgt_loader.py +90 -0
- k3_node/loader/imbalanced_sampler.py +87 -0
- k3_node/loader/keras_dataset.py +334 -0
- k3_node/loader/link_loader.py +179 -0
- k3_node/loader/link_neighbor_loader.py +202 -0
- k3_node/loader/mixin.py +190 -0
- k3_node/loader/neighbor_loader.py +159 -0
- k3_node/loader/neighbor_sampler.py +167 -0
- k3_node/loader/node_loader.py +185 -0
- k3_node/loader/prefetch.py +115 -0
- k3_node/loader/random_node_loader.py +89 -0
- k3_node/loader/sampler_utils.py +499 -0
- k3_node/loader/shadow.py +115 -0
- k3_node/loader/temporal_dataloader.py +98 -0
- k3_node/loader/test_dataloader.py +113 -0
- k3_node/loader/test_keras_dataset.py +221 -0
- k3_node/loader/test_neighbor_loader.py +122 -0
- k3_node/loader/test_sampler_utils.py +82 -0
- k3_node/loader/test_samplers.py +96 -0
- k3_node/loader/test_subgraph_loaders.py +89 -0
- k3_node/loader/utils.py +232 -0
- k3_node/loader/zip_loader.py +88 -0
- k3_node/metrics.py +94 -0
- k3_node/models/__init__.py +424 -0
- k3_node/models/attentive_fp.py +232 -0
- k3_node/models/attract_repel.py +108 -0
- k3_node/models/autoencoder.py +318 -0
- k3_node/models/basic_gnn.py +443 -0
- k3_node/models/bio/__init__.py +4 -0
- k3_node/models/captum.py +52 -0
- k3_node/models/chemistry/__init__.py +4 -0
- k3_node/models/correct_and_smooth.py +146 -0
- k3_node/models/deep_graph_infomax.py +113 -0
- k3_node/models/deepgcn.py +121 -0
- k3_node/models/dimenet.py +737 -0
- k3_node/models/dimenet_utils.py +153 -0
- k3_node/models/gnnff.py +263 -0
- k3_node/models/gps_model.py +1122 -0
- k3_node/models/gpse.py +638 -0
- k3_node/models/graph_unet.py +199 -0
- k3_node/models/graphmae2.py +954 -0
- k3_node/models/graphormer.py +1258 -0
- k3_node/models/graphormer_3d.py +868 -0
- k3_node/models/grover.py +1066 -0
- k3_node/models/jumping_knowledge.py +200 -0
- k3_node/models/label_prop.py +110 -0
- k3_node/models/lightgcn.py +171 -0
- k3_node/models/linkx.py +181 -0
- k3_node/models/lpformer.py +404 -0
- k3_node/models/mask_label.py +114 -0
- k3_node/models/materials/__init__.py +33 -0
- k3_node/models/meta.py +133 -0
- k3_node/models/metapath2vec.py +234 -0
- k3_node/models/mlp.py +264 -0
- k3_node/models/mole_bert.py +379 -0
- k3_node/models/neural_fingerprint.py +95 -0
- k3_node/models/node2vec.py +213 -0
- k3_node/models/pmlp.py +157 -0
- k3_node/models/polynormer.py +229 -0
- k3_node/models/rect.py +93 -0
- k3_node/models/renet.py +221 -0
- k3_node/models/rev_gnn.py +128 -0
- k3_node/models/schnet.py +484 -0
- k3_node/models/sgformer.py +195 -0
- k3_node/models/signed_gcn.py +185 -0
- k3_node/models/test_attentive_fp.py +32 -0
- k3_node/models/test_attract_repel.py +33 -0
- k3_node/models/test_autoencoder.py +119 -0
- k3_node/models/test_basic_gnn.py +102 -0
- k3_node/models/test_correct_and_smooth.py +40 -0
- k3_node/models/test_deep_graph_infomax.py +68 -0
- k3_node/models/test_deepgcn.py +21 -0
- k3_node/models/test_dimenet.py +86 -0
- k3_node/models/test_domain_apis.py +138 -0
- k3_node/models/test_gnnff.py +24 -0
- k3_node/models/test_gps_model.py +271 -0
- k3_node/models/test_gpse.py +34 -0
- k3_node/models/test_graph_unet.py +26 -0
- k3_node/models/test_graphmae2.py +226 -0
- k3_node/models/test_graphormer.py +233 -0
- k3_node/models/test_graphormer3d.py +163 -0
- k3_node/models/test_grover.py +287 -0
- k3_node/models/test_jumping_knowledge.py +129 -0
- k3_node/models/test_label_prop.py +37 -0
- k3_node/models/test_lightgcn.py +38 -0
- k3_node/models/test_linkx.py +31 -0
- k3_node/models/test_lpformer.py +22 -0
- k3_node/models/test_mask_label.py +90 -0
- k3_node/models/test_meta.py +159 -0
- k3_node/models/test_metapath2vec.py +45 -0
- k3_node/models/test_mlp.py +62 -0
- k3_node/models/test_mole_bert.py +164 -0
- k3_node/models/test_neural_fingerprint.py +13 -0
- k3_node/models/test_node2vec.py +57 -0
- k3_node/models/test_pmlp.py +81 -0
- k3_node/models/test_polynormer.py +104 -0
- k3_node/models/test_rect.py +23 -0
- k3_node/models/test_renet.py +32 -0
- k3_node/models/test_rev_gnn.py +24 -0
- k3_node/models/test_schnet.py +43 -0
- k3_node/models/test_sgformer.py +48 -0
- k3_node/models/test_signed_gcn.py +28 -0
- k3_node/models/test_tgn.py +77 -0
- k3_node/models/test_unimol.py +179 -0
- k3_node/models/test_unimol2.py +114 -0
- k3_node/models/test_unimol_plus.py +131 -0
- k3_node/models/test_visnet.py +44 -0
- k3_node/models/tgn.py +382 -0
- k3_node/models/unimol.py +1156 -0
- k3_node/models/unimol2.py +616 -0
- k3_node/models/unimol_docking_v2.py +301 -0
- k3_node/models/unimol_plus.py +456 -0
- k3_node/models/utils.py +97 -0
- k3_node/models/visnet.py +759 -0
- k3_node/ops/__init__.py +4 -0
- k3_node/ops/conv.py +56 -0
- k3_node/ops/creation.py +43 -0
- k3_node/ops/graph.py +27 -0
- k3_node/ops/host.py +41 -0
- k3_node/ops/matmul.py +49 -0
- k3_node/ops/numpy.py +24 -0
- k3_node/ops/segment.py +54 -0
- k3_node/ops/sparse.py +51 -0
- k3_node/rag/__init__.py +49 -0
- k3_node/rag/encoders.py +312 -0
- k3_node/rag/pipeline.py +192 -0
- k3_node/rag/projector.py +184 -0
- k3_node/rag/subgraph.py +270 -0
- k3_node/rag/test_rag.py +347 -0
- k3_node/rag/verbalizer.py +162 -0
- k3_node/tasks/__init__.py +19 -0
- k3_node/tasks/backbone_resolver.py +125 -0
- k3_node/tasks/base.py +67 -0
- k3_node/tasks/graph_classification.py +270 -0
- k3_node/tasks/graph_regression.py +228 -0
- k3_node/tasks/link_prediction.py +306 -0
- k3_node/tasks/node_classification.py +194 -0
- k3_node/tasks/node_regression.py +138 -0
- k3_node/tasks/test_tasks.py +319 -0
- k3_node/test_docstring_examples.py +106 -0
- k3_node/test_training_forwarding.py +116 -0
- k3_node/training.py +115 -0
- k3_node/transforms/__init__.py +166 -0
- k3_node/transforms/base_transform.py +32 -0
- k3_node/transforms/compose.py +58 -0
- k3_node/transforms/general.py +676 -0
- k3_node/transforms/graph.py +1070 -0
- k3_node/transforms/spatial.py +797 -0
- k3_node/transforms/test_random_link_split.py +45 -0
- k3_node/transforms/test_spatial_transforms.py +65 -0
- k3_node/transforms/test_transforms.py +253 -0
- k3_node/transforms/utils.py +102 -0
- k3_node/utils/__init__.py +5 -0
- k3_node/utils/backend_import.py +12 -0
- k3_node/utils/graph.py +286 -0
- k3_node/utils/keras.py +94 -0
- k3_node/utils/random.py +103 -0
- k3_node/utils/smiles.py +235 -0
- k3_node-1.0.0.dist-info/METADATA +284 -0
- k3_node-1.0.0.dist-info/RECORD +459 -0
- k3_node-1.0.0.dist-info/WHEEL +5 -0
- k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
- k3_node-1.0.0.dist-info/top_level.txt +1 -0
k3_node/models/pmlp.py
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv import SimpleConv
|
|
7
|
+
from k3_node.layers.norm import BatchNorm
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class PMLP(keras.Model):
|
|
11
|
+
r"""The P(ropagational)MLP model from the `"Graph Neural Networks are
|
|
12
|
+
Inherently Good Generalizers: Insights by Bridging GNNs and MLPs"
|
|
13
|
+
<https://arxiv.org/abs/2212.09034>`_ paper.
|
|
14
|
+
|
|
15
|
+
:class:`PMLP` is identical to a standard MLP during training, but then
|
|
16
|
+
adopts a GNN architecture during testing.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
in_channels (int): Size of each input sample.
|
|
20
|
+
hidden_channels (int): Size of each hidden sample.
|
|
21
|
+
out_channels (int): Size of each output sample.
|
|
22
|
+
num_layers (int): The number of layers.
|
|
23
|
+
dropout (float, optional): Dropout probability of each hidden
|
|
24
|
+
embedding. (default: :obj:`0.`)
|
|
25
|
+
norm (bool, optional): If set to :obj:`False`, will not apply batch
|
|
26
|
+
normalization. (default: :obj:`True`)
|
|
27
|
+
bias (bool, optional): If set to :obj:`False`, the module will not
|
|
28
|
+
learn additive biases. (default: :obj:`True`)
|
|
29
|
+
|
|
30
|
+
Example:
|
|
31
|
+
```python
|
|
32
|
+
import numpy as np
|
|
33
|
+
from k3_node.models import PMLP
|
|
34
|
+
|
|
35
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
36
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
37
|
+
|
|
38
|
+
model = PMLP(in_channels=8, hidden_channels=32, out_channels=4, num_layers=2)
|
|
39
|
+
print(tuple(model(x, edge_index).shape)) # (10, 4): message passing is used at inference
|
|
40
|
+
```
|
|
41
|
+
"""
|
|
42
|
+
|
|
43
|
+
def __init__(
|
|
44
|
+
self,
|
|
45
|
+
in_channels: int,
|
|
46
|
+
hidden_channels: int,
|
|
47
|
+
out_channels: int,
|
|
48
|
+
num_layers: int,
|
|
49
|
+
dropout: float = 0.,
|
|
50
|
+
norm: bool = True,
|
|
51
|
+
bias: bool = True,
|
|
52
|
+
**kwargs,
|
|
53
|
+
):
|
|
54
|
+
super().__init__(**kwargs)
|
|
55
|
+
|
|
56
|
+
self.in_channels = in_channels
|
|
57
|
+
self.hidden_channels = hidden_channels
|
|
58
|
+
self.out_channels = out_channels
|
|
59
|
+
self.num_layers = num_layers
|
|
60
|
+
self.dropout_rate = dropout
|
|
61
|
+
self.use_norm = norm
|
|
62
|
+
self.use_bias = bias
|
|
63
|
+
# Instance-level training flag matching PyG's API: pmlp.training = False
|
|
64
|
+
self.training = True
|
|
65
|
+
|
|
66
|
+
# Build weight dimensions
|
|
67
|
+
dims = [in_channels] + [hidden_channels] * (num_layers - 1) + [out_channels]
|
|
68
|
+
self._weight_shapes = [(dims[i], dims[i + 1]) for i in range(num_layers)]
|
|
69
|
+
|
|
70
|
+
# We store weights as raw keras Variables so they are always available
|
|
71
|
+
self._weights_list = []
|
|
72
|
+
self._biases_list = []
|
|
73
|
+
for i, (in_d, out_d) in enumerate(self._weight_shapes):
|
|
74
|
+
w = self.add_weight(
|
|
75
|
+
shape=(in_d, out_d),
|
|
76
|
+
initializer=keras.initializers.GlorotUniform(),
|
|
77
|
+
trainable=True,
|
|
78
|
+
name=f"weight_{i}",
|
|
79
|
+
)
|
|
80
|
+
self._weights_list.append(w)
|
|
81
|
+
if bias:
|
|
82
|
+
b = self.add_weight(
|
|
83
|
+
shape=(out_d,),
|
|
84
|
+
initializer="zeros",
|
|
85
|
+
trainable=True,
|
|
86
|
+
name=f"bias_{i}",
|
|
87
|
+
)
|
|
88
|
+
self._biases_list.append(b)
|
|
89
|
+
else:
|
|
90
|
+
self._biases_list.append(None)
|
|
91
|
+
|
|
92
|
+
self._norm_layer = None
|
|
93
|
+
if norm:
|
|
94
|
+
self._norm_layer = BatchNorm(
|
|
95
|
+
hidden_channels,
|
|
96
|
+
affine=False,
|
|
97
|
+
track_running_stats=False,
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
self.conv = SimpleConv(aggr='mean', combine_root='self_loop')
|
|
101
|
+
self._dropout = keras.layers.Dropout(dropout)
|
|
102
|
+
|
|
103
|
+
def build(self, input_shape=None):
|
|
104
|
+
self.built = True
|
|
105
|
+
|
|
106
|
+
def reset_parameters(self) -> None:
|
|
107
|
+
r"""Resets all learnable parameters of the module."""
|
|
108
|
+
for i, (in_d, out_d) in enumerate(self._weight_shapes):
|
|
109
|
+
self._weights_list[i].assign(
|
|
110
|
+
keras.initializers.GlorotUniform()(shape=(in_d, out_d))
|
|
111
|
+
)
|
|
112
|
+
if self.use_bias and self._biases_list[i] is not None:
|
|
113
|
+
self._biases_list[i].assign(ops.zeros((out_d,)))
|
|
114
|
+
|
|
115
|
+
def call(self, x, edge_index=None, training=None):
|
|
116
|
+
"""Forward pass.
|
|
117
|
+
|
|
118
|
+
Args:
|
|
119
|
+
x (Tensor): The node features of shape ``[N, in_channels]``.
|
|
120
|
+
edge_index (Tensor, optional): The edge indices. Required during
|
|
121
|
+
inference. (default: :obj:`None`)
|
|
122
|
+
training (bool, optional): Override the instance-level
|
|
123
|
+
``self.training`` flag. (default: :obj:`None`)
|
|
124
|
+
"""
|
|
125
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
126
|
+
if len(x) >= 2:
|
|
127
|
+
x, edge_index = x[0], x[1]
|
|
128
|
+
# Respect both call-time kwarg and instance-level flag (PyG compat)
|
|
129
|
+
is_training = training if training is not None else self.training
|
|
130
|
+
|
|
131
|
+
if not is_training and edge_index is None:
|
|
132
|
+
raise ValueError(
|
|
133
|
+
f"'edge_index' needs to be present during inference "
|
|
134
|
+
f"in '{self.__class__.__name__}'"
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
for i in range(self.num_layers):
|
|
138
|
+
# Apply weight multiplication (like x @ W^T where W is [in, out])
|
|
139
|
+
x = x @ self._weights_list[i]
|
|
140
|
+
|
|
141
|
+
if not is_training:
|
|
142
|
+
x = self.conv(x, edge_index)
|
|
143
|
+
|
|
144
|
+
if self.use_bias and self._biases_list[i] is not None:
|
|
145
|
+
x = x + self._biases_list[i]
|
|
146
|
+
|
|
147
|
+
if i != self.num_layers - 1:
|
|
148
|
+
if self._norm_layer is not None:
|
|
149
|
+
x = self._norm_layer(x, training=is_training)
|
|
150
|
+
x = ops.relu(x)
|
|
151
|
+
x = self._dropout(x, training=is_training)
|
|
152
|
+
|
|
153
|
+
return x
|
|
154
|
+
|
|
155
|
+
def __repr__(self) -> str:
|
|
156
|
+
return (f'{self.__class__.__name__}({self.in_channels}, '
|
|
157
|
+
f'{self.out_channels}, num_layers={self.num_layers})')
|
|
@@ -0,0 +1,229 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.aggr.base import from_dense_batch, to_dense_batch
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
from k3_node.layers.conv import GATConv, GCNConv
|
|
10
|
+
from k3_node.layers.attention import PolynormerAttention
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class Polynormer(keras.layers.Layer):
|
|
14
|
+
r"""The Polynormer module from the `"Polynormer: polynomial-expressive
|
|
15
|
+
graph transformer in linear time"
|
|
16
|
+
<https://arxiv.org/abs/2403.01232>`_ paper.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
in_channels (int): Input channels.
|
|
20
|
+
hidden_channels (int): Hidden channels.
|
|
21
|
+
out_channels (int): Output channels.
|
|
22
|
+
local_layers (int): The number of local attention layers.
|
|
23
|
+
(default: :obj:`7`)
|
|
24
|
+
global_layers (int): The number of global attention layers.
|
|
25
|
+
(default: :obj:`2`)
|
|
26
|
+
in_dropout (float): Input dropout rate.
|
|
27
|
+
(default: :obj:`0.15`)
|
|
28
|
+
dropout (float): Dropout rate.
|
|
29
|
+
(default: :obj:`0.5`)
|
|
30
|
+
global_dropout (float): Global dropout rate.
|
|
31
|
+
(default: :obj:`0.5`)
|
|
32
|
+
heads (int): The number of heads.
|
|
33
|
+
(default: :obj:`1`)
|
|
34
|
+
beta (float): Aggregate type.
|
|
35
|
+
(default: :obj:`0.9`)
|
|
36
|
+
qk_shared (bool, optional): Whether weight of query and key are shared.
|
|
37
|
+
(default: :obj:`True`)
|
|
38
|
+
pre_ln (bool): Pre layer normalization.
|
|
39
|
+
(default: :obj:`False`)
|
|
40
|
+
post_bn (bool): Post batch normalization.
|
|
41
|
+
(default: :obj:`True`)
|
|
42
|
+
local_attn (bool): Whether use local attention (GATConv vs GCNConv).
|
|
43
|
+
(default: :obj:`False`)
|
|
44
|
+
|
|
45
|
+
Example:
|
|
46
|
+
```python
|
|
47
|
+
import numpy as np
|
|
48
|
+
from k3_node.models import Polynormer
|
|
49
|
+
|
|
50
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
51
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
52
|
+
|
|
53
|
+
batch = np.repeat([0, 1], 5) # two graphs with 5 nodes each
|
|
54
|
+
model = Polynormer(in_channels=8, hidden_channels=32, out_channels=4,
|
|
55
|
+
local_layers=2, global_layers=1, heads=2)
|
|
56
|
+
out = model(x, edge_index, batch)
|
|
57
|
+
print(tuple(out.shape)) # (10, 4)
|
|
58
|
+
```
|
|
59
|
+
"""
|
|
60
|
+
|
|
61
|
+
def __init__(
|
|
62
|
+
self,
|
|
63
|
+
in_channels: int,
|
|
64
|
+
hidden_channels: int,
|
|
65
|
+
out_channels: int,
|
|
66
|
+
local_layers: int = 7,
|
|
67
|
+
global_layers: int = 2,
|
|
68
|
+
in_dropout: float = 0.15,
|
|
69
|
+
dropout: float = 0.5,
|
|
70
|
+
global_dropout: float = 0.5,
|
|
71
|
+
heads: int = 1,
|
|
72
|
+
beta: float = 0.9,
|
|
73
|
+
qk_shared: bool = False,
|
|
74
|
+
pre_ln: bool = False,
|
|
75
|
+
post_bn: bool = True,
|
|
76
|
+
local_attn: bool = False,
|
|
77
|
+
**kwargs,
|
|
78
|
+
) -> None:
|
|
79
|
+
super().__init__(**kwargs)
|
|
80
|
+
|
|
81
|
+
self._global = False
|
|
82
|
+
self.in_drop = in_dropout
|
|
83
|
+
self.dropout = dropout
|
|
84
|
+
self.pre_ln = pre_ln
|
|
85
|
+
self.post_bn = post_bn
|
|
86
|
+
self.beta = beta
|
|
87
|
+
self.heads = heads
|
|
88
|
+
self.hidden_channels = hidden_channels
|
|
89
|
+
self.local_attn = local_attn
|
|
90
|
+
|
|
91
|
+
inner_channels = heads * hidden_channels
|
|
92
|
+
|
|
93
|
+
self.h_lins = []
|
|
94
|
+
self.local_convs = []
|
|
95
|
+
self.lins = []
|
|
96
|
+
self.lns = []
|
|
97
|
+
self.pre_lns = [] if pre_ln else None
|
|
98
|
+
self.post_bns = [] if post_bn else None
|
|
99
|
+
|
|
100
|
+
# ---- First local layer ----
|
|
101
|
+
self.h_lins.append(keras.layers.Dense(inner_channels))
|
|
102
|
+
if local_attn:
|
|
103
|
+
self.local_convs.append(
|
|
104
|
+
GATConv(in_channels, hidden_channels, heads=heads, concat=True,
|
|
105
|
+
add_self_loops=False, bias=False)
|
|
106
|
+
)
|
|
107
|
+
else:
|
|
108
|
+
self.local_convs.append(
|
|
109
|
+
GCNConv(in_channels, inner_channels, cached=False, normalize=True)
|
|
110
|
+
)
|
|
111
|
+
self.lins.append(keras.layers.Dense(inner_channels))
|
|
112
|
+
self.lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
|
|
113
|
+
if pre_ln:
|
|
114
|
+
self.pre_lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
|
|
115
|
+
if post_bn:
|
|
116
|
+
self.post_bns.append(
|
|
117
|
+
keras.layers.BatchNormalization(
|
|
118
|
+
center=True, scale=True, momentum=0.9, epsilon=1e-5,
|
|
119
|
+
)
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
# ---- Subsequent local layers ----
|
|
123
|
+
for _ in range(local_layers - 1):
|
|
124
|
+
self.h_lins.append(keras.layers.Dense(inner_channels))
|
|
125
|
+
if local_attn:
|
|
126
|
+
self.local_convs.append(
|
|
127
|
+
GATConv(inner_channels, hidden_channels, heads=heads,
|
|
128
|
+
concat=True, add_self_loops=False, bias=False)
|
|
129
|
+
)
|
|
130
|
+
else:
|
|
131
|
+
self.local_convs.append(
|
|
132
|
+
GCNConv(inner_channels, inner_channels, cached=False,
|
|
133
|
+
normalize=True)
|
|
134
|
+
)
|
|
135
|
+
self.lins.append(keras.layers.Dense(inner_channels))
|
|
136
|
+
self.lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
|
|
137
|
+
if pre_ln:
|
|
138
|
+
self.pre_lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
|
|
139
|
+
if post_bn:
|
|
140
|
+
self.post_bns.append(
|
|
141
|
+
keras.layers.BatchNormalization(
|
|
142
|
+
center=True, scale=True, momentum=0.9, epsilon=1e-5,
|
|
143
|
+
)
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
self.lin_in = keras.layers.Dense(inner_channels)
|
|
147
|
+
self.ln = keras.layers.LayerNormalization(epsilon=1e-5)
|
|
148
|
+
|
|
149
|
+
self.global_attn = [
|
|
150
|
+
PolynormerAttention(
|
|
151
|
+
channels=hidden_channels,
|
|
152
|
+
heads=heads,
|
|
153
|
+
head_channels=hidden_channels,
|
|
154
|
+
beta=beta,
|
|
155
|
+
dropout=global_dropout,
|
|
156
|
+
qk_shared=qk_shared,
|
|
157
|
+
)
|
|
158
|
+
for _ in range(global_layers)
|
|
159
|
+
]
|
|
160
|
+
|
|
161
|
+
self.pred_local = keras.layers.Dense(out_channels)
|
|
162
|
+
self.pred_global = keras.layers.Dense(out_channels)
|
|
163
|
+
|
|
164
|
+
self._in_dropout = keras.layers.Dropout(in_dropout)
|
|
165
|
+
self._dropout = keras.layers.Dropout(dropout)
|
|
166
|
+
|
|
167
|
+
def build(self, input_shape=None):
|
|
168
|
+
self.built = True
|
|
169
|
+
|
|
170
|
+
def reset_parameters(self) -> None:
|
|
171
|
+
r"""Resets all learnable parameters of the module."""
|
|
172
|
+
# Keras layers reinitialize on next forward; no-op for unbuilt layers.
|
|
173
|
+
pass
|
|
174
|
+
|
|
175
|
+
def call(self, x, edge_index, batch: Optional[object] = None, training=None):
|
|
176
|
+
r"""Forward pass.
|
|
177
|
+
|
|
178
|
+
Args:
|
|
179
|
+
x (Tensor): The input node features.
|
|
180
|
+
edge_index (Tensor): The edge indices.
|
|
181
|
+
batch (Tensor, optional): The batch vector assigning each node to
|
|
182
|
+
a graph. (default: :obj:`None`)
|
|
183
|
+
training (bool, optional): Whether in training mode.
|
|
184
|
+
(default: :obj:`None`)
|
|
185
|
+
"""
|
|
186
|
+
x = self._in_dropout(x, training=training)
|
|
187
|
+
|
|
188
|
+
# ---- Equivariant local attention ----
|
|
189
|
+
x_local = 0
|
|
190
|
+
for i, local_conv in enumerate(self.local_convs):
|
|
191
|
+
if self.pre_ln:
|
|
192
|
+
x = self.pre_lns[i](x)
|
|
193
|
+
h = self.h_lins[i](x)
|
|
194
|
+
h = ops.relu(h)
|
|
195
|
+
x = local_conv(x, edge_index) + self.lins[i](x)
|
|
196
|
+
if self.post_bn:
|
|
197
|
+
x = self.post_bns[i](x, training=training)
|
|
198
|
+
x = ops.relu(x)
|
|
199
|
+
x = self._dropout(x, training=training)
|
|
200
|
+
x = (1 - self.beta) * self.lns[i](h * x) + self.beta * x
|
|
201
|
+
x_local = x_local + x
|
|
202
|
+
|
|
203
|
+
# ---- Equivariant global attention ----
|
|
204
|
+
if self._global:
|
|
205
|
+
# Sort nodes by batch assignment (required by to_dense_batch)
|
|
206
|
+
batch_i = ops.cast(batch, "int32")
|
|
207
|
+
indices = ops.argsort(batch_i)
|
|
208
|
+
rev_perm = ops.argsort(indices)
|
|
209
|
+
batch_sorted = ops.take(batch_i, indices, axis=0)
|
|
210
|
+
x_local_sorted = self.ln(ops.take(x_local, indices, axis=0))
|
|
211
|
+
|
|
212
|
+
x_global, mask = to_dense_batch(x_local_sorted, batch_sorted)
|
|
213
|
+
for attn in self.global_attn:
|
|
214
|
+
x_global = attn(x_global, mask=mask, training=training)
|
|
215
|
+
|
|
216
|
+
# Flatten and undo the sort
|
|
217
|
+
x = ops.take(from_dense_batch(x_global, batch_sorted), rev_perm, axis=0)
|
|
218
|
+
x = self.pred_global(x)
|
|
219
|
+
else:
|
|
220
|
+
x = self.pred_local(x_local)
|
|
221
|
+
|
|
222
|
+
return ops.log_softmax(x, axis=-1)
|
|
223
|
+
|
|
224
|
+
def __repr__(self) -> str:
|
|
225
|
+
return (f'{self.__class__.__name__}('
|
|
226
|
+
f'in_channels={self.hidden_channels}, '
|
|
227
|
+
f'hidden_channels={self.hidden_channels}, '
|
|
228
|
+
f'heads={self.heads})')
|
|
229
|
+
|
k3_node/models/rect.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv import GCNConv
|
|
7
|
+
from k3_node.layers.conv.utils import scatter
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class RECT_L(keras.layers.Layer):
|
|
11
|
+
r"""The RECT model, *i.e.* its supervised RECT-L part, from the
|
|
12
|
+
`"Network Embedding with Completely-imbalanced Labels"
|
|
13
|
+
<https://arxiv.org/abs/2007.03545>`_ paper.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
in_channels (int): Size of each input sample.
|
|
17
|
+
hidden_channels (int): Intermediate size of each sample.
|
|
18
|
+
normalize (bool, optional): Whether to add self-loops and compute
|
|
19
|
+
symmetric normalization coefficients on-the-fly.
|
|
20
|
+
(default: :obj:`True`)
|
|
21
|
+
dropout (float, optional): The dropout probability.
|
|
22
|
+
(default: :obj:`0.0`)
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```python
|
|
26
|
+
import numpy as np
|
|
27
|
+
from k3_node.models import RECT_L
|
|
28
|
+
|
|
29
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
30
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
31
|
+
|
|
32
|
+
model = RECT_L(in_channels=8, hidden_channels=16)
|
|
33
|
+
out = model(x, edge_index) # reconstructs the (semantic) input features
|
|
34
|
+
print(tuple(out.shape)) # (10, 8)
|
|
35
|
+
print(tuple(model.embed(x, edge_index).shape)) # (10, 16): node embeddings
|
|
36
|
+
```
|
|
37
|
+
"""
|
|
38
|
+
def __init__(
|
|
39
|
+
self,
|
|
40
|
+
in_channels: int,
|
|
41
|
+
hidden_channels: int,
|
|
42
|
+
normalize: bool = True,
|
|
43
|
+
dropout: float = 0.0,
|
|
44
|
+
**kwargs,
|
|
45
|
+
):
|
|
46
|
+
super().__init__(**kwargs)
|
|
47
|
+
self.in_channels = in_channels
|
|
48
|
+
self.hidden_channels = hidden_channels
|
|
49
|
+
self.dropout = dropout
|
|
50
|
+
|
|
51
|
+
self.conv = GCNConv(in_channels, hidden_channels, normalize=normalize)
|
|
52
|
+
self.lin = keras.layers.Dense(in_channels)
|
|
53
|
+
self._dropout = keras.layers.Dropout(dropout) if dropout > 0 else None
|
|
54
|
+
|
|
55
|
+
def build(self, input_shape=None):
|
|
56
|
+
self.built = True
|
|
57
|
+
|
|
58
|
+
def reset_parameters(self):
|
|
59
|
+
r"""Resets all learnable parameters of the module."""
|
|
60
|
+
self.conv.reset_parameters()
|
|
61
|
+
if self.lin.built:
|
|
62
|
+
self.lin.kernel.assign(
|
|
63
|
+
keras.initializers.GlorotUniform()(self.lin.kernel.shape)
|
|
64
|
+
)
|
|
65
|
+
if self.lin.bias is not None:
|
|
66
|
+
self.lin.bias.assign(ops.zeros(self.lin.bias.shape))
|
|
67
|
+
|
|
68
|
+
def call(self, x, edge_index, edge_weight=None, training=None):
|
|
69
|
+
x = self.conv(x, edge_index, edge_weight=edge_weight)
|
|
70
|
+
if self._dropout is not None:
|
|
71
|
+
x = self._dropout(x, training=training)
|
|
72
|
+
return self.lin(x)
|
|
73
|
+
|
|
74
|
+
def embed(self, x, edge_index, edge_weight=None):
|
|
75
|
+
return self.conv(x, edge_index, edge_weight=edge_weight)
|
|
76
|
+
|
|
77
|
+
def get_semantic_labels(self, x, y, mask):
|
|
78
|
+
r"""Replaces the original labels by their class-centers."""
|
|
79
|
+
mask_shape = ops.shape(mask)
|
|
80
|
+
if len(mask_shape) == 1 and 'bool' in str(mask.dtype):
|
|
81
|
+
y_sub = y[mask]
|
|
82
|
+
x_sub = x[mask]
|
|
83
|
+
else:
|
|
84
|
+
y_sub = ops.take(y, mask, axis=0)
|
|
85
|
+
x_sub = ops.take(x, mask, axis=0)
|
|
86
|
+
|
|
87
|
+
num_classes = int(ops.max(y_sub)) + 1
|
|
88
|
+
mean = scatter(x_sub, y_sub, dim=0, dim_size=num_classes, reduce='mean')
|
|
89
|
+
return ops.take(mean, y_sub, axis=0)
|
|
90
|
+
|
|
91
|
+
def __repr__(self) -> str:
|
|
92
|
+
return (f'{self.__class__.__name__}({self.in_channels}, '
|
|
93
|
+
f'{self.hidden_channels})')
|
k3_node/models/renet.py
ADDED
|
@@ -0,0 +1,221 @@
|
|
|
1
|
+
from typing import Callable, List, Optional, Tuple
|
|
2
|
+
import math
|
|
3
|
+
import numpy as np
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.layers.aggr import MeanAggregation
|
|
8
|
+
from k3_node.ops.creation import repeat
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class RENet(keras.Model):
|
|
12
|
+
r"""The Recurrent Event Network model from the `"Recurrent Event Network
|
|
13
|
+
for Reasoning over Temporal Knowledge Graphs"
|
|
14
|
+
<https://arxiv.org/abs/1904.05530>`_ paper.
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
num_nodes (int): The number of nodes in the knowledge graph.
|
|
18
|
+
num_rels (int): The number of relations in the knowledge graph.
|
|
19
|
+
hidden_channels (int): Hidden size of node and relation embeddings.
|
|
20
|
+
seq_len (int): The sequence length of past events.
|
|
21
|
+
num_layers (int, optional): The number of recurrent layers.
|
|
22
|
+
(default: :obj:`1`)
|
|
23
|
+
dropout (float, optional): Dropout rate before final prediction.
|
|
24
|
+
(default: :obj:`0.0`)
|
|
25
|
+
bias (bool, optional): If set to :obj:`False`, all layers will not
|
|
26
|
+
learn an additive bias. (default: :obj:`True`)
|
|
27
|
+
|
|
28
|
+
Example:
|
|
29
|
+
```python
|
|
30
|
+
import numpy as np
|
|
31
|
+
from k3_node.models import RENet
|
|
32
|
+
|
|
33
|
+
model = RENet(num_nodes=5, num_rels=4, hidden_channels=16, seq_len=3)
|
|
34
|
+
sub, rel, obj = np.array([0, 1]), np.array([0, 1]), np.array([2, 3]) # queries at time t
|
|
35
|
+
# Neighbor histories of subjects and objects: neighbor id, timestep and query index
|
|
36
|
+
h_sub, h_sub_t, h_sub_batch = np.array([0, 1, 2]), np.array([0, 1, 0]), np.array([0, 0, 1])
|
|
37
|
+
h_obj, h_obj_t, h_obj_batch = np.array([1, 2, 3]), np.array([1, 2, 0]), np.array([0, 0, 1])
|
|
38
|
+
log_prob_obj, log_prob_sub = model(sub, rel, obj, h_sub, h_sub_t, h_sub_batch, h_obj, h_obj_t, h_obj_batch)
|
|
39
|
+
print(tuple(log_prob_obj.shape)) # (2, 5): scores over all entities for each query
|
|
40
|
+
```
|
|
41
|
+
"""
|
|
42
|
+
def __init__(
|
|
43
|
+
self,
|
|
44
|
+
num_nodes: int,
|
|
45
|
+
num_rels: int,
|
|
46
|
+
hidden_channels: int,
|
|
47
|
+
seq_len: int,
|
|
48
|
+
num_layers: int = 1,
|
|
49
|
+
dropout: float = 0.0,
|
|
50
|
+
bias: bool = True,
|
|
51
|
+
**kwargs,
|
|
52
|
+
):
|
|
53
|
+
super().__init__(**kwargs)
|
|
54
|
+
|
|
55
|
+
self.num_nodes = num_nodes
|
|
56
|
+
self.num_rels = num_rels
|
|
57
|
+
self.hidden_channels = hidden_channels
|
|
58
|
+
self.seq_len = seq_len
|
|
59
|
+
self.dropout_rate = dropout
|
|
60
|
+
self.num_layers = num_layers
|
|
61
|
+
|
|
62
|
+
self.ent = self.add_weight(
|
|
63
|
+
name="ent",
|
|
64
|
+
shape=(num_nodes, hidden_channels),
|
|
65
|
+
initializer=keras.initializers.GlorotUniform(),
|
|
66
|
+
)
|
|
67
|
+
self.rel = self.add_weight(
|
|
68
|
+
name="rel",
|
|
69
|
+
shape=(num_rels, hidden_channels),
|
|
70
|
+
initializer=keras.initializers.GlorotUniform(),
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
self.sub_gru = keras.layers.GRU(
|
|
74
|
+
hidden_channels,
|
|
75
|
+
return_sequences=False,
|
|
76
|
+
use_bias=bias,
|
|
77
|
+
)
|
|
78
|
+
self.obj_gru = keras.layers.GRU(
|
|
79
|
+
hidden_channels,
|
|
80
|
+
return_sequences=False,
|
|
81
|
+
use_bias=bias,
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
self.sub_lin = keras.layers.Dense(num_nodes, use_bias=bias)
|
|
85
|
+
self.obj_lin = keras.layers.Dense(num_nodes, use_bias=bias)
|
|
86
|
+
self.drop = keras.layers.Dropout(dropout)
|
|
87
|
+
self.mean_aggr = MeanAggregation()
|
|
88
|
+
|
|
89
|
+
def reset_parameters(self):
|
|
90
|
+
self.ent.assign(
|
|
91
|
+
keras.initializers.GlorotUniform()(self.ent.shape)
|
|
92
|
+
)
|
|
93
|
+
self.rel.assign(
|
|
94
|
+
keras.initializers.GlorotUniform()(self.rel.shape)
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
def build(self, input_shape=None):
|
|
98
|
+
self.built = True
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def call(
|
|
102
|
+
self,
|
|
103
|
+
sub,
|
|
104
|
+
rel,
|
|
105
|
+
obj,
|
|
106
|
+
h_sub,
|
|
107
|
+
h_sub_t,
|
|
108
|
+
h_sub_batch,
|
|
109
|
+
h_obj,
|
|
110
|
+
h_obj_t,
|
|
111
|
+
h_obj_batch,
|
|
112
|
+
training=False,
|
|
113
|
+
):
|
|
114
|
+
batch_size = ops.shape(sub)[0]
|
|
115
|
+
seq_len = self.seq_len
|
|
116
|
+
|
|
117
|
+
h_sub_t = h_sub_t + h_sub_batch * seq_len
|
|
118
|
+
h_obj_t = h_obj_t + h_obj_batch * seq_len
|
|
119
|
+
|
|
120
|
+
ent_h_sub = ops.take(self.ent, h_sub, axis=0)
|
|
121
|
+
ent_h_obj = ops.take(self.ent, h_obj, axis=0)
|
|
122
|
+
|
|
123
|
+
h_sub_scatter = self.mean_aggr(
|
|
124
|
+
ent_h_sub, index=h_sub_t, dim_size=batch_size * seq_len, dim=0
|
|
125
|
+
)
|
|
126
|
+
h_sub = ops.reshape(h_sub_scatter, (-1, seq_len, self.hidden_channels)) # static feature size
|
|
127
|
+
|
|
128
|
+
h_obj_scatter = self.mean_aggr(
|
|
129
|
+
ent_h_obj, index=h_obj_t, dim_size=batch_size * seq_len, dim=0
|
|
130
|
+
)
|
|
131
|
+
h_obj = ops.reshape(h_obj_scatter, (-1, seq_len, self.hidden_channels)) # static feature size
|
|
132
|
+
|
|
133
|
+
sub_emb = ops.take(self.ent, sub, axis=0)
|
|
134
|
+
rel_emb = ops.take(self.rel, rel, axis=0)
|
|
135
|
+
obj_emb = ops.take(self.ent, obj, axis=0)
|
|
136
|
+
|
|
137
|
+
sub_rep = repeat(ops.expand_dims(sub_emb, 1), seq_len, axis=1)
|
|
138
|
+
rel_rep = repeat(ops.expand_dims(rel_emb, 1), seq_len, axis=1)
|
|
139
|
+
obj_rep = repeat(ops.expand_dims(obj_emb, 1), seq_len, axis=1)
|
|
140
|
+
|
|
141
|
+
gru_sub_in = ops.concatenate([sub_rep, h_sub, rel_rep], axis=-1)
|
|
142
|
+
gru_obj_in = ops.concatenate([obj_rep, h_obj, rel_rep], axis=-1)
|
|
143
|
+
|
|
144
|
+
h_sub = self.sub_gru(gru_sub_in, training=training)
|
|
145
|
+
h_obj = self.obj_gru(gru_obj_in, training=training)
|
|
146
|
+
|
|
147
|
+
h_sub = ops.concatenate([sub_emb, h_sub, rel_emb], axis=-1)
|
|
148
|
+
h_obj = ops.concatenate([obj_emb, h_obj, rel_emb], axis=-1)
|
|
149
|
+
|
|
150
|
+
h_sub = self.drop(h_sub, training=training)
|
|
151
|
+
h_obj = self.drop(h_obj, training=training)
|
|
152
|
+
|
|
153
|
+
log_prob_obj = ops.log_softmax(self.sub_lin(h_sub), axis=-1)
|
|
154
|
+
log_prob_sub = ops.log_softmax(self.obj_lin(h_obj), axis=-1)
|
|
155
|
+
|
|
156
|
+
return log_prob_obj, log_prob_sub
|
|
157
|
+
|
|
158
|
+
@staticmethod
|
|
159
|
+
def pre_transform(seq_len: int) -> Callable:
|
|
160
|
+
r"""Returns a pre-transform that adds to every event (processed in time order) the history
|
|
161
|
+
of its subject and object: the entities they were linked to by the same relation in each of
|
|
162
|
+
the last ``seq_len`` time steps (``h_sub`` / ``h_obj``, with the step in ``h_sub_t`` /
|
|
163
|
+
``h_obj_t``), as in PyG."""
|
|
164
|
+
|
|
165
|
+
class PreTransform:
|
|
166
|
+
def __init__(self, seq_len):
|
|
167
|
+
self.seq_len = seq_len
|
|
168
|
+
self.t_last = 0
|
|
169
|
+
self.sub_hist, self.obj_hist = {}, {} # node -> list of seq_len + 1 steps of (node, rel)
|
|
170
|
+
|
|
171
|
+
def _hist(self, hist, node):
|
|
172
|
+
if node not in hist:
|
|
173
|
+
hist[node] = [[] for _ in range(self.seq_len + 1)]
|
|
174
|
+
return hist[node]
|
|
175
|
+
|
|
176
|
+
def _history(self, hist, node, rel):
|
|
177
|
+
steps = self._hist(hist, node)
|
|
178
|
+
nodes, ts = [], []
|
|
179
|
+
for s in range(self.seq_len):
|
|
180
|
+
for other, r in steps[s]:
|
|
181
|
+
if r == rel:
|
|
182
|
+
nodes.append(other)
|
|
183
|
+
ts.append(s)
|
|
184
|
+
return np.array(nodes, dtype=np.int64), np.array(ts, dtype=np.int64)
|
|
185
|
+
|
|
186
|
+
def __call__(self, data):
|
|
187
|
+
sub, rel, obj, t = int(data.sub), int(data.rel), int(data.obj), int(data.t)
|
|
188
|
+
if t > self.t_last: # a new time step: forget the oldest one
|
|
189
|
+
for hist in (self.sub_hist, self.obj_hist):
|
|
190
|
+
for steps in hist.values():
|
|
191
|
+
steps.pop(0)
|
|
192
|
+
steps.append([])
|
|
193
|
+
self.t_last = t
|
|
194
|
+
data.h_sub, data.h_sub_t = self._history(self.sub_hist, sub, rel)
|
|
195
|
+
data.h_obj, data.h_obj_t = self._history(self.obj_hist, obj, rel)
|
|
196
|
+
self._hist(self.sub_hist, sub)[-1].append((obj, rel))
|
|
197
|
+
self._hist(self.obj_hist, obj)[-1].append((sub, rel))
|
|
198
|
+
return data
|
|
199
|
+
|
|
200
|
+
def __repr__(self):
|
|
201
|
+
return f"{self.__class__.__name__}(seq_len={self.seq_len})"
|
|
202
|
+
|
|
203
|
+
return PreTransform(seq_len)
|
|
204
|
+
|
|
205
|
+
def test(self, logits, y):
|
|
206
|
+
r"""Given ground-truth :obj:`y`, computes Mean Reciprocal Rank (MRR)
|
|
207
|
+
and Hits at 1/3/10.
|
|
208
|
+
"""
|
|
209
|
+
logits_np = ops.convert_to_numpy(logits)
|
|
210
|
+
y_np = ops.convert_to_numpy(y).reshape(-1, 1)
|
|
211
|
+
|
|
212
|
+
perm = np.argsort(-logits_np, axis=1)
|
|
213
|
+
mask = y_np == perm
|
|
214
|
+
|
|
215
|
+
rows, cols = np.nonzero(mask)
|
|
216
|
+
mrr = float(np.mean(1.0 / (cols + 1.0)))
|
|
217
|
+
hits1 = float(np.sum(cols < 1) / len(y_np))
|
|
218
|
+
hits3 = float(np.sum(cols < 3) / len(y_np))
|
|
219
|
+
hits10 = float(np.sum(cols < 10) / len(y_np))
|
|
220
|
+
|
|
221
|
+
return ops.convert_to_tensor([mrr, hits1, hits3, hits10], dtype="float32")
|