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,255 @@
|
|
|
1
|
+
from typing import Tuple
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.layers.kge.loader import KGTripletLoader
|
|
8
|
+
from k3_node.training import no_grad
|
|
9
|
+
|
|
10
|
+
try:
|
|
11
|
+
from tqdm import tqdm
|
|
12
|
+
except ImportError: # pragma: no cover
|
|
13
|
+
def tqdm(iterable, *args, **kwargs):
|
|
14
|
+
return iterable
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def normalize(x, p: float = 2.0, axis: int = -1, eps: float = 1e-12):
|
|
18
|
+
r"""Lp-normalizes `x` along `axis`, mirroring `torch.nn.functional.normalize`."""
|
|
19
|
+
norm = ops.power(ops.sum(ops.power(ops.abs(x), p), axis=axis, keepdims=True), 1.0 / p)
|
|
20
|
+
return x / ops.maximum(norm, eps)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def margin_ranking_loss(pos_score, neg_score, margin: float = 1.0):
|
|
24
|
+
r"""Mirrors `torch.nn.functional.margin_ranking_loss` with `target=1`."""
|
|
25
|
+
return ops.mean(ops.relu(margin - pos_score + neg_score))
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def binary_cross_entropy_with_logits(logits, target):
|
|
29
|
+
r"""Numerically stable sigmoid cross-entropy, mirroring
|
|
30
|
+
`torch.nn.functional.binary_cross_entropy_with_logits`."""
|
|
31
|
+
zeros = ops.zeros_like(logits)
|
|
32
|
+
cond = logits >= zeros
|
|
33
|
+
relu_logits = ops.where(cond, logits, zeros)
|
|
34
|
+
neg_abs_logits = ops.where(cond, -logits, logits)
|
|
35
|
+
loss = relu_logits - logits * target + ops.log(1.0 + ops.exp(neg_abs_logits))
|
|
36
|
+
return ops.mean(loss)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class KGEModel(keras.layers.Layer):
|
|
40
|
+
r"""An abstract base class for implementing custom KGE models.
|
|
41
|
+
|
|
42
|
+
Args:
|
|
43
|
+
num_nodes (int): The number of nodes/entities in the graph.
|
|
44
|
+
num_relations (int): The number of relations in the graph.
|
|
45
|
+
hidden_channels (int): The hidden embedding size.
|
|
46
|
+
sparse (bool, optional): Kept for API compatibility with PyG; has no
|
|
47
|
+
effect since Keras optimizers do not distinguish sparse
|
|
48
|
+
embedding gradients the way PyTorch does. (default: :obj:`False`)
|
|
49
|
+
"""
|
|
50
|
+
def __init__(
|
|
51
|
+
self,
|
|
52
|
+
num_nodes: int,
|
|
53
|
+
num_relations: int,
|
|
54
|
+
hidden_channels: int,
|
|
55
|
+
sparse: bool = False,
|
|
56
|
+
**kwargs,
|
|
57
|
+
):
|
|
58
|
+
super().__init__(**kwargs)
|
|
59
|
+
|
|
60
|
+
self.num_nodes = num_nodes
|
|
61
|
+
self.num_relations = num_relations
|
|
62
|
+
self.hidden_channels = hidden_channels
|
|
63
|
+
self.sparse = sparse
|
|
64
|
+
|
|
65
|
+
self.node_emb = keras.layers.Embedding(num_nodes, hidden_channels)
|
|
66
|
+
self.rel_emb = keras.layers.Embedding(num_relations, hidden_channels)
|
|
67
|
+
self.node_emb.build((None,))
|
|
68
|
+
self.rel_emb.build((None,))
|
|
69
|
+
|
|
70
|
+
def reset_parameters(self):
|
|
71
|
+
r"""Resets all learnable parameters of the module."""
|
|
72
|
+
self.node_emb.embeddings.assign(
|
|
73
|
+
self.node_emb.embeddings_initializer(ops.shape(self.node_emb.embeddings))
|
|
74
|
+
)
|
|
75
|
+
self.rel_emb.embeddings.assign(
|
|
76
|
+
self.rel_emb.embeddings_initializer(ops.shape(self.rel_emb.embeddings))
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
def call(self, head_index, rel_type, tail_index):
|
|
80
|
+
r"""Returns the score for the given triplet.
|
|
81
|
+
|
|
82
|
+
Args:
|
|
83
|
+
head_index: The head indices.
|
|
84
|
+
rel_type: The relation type.
|
|
85
|
+
tail_index: The tail indices.
|
|
86
|
+
"""
|
|
87
|
+
raise NotImplementedError
|
|
88
|
+
|
|
89
|
+
def loss(self, head_index, rel_type, tail_index):
|
|
90
|
+
r"""Returns the loss value for the given triplet."""
|
|
91
|
+
raise NotImplementedError
|
|
92
|
+
|
|
93
|
+
def loader(self, head_index, rel_type, tail_index, **kwargs):
|
|
94
|
+
r"""Returns a mini-batch loader that samples a subset of triplets.
|
|
95
|
+
|
|
96
|
+
Args:
|
|
97
|
+
head_index: The head indices.
|
|
98
|
+
rel_type: The relation type.
|
|
99
|
+
tail_index: The tail indices.
|
|
100
|
+
**kwargs (optional): Additional arguments of
|
|
101
|
+
:class:`k3_node.layers.kge.KGTripletLoader`, such as
|
|
102
|
+
`batch_size`, `shuffle` or `drop_last`.
|
|
103
|
+
"""
|
|
104
|
+
return KGTripletLoader(head_index, rel_type, tail_index, **kwargs)
|
|
105
|
+
|
|
106
|
+
def test(
|
|
107
|
+
self,
|
|
108
|
+
head_index,
|
|
109
|
+
rel_type,
|
|
110
|
+
tail_index,
|
|
111
|
+
batch_size: int,
|
|
112
|
+
k: int = 10,
|
|
113
|
+
log: bool = True,
|
|
114
|
+
) -> Tuple[float, float, float]:
|
|
115
|
+
r"""Evaluates the model quality by computing Mean Rank, MRR and
|
|
116
|
+
Hits@:math:`k` across all possible tail entities.
|
|
117
|
+
|
|
118
|
+
Args:
|
|
119
|
+
head_index: The head indices.
|
|
120
|
+
rel_type: The relation type.
|
|
121
|
+
tail_index: The tail indices.
|
|
122
|
+
batch_size (int): The batch size to use for evaluating.
|
|
123
|
+
k (int, optional): The :math:`k` in Hits @ :math:`k`.
|
|
124
|
+
(default: :obj:`10`)
|
|
125
|
+
log (bool, optional): If set to :obj:`False`, will not print a
|
|
126
|
+
progress bar to the console. (default: :obj:`True`)
|
|
127
|
+
"""
|
|
128
|
+
head_index = ops.convert_to_numpy(head_index)
|
|
129
|
+
rel_type = ops.convert_to_numpy(rel_type)
|
|
130
|
+
tail_index = ops.convert_to_numpy(tail_index)
|
|
131
|
+
|
|
132
|
+
arange = range(head_index.shape[0])
|
|
133
|
+
arange = tqdm(arange) if log else arange
|
|
134
|
+
|
|
135
|
+
mean_ranks, reciprocal_ranks, hits_at_k = [], [], []
|
|
136
|
+
for i in arange:
|
|
137
|
+
h, r, t = int(head_index[i]), int(rel_type[i]), int(tail_index[i])
|
|
138
|
+
|
|
139
|
+
scores = []
|
|
140
|
+
tail_indices = np.arange(self.num_nodes)
|
|
141
|
+
for start in range(0, self.num_nodes, batch_size):
|
|
142
|
+
ts = tail_indices[start:start + batch_size]
|
|
143
|
+
hs = np.full_like(ts, h)
|
|
144
|
+
rs = np.full_like(ts, r)
|
|
145
|
+
with no_grad(): # as PyG's @torch.no_grad()
|
|
146
|
+
out = self(
|
|
147
|
+
ops.convert_to_tensor(hs),
|
|
148
|
+
ops.convert_to_tensor(rs),
|
|
149
|
+
ops.convert_to_tensor(ts),
|
|
150
|
+
)
|
|
151
|
+
scores.append(ops.convert_to_numpy(out))
|
|
152
|
+
scores = np.concatenate(scores)
|
|
153
|
+
rank = int(np.nonzero(np.argsort(-scores) == t)[0][0])
|
|
154
|
+
|
|
155
|
+
mean_ranks.append(rank)
|
|
156
|
+
reciprocal_ranks.append(1.0 / (rank + 1))
|
|
157
|
+
hits_at_k.append(rank < k)
|
|
158
|
+
|
|
159
|
+
mean_rank = float(np.mean(mean_ranks))
|
|
160
|
+
mrr = float(np.mean(reciprocal_ranks))
|
|
161
|
+
hits_at_k = float(np.mean(hits_at_k))
|
|
162
|
+
|
|
163
|
+
return mean_rank, mrr, hits_at_k
|
|
164
|
+
|
|
165
|
+
# ---- Keras-style training -------------------------------------------------------------------
|
|
166
|
+
def compile(self, optimizer):
|
|
167
|
+
r"""Sets the optimizer used by :meth:`fit`."""
|
|
168
|
+
self.optimizer = optimizer
|
|
169
|
+
|
|
170
|
+
def fit(self, data, epochs: int = 1, batch_size: int = 1000, validation_data=None,
|
|
171
|
+
validation_batch_size: int = 20000, verbose: int = 1):
|
|
172
|
+
r"""Trains on the triplets of ``data`` (``edge_index`` holds heads and tails, ``edge_type``
|
|
173
|
+
the relations), in shuffled mini-batches.
|
|
174
|
+
|
|
175
|
+
Args:
|
|
176
|
+
data (Data): The training knowledge graph.
|
|
177
|
+
epochs (int): Passes over all triplets. (default: ``1``)
|
|
178
|
+
batch_size (int): Triplets per gradient step. (default: ``1000``)
|
|
179
|
+
validation_data (Data, optional): Evaluated with :meth:`evaluate` after every epoch
|
|
180
|
+
(this ranks all entities for every triplet, so keep it small).
|
|
181
|
+
validation_batch_size (int): Entities scored at once during validation.
|
|
182
|
+
verbose (int): ``0`` is silent, otherwise one line is printed per epoch.
|
|
183
|
+
|
|
184
|
+
Returns:
|
|
185
|
+
dict: The mean loss (and validation metrics) of every epoch.
|
|
186
|
+
"""
|
|
187
|
+
from k3_node.training import gradient_step
|
|
188
|
+
|
|
189
|
+
if getattr(self, "optimizer", None) is None:
|
|
190
|
+
raise ValueError("Call `compile(optimizer=...)` before `fit`.")
|
|
191
|
+
head, tail = data.edge_index[0], data.edge_index[1]
|
|
192
|
+
loader = self.loader(head, data.edge_type, tail, batch_size=batch_size, shuffle=True)
|
|
193
|
+
history = {"loss": []}
|
|
194
|
+
for epoch in range(1, epochs + 1):
|
|
195
|
+
total = count = 0
|
|
196
|
+
for h, r, t in loader:
|
|
197
|
+
loss = gradient_step(lambda: self.loss(h, r, t), self.trainable_variables, self.optimizer)
|
|
198
|
+
total += loss * int(h.shape[0])
|
|
199
|
+
count += int(h.shape[0])
|
|
200
|
+
logs = {"loss": total / count}
|
|
201
|
+
if validation_data is not None:
|
|
202
|
+
logs.update({f"val_{k}": v for k, v in
|
|
203
|
+
self.evaluate(validation_data, batch_size=validation_batch_size).items()})
|
|
204
|
+
for key, value in logs.items():
|
|
205
|
+
history.setdefault(key, []).append(value)
|
|
206
|
+
if verbose:
|
|
207
|
+
print(f"Epoch {epoch:03d}: " + ", ".join(f"{k}: {v:.4f}" for k, v in logs.items()))
|
|
208
|
+
return history
|
|
209
|
+
|
|
210
|
+
def evaluate(self, data, batch_size: int = 20000, k: int = 10):
|
|
211
|
+
r"""Ranks the true tail of every triplet in ``data`` among all entities and returns the
|
|
212
|
+
mean rank, the mean reciprocal rank (MRR) and Hits@``k``."""
|
|
213
|
+
mean_rank, mrr, hits = self.test(data.edge_index[0], data.edge_type, data.edge_index[1],
|
|
214
|
+
batch_size=batch_size, k=k, log=False)
|
|
215
|
+
return {"mean_rank": mean_rank, "mrr": mrr, f"hits@{k}": hits}
|
|
216
|
+
|
|
217
|
+
def random_sample(
|
|
218
|
+
self,
|
|
219
|
+
head_index,
|
|
220
|
+
rel_type,
|
|
221
|
+
tail_index,
|
|
222
|
+
):
|
|
223
|
+
r"""Randomly samples negative triplets by either replacing the head or
|
|
224
|
+
the tail (but not both).
|
|
225
|
+
|
|
226
|
+
Args:
|
|
227
|
+
head_index: The head indices.
|
|
228
|
+
rel_type: The relation type.
|
|
229
|
+
tail_index: The tail indices.
|
|
230
|
+
"""
|
|
231
|
+
num_triplets = ops.shape(head_index)[0]
|
|
232
|
+
num_negatives = num_triplets // 2
|
|
233
|
+
|
|
234
|
+
rnd_index = keras.random.randint(
|
|
235
|
+
ops.shape(head_index), 0, self.num_nodes, dtype="int32"
|
|
236
|
+
)
|
|
237
|
+
rnd_index = ops.cast(rnd_index, head_index.dtype)
|
|
238
|
+
|
|
239
|
+
head_index = ops.concatenate(
|
|
240
|
+
[rnd_index[:num_negatives], head_index[num_negatives:]], axis=0
|
|
241
|
+
)
|
|
242
|
+
tail_index = ops.concatenate(
|
|
243
|
+
[tail_index[:num_negatives], rnd_index[num_negatives:]], axis=0
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
return head_index, rel_type, tail_index
|
|
247
|
+
|
|
248
|
+
def __repr__(self) -> str:
|
|
249
|
+
return (
|
|
250
|
+
f"{self.__class__.__name__}({self.num_nodes}, "
|
|
251
|
+
f"num_relations={self.num_relations}, "
|
|
252
|
+
f"hidden_channels={self.hidden_channels})"
|
|
253
|
+
)
|
|
254
|
+
|
|
255
|
+
__str__ = __repr__
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.kge.base import KGEModel, binary_cross_entropy_with_logits
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def triple_dot(x, y, z):
|
|
8
|
+
return ops.sum(x * y * z, axis=-1)
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class ComplEx(KGEModel):
|
|
12
|
+
r"""The ComplEx model from the `"Complex Embeddings for Simple Link
|
|
13
|
+
Prediction" <https://arxiv.org/abs/1606.06357>`_ paper.
|
|
14
|
+
|
|
15
|
+
:class:`ComplEx` models relations as complex-valued bilinear mappings
|
|
16
|
+
between head and tail entities using the Hermetian dot product.
|
|
17
|
+
The entities and relations are embedded in different dimensional spaces,
|
|
18
|
+
resulting in the scoring function:
|
|
19
|
+
|
|
20
|
+
.. math::
|
|
21
|
+
d(h, r, t) = Re(< \mathbf{e}_h, \mathbf{e}_r, \mathbf{e}_t>)
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
num_nodes (int): The number of nodes/entities in the graph.
|
|
25
|
+
num_relations (int): The number of relations in the graph.
|
|
26
|
+
hidden_channels (int): The hidden embedding size.
|
|
27
|
+
sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
|
|
28
|
+
|
|
29
|
+
Example:
|
|
30
|
+
```python
|
|
31
|
+
import numpy as np
|
|
32
|
+
from k3_node.layers import ComplEx
|
|
33
|
+
|
|
34
|
+
head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
|
|
35
|
+
rel = np.random.randint(0, 5, size=(10,))
|
|
36
|
+
tail = np.random.randint(0, 20, size=(10,))
|
|
37
|
+
|
|
38
|
+
model = ComplEx(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
39
|
+
score = model(head, rel, tail) # plausibility score of every triple
|
|
40
|
+
print(tuple(score.shape)) # (10,)
|
|
41
|
+
loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
|
|
42
|
+
print(tuple(loss.shape)) # (): a scalar
|
|
43
|
+
```
|
|
44
|
+
"""
|
|
45
|
+
def __init__(
|
|
46
|
+
self,
|
|
47
|
+
num_nodes: int,
|
|
48
|
+
num_relations: int,
|
|
49
|
+
hidden_channels: int,
|
|
50
|
+
sparse: bool = False,
|
|
51
|
+
**kwargs,
|
|
52
|
+
):
|
|
53
|
+
super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
|
|
54
|
+
|
|
55
|
+
self.node_emb_im = keras.layers.Embedding(num_nodes, hidden_channels)
|
|
56
|
+
self.rel_emb_im = keras.layers.Embedding(num_relations, hidden_channels)
|
|
57
|
+
self.node_emb_im.build((None,))
|
|
58
|
+
self.rel_emb_im.build((None,))
|
|
59
|
+
|
|
60
|
+
self.reset_parameters()
|
|
61
|
+
|
|
62
|
+
def reset_parameters(self):
|
|
63
|
+
# A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
|
|
64
|
+
glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
|
|
65
|
+
self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
|
|
66
|
+
self.node_emb_im.embeddings.assign(glorot(ops.shape(self.node_emb_im.embeddings)))
|
|
67
|
+
self.rel_emb.embeddings.assign(glorot(ops.shape(self.rel_emb.embeddings)))
|
|
68
|
+
self.rel_emb_im.embeddings.assign(glorot(ops.shape(self.rel_emb_im.embeddings)))
|
|
69
|
+
|
|
70
|
+
def call(self, head_index, rel_type, tail_index):
|
|
71
|
+
head_index = ops.cast(head_index, "int32")
|
|
72
|
+
rel_type = ops.cast(rel_type, "int32")
|
|
73
|
+
tail_index = ops.cast(tail_index, "int32")
|
|
74
|
+
|
|
75
|
+
head_re = self.node_emb(head_index)
|
|
76
|
+
head_im = self.node_emb_im(head_index)
|
|
77
|
+
rel_re = self.rel_emb(rel_type)
|
|
78
|
+
rel_im = self.rel_emb_im(rel_type)
|
|
79
|
+
tail_re = self.node_emb(tail_index)
|
|
80
|
+
tail_im = self.node_emb_im(tail_index)
|
|
81
|
+
|
|
82
|
+
return (
|
|
83
|
+
triple_dot(head_re, rel_re, tail_re)
|
|
84
|
+
+ triple_dot(head_im, rel_re, tail_im)
|
|
85
|
+
+ triple_dot(head_re, rel_im, tail_im)
|
|
86
|
+
- triple_dot(head_im, rel_im, tail_re)
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
def loss(self, head_index, rel_type, tail_index):
|
|
90
|
+
pos_score = self(head_index, rel_type, tail_index)
|
|
91
|
+
neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
|
|
92
|
+
scores = ops.concatenate([pos_score, neg_score], axis=0)
|
|
93
|
+
|
|
94
|
+
pos_target = ops.ones_like(pos_score)
|
|
95
|
+
neg_target = ops.zeros_like(neg_score)
|
|
96
|
+
target = ops.concatenate([pos_target, neg_target], axis=0)
|
|
97
|
+
|
|
98
|
+
return binary_cross_entropy_with_logits(scores, target)
|
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.kge.base import KGEModel, margin_ranking_loss
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class DistMult(KGEModel):
|
|
8
|
+
r"""The DistMult model from the `"Embedding Entities and Relations for
|
|
9
|
+
Learning and Inference in Knowledge Bases"
|
|
10
|
+
<https://arxiv.org/abs/1412.6575>`_ paper.
|
|
11
|
+
|
|
12
|
+
:class:`DistMult` models relations as diagonal matrices, which simplifies
|
|
13
|
+
the bi-linear interaction between the head and tail entities to the score
|
|
14
|
+
function:
|
|
15
|
+
|
|
16
|
+
.. math::
|
|
17
|
+
d(h, r, t) = < \mathbf{e}_h, \mathbf{e}_r, \mathbf{e}_t >
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
num_nodes (int): The number of nodes/entities in the graph.
|
|
21
|
+
num_relations (int): The number of relations in the graph.
|
|
22
|
+
hidden_channels (int): The hidden embedding size.
|
|
23
|
+
margin (float, optional): The margin of the ranking loss.
|
|
24
|
+
(default: :obj:`1.0`)
|
|
25
|
+
sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
|
|
26
|
+
|
|
27
|
+
Example:
|
|
28
|
+
```python
|
|
29
|
+
import numpy as np
|
|
30
|
+
from k3_node.layers import DistMult
|
|
31
|
+
|
|
32
|
+
head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
|
|
33
|
+
rel = np.random.randint(0, 5, size=(10,))
|
|
34
|
+
tail = np.random.randint(0, 20, size=(10,))
|
|
35
|
+
|
|
36
|
+
model = DistMult(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
37
|
+
score = model(head, rel, tail) # plausibility score of every triple
|
|
38
|
+
print(tuple(score.shape)) # (10,)
|
|
39
|
+
loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
|
|
40
|
+
print(tuple(loss.shape)) # (): a scalar
|
|
41
|
+
```
|
|
42
|
+
"""
|
|
43
|
+
def __init__(
|
|
44
|
+
self,
|
|
45
|
+
num_nodes: int,
|
|
46
|
+
num_relations: int,
|
|
47
|
+
hidden_channels: int,
|
|
48
|
+
margin: float = 1.0,
|
|
49
|
+
sparse: bool = False,
|
|
50
|
+
**kwargs,
|
|
51
|
+
):
|
|
52
|
+
super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
|
|
53
|
+
|
|
54
|
+
self.margin = margin
|
|
55
|
+
|
|
56
|
+
self.reset_parameters()
|
|
57
|
+
|
|
58
|
+
def reset_parameters(self):
|
|
59
|
+
# A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
|
|
60
|
+
glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
|
|
61
|
+
self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
|
|
62
|
+
self.rel_emb.embeddings.assign(glorot(ops.shape(self.rel_emb.embeddings)))
|
|
63
|
+
|
|
64
|
+
def call(self, head_index, rel_type, tail_index):
|
|
65
|
+
head_index = ops.cast(head_index, "int32")
|
|
66
|
+
rel_type = ops.cast(rel_type, "int32")
|
|
67
|
+
tail_index = ops.cast(tail_index, "int32")
|
|
68
|
+
|
|
69
|
+
head = self.node_emb(head_index)
|
|
70
|
+
rel = self.rel_emb(rel_type)
|
|
71
|
+
tail = self.node_emb(tail_index)
|
|
72
|
+
|
|
73
|
+
return ops.sum(head * rel * tail, axis=-1)
|
|
74
|
+
|
|
75
|
+
def loss(self, head_index, rel_type, tail_index):
|
|
76
|
+
pos_score = self(head_index, rel_type, tail_index)
|
|
77
|
+
neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
|
|
78
|
+
|
|
79
|
+
return margin_ranking_loss(pos_score, neg_score, margin=self.margin)
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
import math
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class KGTripletLoader:
|
|
8
|
+
r"""A minimal, framework-agnostic batching iterator over knowledge-graph
|
|
9
|
+
triplets, mirroring :class:`torch.utils.data.DataLoader` usage in
|
|
10
|
+
:meth:`k3_node.layers.kge.KGEModel.loader`.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
head_index: The head indices.
|
|
14
|
+
rel_type: The relation type.
|
|
15
|
+
tail_index: The tail indices.
|
|
16
|
+
batch_size (int, optional): The batch size. (default: :obj:`1`)
|
|
17
|
+
shuffle (bool, optional): If set to :obj:`True`, shuffles the
|
|
18
|
+
triplets at every epoch. (default: :obj:`False`)
|
|
19
|
+
drop_last (bool, optional): If set to :obj:`True`, drops the last
|
|
20
|
+
incomplete batch. (default: :obj:`False`)
|
|
21
|
+
"""
|
|
22
|
+
def __init__(self, head_index, rel_type, tail_index, batch_size: int = 1,
|
|
23
|
+
shuffle: bool = False, drop_last: bool = False, **kwargs):
|
|
24
|
+
self.head_index = ops.convert_to_numpy(head_index)
|
|
25
|
+
self.rel_type = ops.convert_to_numpy(rel_type)
|
|
26
|
+
self.tail_index = ops.convert_to_numpy(tail_index)
|
|
27
|
+
self.batch_size = batch_size
|
|
28
|
+
self.shuffle = shuffle
|
|
29
|
+
self.drop_last = drop_last
|
|
30
|
+
self.num_triplets = self.head_index.shape[0]
|
|
31
|
+
|
|
32
|
+
def __len__(self):
|
|
33
|
+
if self.drop_last:
|
|
34
|
+
return self.num_triplets // self.batch_size
|
|
35
|
+
return math.ceil(self.num_triplets / self.batch_size)
|
|
36
|
+
|
|
37
|
+
def __iter__(self):
|
|
38
|
+
indices = np.arange(self.num_triplets)
|
|
39
|
+
if self.shuffle:
|
|
40
|
+
np.random.shuffle(indices)
|
|
41
|
+
|
|
42
|
+
for start in range(0, self.num_triplets, self.batch_size):
|
|
43
|
+
batch_idx = indices[start:start + self.batch_size]
|
|
44
|
+
if self.drop_last and batch_idx.shape[0] < self.batch_size:
|
|
45
|
+
continue
|
|
46
|
+
yield (
|
|
47
|
+
ops.convert_to_tensor(self.head_index[batch_idx]),
|
|
48
|
+
ops.convert_to_tensor(self.rel_type[batch_idx]),
|
|
49
|
+
ops.convert_to_tensor(self.tail_index[batch_idx]),
|
|
50
|
+
)
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
import math
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.kge.base import KGEModel, binary_cross_entropy_with_logits
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class RotatE(KGEModel):
|
|
10
|
+
r"""The RotatE model from the `"RotatE: Knowledge Graph Embedding by
|
|
11
|
+
Relational Rotation in Complex Space" <https://arxiv.org/abs/
|
|
12
|
+
1902.10197>`_ paper.
|
|
13
|
+
|
|
14
|
+
:class:`RotatE` models relations as a rotation in complex space
|
|
15
|
+
from head to tail such that
|
|
16
|
+
|
|
17
|
+
.. math::
|
|
18
|
+
\mathbf{e}_t = \mathbf{e}_h \circ \mathbf{e}_r,
|
|
19
|
+
|
|
20
|
+
resulting in the scoring function
|
|
21
|
+
|
|
22
|
+
.. math::
|
|
23
|
+
d(h, r, t) = - {\| \mathbf{e}_h \circ \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
|
+
sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
|
|
32
|
+
|
|
33
|
+
Example:
|
|
34
|
+
```python
|
|
35
|
+
import numpy as np
|
|
36
|
+
from k3_node.layers import RotatE
|
|
37
|
+
|
|
38
|
+
head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
|
|
39
|
+
rel = np.random.randint(0, 5, size=(10,))
|
|
40
|
+
tail = np.random.randint(0, 20, size=(10,))
|
|
41
|
+
|
|
42
|
+
model = RotatE(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
43
|
+
score = model(head, rel, tail) # plausibility score of every triple
|
|
44
|
+
print(tuple(score.shape)) # (10,)
|
|
45
|
+
loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
|
|
46
|
+
print(tuple(loss.shape)) # (): a scalar
|
|
47
|
+
```
|
|
48
|
+
"""
|
|
49
|
+
def __init__(
|
|
50
|
+
self,
|
|
51
|
+
num_nodes: int,
|
|
52
|
+
num_relations: int,
|
|
53
|
+
hidden_channels: int,
|
|
54
|
+
margin: float = 1.0,
|
|
55
|
+
sparse: bool = False,
|
|
56
|
+
**kwargs,
|
|
57
|
+
):
|
|
58
|
+
super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
|
|
59
|
+
|
|
60
|
+
self.margin = margin
|
|
61
|
+
self.node_emb_im = keras.layers.Embedding(num_nodes, hidden_channels)
|
|
62
|
+
self.node_emb_im.build((None,))
|
|
63
|
+
|
|
64
|
+
self.reset_parameters()
|
|
65
|
+
|
|
66
|
+
def reset_parameters(self):
|
|
67
|
+
# A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
|
|
68
|
+
glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
|
|
69
|
+
self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
|
|
70
|
+
self.node_emb_im.embeddings.assign(glorot(ops.shape(self.node_emb_im.embeddings)))
|
|
71
|
+
uniform = keras.initializers.RandomUniform(0, 2 * math.pi)
|
|
72
|
+
self.rel_emb.embeddings.assign(uniform(ops.shape(self.rel_emb.embeddings)))
|
|
73
|
+
|
|
74
|
+
def call(self, head_index, rel_type, tail_index):
|
|
75
|
+
head_index = ops.cast(head_index, "int32")
|
|
76
|
+
rel_type = ops.cast(rel_type, "int32")
|
|
77
|
+
tail_index = ops.cast(tail_index, "int32")
|
|
78
|
+
|
|
79
|
+
head_re = self.node_emb(head_index)
|
|
80
|
+
head_im = self.node_emb_im(head_index)
|
|
81
|
+
tail_re = self.node_emb(tail_index)
|
|
82
|
+
tail_im = self.node_emb_im(tail_index)
|
|
83
|
+
|
|
84
|
+
rel_theta = self.rel_emb(rel_type)
|
|
85
|
+
rel_re, rel_im = ops.cos(rel_theta), ops.sin(rel_theta)
|
|
86
|
+
|
|
87
|
+
re_score = (rel_re * head_re - rel_im * head_im) - tail_re
|
|
88
|
+
im_score = (rel_re * head_im + rel_im * head_re) - tail_im
|
|
89
|
+
complex_score = ops.stack([re_score, im_score], axis=2)
|
|
90
|
+
score = ops.sqrt(ops.sum(ops.square(complex_score), axis=(1, 2)))
|
|
91
|
+
|
|
92
|
+
return self.margin - score
|
|
93
|
+
|
|
94
|
+
def loss(self, head_index, rel_type, tail_index):
|
|
95
|
+
pos_score = self(head_index, rel_type, tail_index)
|
|
96
|
+
neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
|
|
97
|
+
scores = ops.concatenate([pos_score, neg_score], axis=0)
|
|
98
|
+
|
|
99
|
+
pos_target = ops.ones_like(pos_score)
|
|
100
|
+
neg_target = ops.zeros_like(neg_score)
|
|
101
|
+
target = ops.concatenate([pos_target, neg_target], axis=0)
|
|
102
|
+
|
|
103
|
+
return binary_cross_entropy_with_logits(scores, target)
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.kge import TransE, DistMult, ComplEx, RotatE
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def _run_model(cls):
|
|
7
|
+
model = cls(num_nodes=10, num_relations=5, hidden_channels=32)
|
|
8
|
+
assert str(model) == f"{cls.__name__}(10, num_relations=5, hidden_channels=32)"
|
|
9
|
+
|
|
10
|
+
head_index = ops.convert_to_tensor([0, 2, 4, 6, 8])
|
|
11
|
+
rel_type = ops.convert_to_tensor([0, 1, 2, 3, 4])
|
|
12
|
+
tail_index = ops.convert_to_tensor([1, 3, 5, 7, 9])
|
|
13
|
+
|
|
14
|
+
loader = model.loader(head_index, rel_type, tail_index, batch_size=5)
|
|
15
|
+
for h, r, t in loader:
|
|
16
|
+
out = model(h, r, t)
|
|
17
|
+
assert ops.shape(out) == (5,)
|
|
18
|
+
|
|
19
|
+
loss = model.loss(h, r, t)
|
|
20
|
+
assert float(ops.convert_to_numpy(loss)) >= 0.0
|
|
21
|
+
|
|
22
|
+
mean_rank, mrr, hits = model.test(h, r, t, batch_size=5, log=False)
|
|
23
|
+
assert 0 <= mean_rank <= 10
|
|
24
|
+
assert 0 < mrr <= 1
|
|
25
|
+
assert hits == 1.0
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def test_transe():
|
|
29
|
+
_run_model(TransE)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def test_distmult():
|
|
33
|
+
_run_model(DistMult)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def test_complex():
|
|
37
|
+
_run_model(ComplEx)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def test_rotate():
|
|
41
|
+
_run_model(RotatE)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def test_complex_scoring():
|
|
45
|
+
model = ComplEx(num_nodes=5, num_relations=2, hidden_channels=1)
|
|
46
|
+
|
|
47
|
+
model.node_emb.embeddings.assign(
|
|
48
|
+
ops.convert_to_tensor([[2.0], [3.0], [5.0], [1.0], [2.0]], dtype="float32")
|
|
49
|
+
)
|
|
50
|
+
model.node_emb_im.embeddings.assign(
|
|
51
|
+
ops.convert_to_tensor([[4.0], [1.0], [3.0], [1.0], [2.0]], dtype="float32")
|
|
52
|
+
)
|
|
53
|
+
model.rel_emb.embeddings.assign(ops.convert_to_tensor([[2.0], [3.0]], dtype="float32"))
|
|
54
|
+
model.rel_emb_im.embeddings.assign(ops.convert_to_tensor([[3.0], [1.0]], dtype="float32"))
|
|
55
|
+
|
|
56
|
+
score = model(
|
|
57
|
+
ops.convert_to_tensor([1, 3]),
|
|
58
|
+
ops.convert_to_tensor([1, 0]),
|
|
59
|
+
ops.convert_to_tensor([2, 4]),
|
|
60
|
+
)
|
|
61
|
+
assert ops.convert_to_numpy(score).tolist() == [58.0, 8.0]
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def test_same_shape_embeddings_are_initialized_independently():
|
|
65
|
+
# A reused Keras 3 initializer instance repeats its values. For ComplEx, identical real and
|
|
66
|
+
# imaginary parts cancel the asymmetric term, collapsing the model to DistMult.
|
|
67
|
+
import numpy as np
|
|
68
|
+
|
|
69
|
+
complex_model = ComplEx(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
70
|
+
rotate_model = RotatE(num_nodes=20, num_relations=5, hidden_channels=8)
|
|
71
|
+
for a, b in [
|
|
72
|
+
(complex_model.node_emb, complex_model.node_emb_im),
|
|
73
|
+
(complex_model.rel_emb, complex_model.rel_emb_im),
|
|
74
|
+
(rotate_model.node_emb, rotate_model.node_emb_im),
|
|
75
|
+
]:
|
|
76
|
+
assert not np.allclose(ops.convert_to_numpy(a.embeddings), ops.convert_to_numpy(b.embeddings))
|