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/rag/encoders.py
ADDED
|
@@ -0,0 +1,312 @@
|
|
|
1
|
+
"""GNN and KGE Subgraph Encoders for GraphRAG and KG-LLM augmentation."""
|
|
2
|
+
|
|
3
|
+
from typing import List, Literal, Optional, Tuple, Union
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import layers, ops
|
|
7
|
+
|
|
8
|
+
from k3_node.layers.conv import GCNConv, RGCNConv, SAGEConv
|
|
9
|
+
from k3_node.layers.kge import KGEModel, TransE
|
|
10
|
+
from k3_node.rag.subgraph import SubgraphResult
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class RGCNSubGraphEncoder(keras.Model):
|
|
14
|
+
r"""Relational Graph Convolutional Network (RGCN) encoder for multi-relational subgraphs.
|
|
15
|
+
|
|
16
|
+
Processes subgraphs extracted from Knowledge Graphs with multiple relation types,
|
|
17
|
+
updating entity representations through relational message passing and pooling them
|
|
18
|
+
into a dense graph embedding.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
in_channels: Dimensionality of input node features.
|
|
22
|
+
hidden_channels: Hidden representation dimension.
|
|
23
|
+
out_channels: Output graph embedding dimension.
|
|
24
|
+
num_relations: Total number of relation types in the Knowledge Graph.
|
|
25
|
+
num_layers: Number of RGCN message passing layers. (default: 2)
|
|
26
|
+
num_bases: Optional number of basis decomposition components for relation weights.
|
|
27
|
+
pooling: Readout pooling strategy across subgraph nodes:
|
|
28
|
+
- `"mean"`: Global average over all subgraph nodes.
|
|
29
|
+
- `"sum"`: Global sum over all subgraph nodes.
|
|
30
|
+
- `"max"`: Global maximum over all subgraph nodes.
|
|
31
|
+
- `"center"`: Pool only the retrieved center seed entities.
|
|
32
|
+
- `"none"`: Return all node embeddings without pooling.
|
|
33
|
+
dropout: Dropout rate applied between convolution layers. (default: 0.0)
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
def __init__(
|
|
37
|
+
self,
|
|
38
|
+
in_channels: int,
|
|
39
|
+
hidden_channels: int,
|
|
40
|
+
out_channels: int,
|
|
41
|
+
num_relations: int,
|
|
42
|
+
num_layers: int = 2,
|
|
43
|
+
num_bases: Optional[int] = None,
|
|
44
|
+
pooling: Literal["mean", "sum", "max", "center", "none"] = "mean",
|
|
45
|
+
dropout: float = 0.0,
|
|
46
|
+
**kwargs,
|
|
47
|
+
):
|
|
48
|
+
super().__init__(**kwargs)
|
|
49
|
+
self.in_channels = in_channels
|
|
50
|
+
self.hidden_channels = hidden_channels
|
|
51
|
+
self.out_channels = out_channels
|
|
52
|
+
self.num_relations = num_relations
|
|
53
|
+
self.num_layers = num_layers
|
|
54
|
+
self.pooling = pooling
|
|
55
|
+
self.dropout_rate = dropout
|
|
56
|
+
|
|
57
|
+
self.convs = []
|
|
58
|
+
for i in range(num_layers):
|
|
59
|
+
c_in = in_channels if i == 0 else hidden_channels
|
|
60
|
+
c_out = out_channels if i == num_layers - 1 else hidden_channels
|
|
61
|
+
self.convs.append(
|
|
62
|
+
RGCNConv(
|
|
63
|
+
in_channels=c_in,
|
|
64
|
+
out_channels=c_out,
|
|
65
|
+
num_relations=num_relations,
|
|
66
|
+
num_bases=num_bases,
|
|
67
|
+
)
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
self.act = layers.Activation("relu")
|
|
71
|
+
self.drop = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
72
|
+
|
|
73
|
+
def call(
|
|
74
|
+
self,
|
|
75
|
+
x: any,
|
|
76
|
+
edge_index: any,
|
|
77
|
+
edge_type: Optional[any] = None,
|
|
78
|
+
center_nodes: Optional[any] = None,
|
|
79
|
+
training: Optional[bool] = None,
|
|
80
|
+
) -> any:
|
|
81
|
+
"""Forward pass encoding the subgraph into a graph embedding.
|
|
82
|
+
|
|
83
|
+
Args:
|
|
84
|
+
x: Node feature tensor of shape `(num_nodes, in_channels)`.
|
|
85
|
+
edge_index: Graph edge indices of shape `(2, num_edges)`.
|
|
86
|
+
edge_type: 1D tensor of relation IDs for each edge `(num_edges,)`.
|
|
87
|
+
center_nodes: Optional 1D tensor of seed entity indices in the subgraph.
|
|
88
|
+
training: Whether running in training mode.
|
|
89
|
+
|
|
90
|
+
Returns:
|
|
91
|
+
Tensor of shape `(1, out_channels)` (or `(num_nodes, out_channels)` if `pooling="none"`).
|
|
92
|
+
"""
|
|
93
|
+
h = x
|
|
94
|
+
for i, conv in enumerate(self.convs):
|
|
95
|
+
h = conv(h, edge_index, edge_type=edge_type)
|
|
96
|
+
if i < self.num_layers - 1:
|
|
97
|
+
h = self.act(h)
|
|
98
|
+
if self.drop is not None:
|
|
99
|
+
h = self.drop(h, training=training)
|
|
100
|
+
|
|
101
|
+
if self.pooling == "none":
|
|
102
|
+
return h
|
|
103
|
+
|
|
104
|
+
if self.pooling == "center" and center_nodes is not None:
|
|
105
|
+
center_nodes = ops.convert_to_tensor(center_nodes, dtype="int32")
|
|
106
|
+
if ops.shape(center_nodes)[0] > 0:
|
|
107
|
+
center_h = ops.take(h, center_nodes, axis=0)
|
|
108
|
+
pooled = ops.mean(center_h, axis=0, keepdims=True)
|
|
109
|
+
return pooled
|
|
110
|
+
|
|
111
|
+
if self.pooling == "sum":
|
|
112
|
+
return ops.sum(h, axis=0, keepdims=True)
|
|
113
|
+
elif self.pooling == "max":
|
|
114
|
+
return ops.max(h, axis=0, keepdims=True)
|
|
115
|
+
else: # "mean" or fallback
|
|
116
|
+
return ops.mean(h, axis=0, keepdims=True)
|
|
117
|
+
|
|
118
|
+
def encode_subgraph(self, subgraph: SubgraphResult, default_x_dim: Optional[int] = None) -> any:
|
|
119
|
+
"""Helper to encode a SubgraphResult directly."""
|
|
120
|
+
x = subgraph.x
|
|
121
|
+
if x is None:
|
|
122
|
+
dim = default_x_dim if default_x_dim is not None else self.in_channels
|
|
123
|
+
num_n = subgraph.num_nodes if subgraph.num_nodes > 0 else 1
|
|
124
|
+
x = ops.ones((num_n, dim), dtype="float32")
|
|
125
|
+
|
|
126
|
+
return self(
|
|
127
|
+
x=x,
|
|
128
|
+
edge_index=subgraph.edge_index,
|
|
129
|
+
edge_type=subgraph.edge_type,
|
|
130
|
+
center_nodes=subgraph.center_nodes,
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
class TransEPrefixEncoder(keras.Model):
|
|
135
|
+
r"""Knowledge Graph Embedding (KGE) prefix encoder using TransE representations.
|
|
136
|
+
|
|
137
|
+
Uses pretrained or end-to-end entity and relation embeddings from TransE ($h + r \approx t$)
|
|
138
|
+
to encode extracted knowledge subgraphs into a unified dense embedding vector.
|
|
139
|
+
|
|
140
|
+
Args:
|
|
141
|
+
num_nodes: Total number of entities in the knowledge graph.
|
|
142
|
+
num_relations: Total number of relation types.
|
|
143
|
+
embedding_dim: Dimension of entity and relation embeddings. (default: 64)
|
|
144
|
+
out_channels: Output projected graph embedding dimension. (default: 128)
|
|
145
|
+
kge_model: Optional pretrained `KGEModel` (e.g. `TransE`). If provided,
|
|
146
|
+
its embeddings are reused.
|
|
147
|
+
pooling: Readout pooling strategy across entities (`"mean"`, `"center"`, `"sum"`).
|
|
148
|
+
"""
|
|
149
|
+
|
|
150
|
+
def __init__(
|
|
151
|
+
self,
|
|
152
|
+
num_nodes: Optional[int] = None,
|
|
153
|
+
num_relations: Optional[int] = None,
|
|
154
|
+
embedding_dim: int = 64,
|
|
155
|
+
out_channels: int = 128,
|
|
156
|
+
kge_model: Optional[KGEModel] = None,
|
|
157
|
+
pooling: Literal["mean", "sum", "center"] = "mean",
|
|
158
|
+
**kwargs,
|
|
159
|
+
):
|
|
160
|
+
super().__init__(**kwargs)
|
|
161
|
+
self.embedding_dim = embedding_dim
|
|
162
|
+
self.out_channels = out_channels
|
|
163
|
+
self.pooling = pooling
|
|
164
|
+
|
|
165
|
+
if kge_model is not None:
|
|
166
|
+
self.node_emb = kge_model.node_emb
|
|
167
|
+
self.rel_emb = kge_model.rel_emb
|
|
168
|
+
self.embedding_dim = kge_model.hidden_channels
|
|
169
|
+
else:
|
|
170
|
+
if num_nodes is None or num_relations is None:
|
|
171
|
+
raise ValueError("Must provide either `kge_model` or both `num_nodes` and `num_relations`.")
|
|
172
|
+
self.node_emb = layers.Embedding(num_nodes, embedding_dim)
|
|
173
|
+
self.rel_emb = layers.Embedding(num_relations, embedding_dim)
|
|
174
|
+
self.node_emb.build((None,))
|
|
175
|
+
self.rel_emb.build((None,))
|
|
176
|
+
|
|
177
|
+
# Projection head combining entity representation and relational context
|
|
178
|
+
self.proj = keras.Sequential(
|
|
179
|
+
[
|
|
180
|
+
layers.Dense(out_channels, activation="relu"),
|
|
181
|
+
layers.LayerNormalization(),
|
|
182
|
+
layers.Dense(out_channels),
|
|
183
|
+
]
|
|
184
|
+
)
|
|
185
|
+
|
|
186
|
+
def call(
|
|
187
|
+
self,
|
|
188
|
+
subgraph_nodes: any,
|
|
189
|
+
edge_index: Optional[any] = None,
|
|
190
|
+
edge_type: Optional[any] = None,
|
|
191
|
+
center_nodes: Optional[any] = None,
|
|
192
|
+
) -> any:
|
|
193
|
+
"""Encode a set of subgraph entities and relations into a dense embedding.
|
|
194
|
+
|
|
195
|
+
Args:
|
|
196
|
+
subgraph_nodes: 1D tensor of original entity IDs in the subgraph `(num_nodes,)`.
|
|
197
|
+
edge_index: Optional edge index tensor `(2, num_edges)`.
|
|
198
|
+
edge_type: Optional relation type tensor `(num_edges,)`.
|
|
199
|
+
center_nodes: Optional indices of seed entities within `subgraph_nodes`.
|
|
200
|
+
|
|
201
|
+
Returns:
|
|
202
|
+
Dense embedding tensor of shape `(1, out_channels)`.
|
|
203
|
+
"""
|
|
204
|
+
subgraph_nodes = ops.convert_to_tensor(subgraph_nodes, dtype="int32")
|
|
205
|
+
entity_embeddings = self.node_emb(subgraph_nodes) # [num_nodes, embedding_dim]
|
|
206
|
+
|
|
207
|
+
# Entity pooling
|
|
208
|
+
if self.pooling == "center" and center_nodes is not None:
|
|
209
|
+
center_idx = ops.convert_to_tensor(center_nodes, dtype="int32")
|
|
210
|
+
if ops.shape(center_idx)[0] > 0:
|
|
211
|
+
ent_h = ops.take(entity_embeddings, center_idx, axis=0)
|
|
212
|
+
node_repr = ops.mean(ent_h, axis=0, keepdims=True)
|
|
213
|
+
else:
|
|
214
|
+
node_repr = ops.mean(entity_embeddings, axis=0, keepdims=True)
|
|
215
|
+
elif self.pooling == "sum":
|
|
216
|
+
node_repr = ops.sum(entity_embeddings, axis=0, keepdims=True)
|
|
217
|
+
else:
|
|
218
|
+
node_repr = ops.mean(entity_embeddings, axis=0, keepdims=True)
|
|
219
|
+
|
|
220
|
+
# Relational translation context (h + r - t in TransE)
|
|
221
|
+
if edge_index is not None and edge_type is not None:
|
|
222
|
+
edge_index = ops.convert_to_tensor(edge_index, dtype="int32")
|
|
223
|
+
edge_type = ops.convert_to_tensor(edge_type, dtype="int32")
|
|
224
|
+
num_e = ops.shape(edge_index)[1]
|
|
225
|
+
|
|
226
|
+
if num_e > 0:
|
|
227
|
+
row, col = edge_index[0], edge_index[1]
|
|
228
|
+
h_sub = ops.take(entity_embeddings, row, axis=0)
|
|
229
|
+
t_sub = ops.take(entity_embeddings, col, axis=0)
|
|
230
|
+
r_sub = self.rel_emb(edge_type)
|
|
231
|
+
# TransE relation representation
|
|
232
|
+
triple_context = ops.mean(h_sub + r_sub - t_sub, axis=0, keepdims=True)
|
|
233
|
+
combined = ops.concatenate([node_repr, triple_context], axis=-1)
|
|
234
|
+
else:
|
|
235
|
+
combined = ops.concatenate([node_repr, ops.zeros_like(node_repr)], axis=-1)
|
|
236
|
+
else:
|
|
237
|
+
combined = ops.concatenate([node_repr, ops.zeros_like(node_repr)], axis=-1)
|
|
238
|
+
|
|
239
|
+
return self.proj(combined)
|
|
240
|
+
|
|
241
|
+
def encode_subgraph(self, subgraph: SubgraphResult) -> any:
|
|
242
|
+
"""Helper to encode a SubgraphResult directly."""
|
|
243
|
+
nodes = subgraph.nodes if subgraph.nodes is not None else np.arange(subgraph.num_nodes)
|
|
244
|
+
return self(
|
|
245
|
+
subgraph_nodes=nodes,
|
|
246
|
+
edge_index=subgraph.edge_index,
|
|
247
|
+
edge_type=subgraph.edge_type,
|
|
248
|
+
center_nodes=subgraph.center_nodes,
|
|
249
|
+
)
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
class GNNSubGraphEncoder(keras.Model):
|
|
253
|
+
r"""General-purpose GNN encoder (GCN / GraphSAGE) for homogeneous subgraphs.
|
|
254
|
+
|
|
255
|
+
Args:
|
|
256
|
+
in_channels: Input node feature dimensionality.
|
|
257
|
+
hidden_channels: Hidden representation dimension.
|
|
258
|
+
out_channels: Output graph embedding dimension.
|
|
259
|
+
conv_type: Type of GNN convolution (`"gcn"` or `"sage"`). (default: `"gcn"`)
|
|
260
|
+
num_layers: Number of convolution layers. (default: 2)
|
|
261
|
+
pooling: Pooling strategy (`"mean"`, `"sum"`, `"max"`, `"center"`).
|
|
262
|
+
"""
|
|
263
|
+
|
|
264
|
+
def __init__(
|
|
265
|
+
self,
|
|
266
|
+
in_channels: int,
|
|
267
|
+
hidden_channels: int,
|
|
268
|
+
out_channels: int,
|
|
269
|
+
conv_type: Literal["gcn", "sage"] = "gcn",
|
|
270
|
+
num_layers: int = 2,
|
|
271
|
+
pooling: Literal["mean", "sum", "max", "center"] = "mean",
|
|
272
|
+
**kwargs,
|
|
273
|
+
):
|
|
274
|
+
super().__init__(**kwargs)
|
|
275
|
+
self.in_channels = in_channels
|
|
276
|
+
self.hidden_channels = hidden_channels
|
|
277
|
+
self.out_channels = out_channels
|
|
278
|
+
self.pooling = pooling
|
|
279
|
+
self.num_layers = num_layers
|
|
280
|
+
|
|
281
|
+
ConvCls = GCNConv if conv_type.lower() == "gcn" else SAGEConv
|
|
282
|
+
self.convs = []
|
|
283
|
+
for i in range(num_layers):
|
|
284
|
+
c_in = in_channels if i == 0 else hidden_channels
|
|
285
|
+
c_out = out_channels if i == num_layers - 1 else hidden_channels
|
|
286
|
+
self.convs.append(ConvCls(c_in, c_out))
|
|
287
|
+
|
|
288
|
+
self.act = layers.Activation("relu")
|
|
289
|
+
|
|
290
|
+
def call(
|
|
291
|
+
self,
|
|
292
|
+
x: any,
|
|
293
|
+
edge_index: any,
|
|
294
|
+
center_nodes: Optional[any] = None,
|
|
295
|
+
) -> any:
|
|
296
|
+
h = x
|
|
297
|
+
for i, conv in enumerate(self.convs):
|
|
298
|
+
h = conv(h, edge_index)
|
|
299
|
+
if i < self.num_layers - 1:
|
|
300
|
+
h = self.act(h)
|
|
301
|
+
|
|
302
|
+
if self.pooling == "center" and center_nodes is not None:
|
|
303
|
+
center_nodes = ops.convert_to_tensor(center_nodes, dtype="int32")
|
|
304
|
+
if ops.shape(center_nodes)[0] > 0:
|
|
305
|
+
return ops.mean(ops.take(h, center_nodes, axis=0), axis=0, keepdims=True)
|
|
306
|
+
|
|
307
|
+
if self.pooling == "sum":
|
|
308
|
+
return ops.sum(h, axis=0, keepdims=True)
|
|
309
|
+
elif self.pooling == "max":
|
|
310
|
+
return ops.max(h, axis=0, keepdims=True)
|
|
311
|
+
else:
|
|
312
|
+
return ops.mean(h, axis=0, keepdims=True)
|
k3_node/rag/pipeline.py
ADDED
|
@@ -0,0 +1,192 @@
|
|
|
1
|
+
"""High-level GraphRAG Pipeline integrating retrieval, verbalization, and LLM prefix injection."""
|
|
2
|
+
|
|
3
|
+
from typing import Dict, List, Literal, Optional, Sequence, Union
|
|
4
|
+
import numpy as np
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.rag.encoders import RGCNSubGraphEncoder, TransEPrefixEncoder
|
|
8
|
+
from k3_node.rag.projector import GraphPrefixProjector, KGLLMConnector
|
|
9
|
+
from k3_node.rag.subgraph import KGEntityRetriever, SubgraphResult, extract_subgraph
|
|
10
|
+
from k3_node.rag.verbalizer import format_llm_prompt, verbalize_subgraph
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class GraphRAG:
|
|
14
|
+
r"""End-to-End Graph-Augmented Generation (GraphRAG) Pipeline for Knowledge Graphs.
|
|
15
|
+
|
|
16
|
+
Provides a unified API for:
|
|
17
|
+
1. Extracting subgraphs around query entities.
|
|
18
|
+
2. Verbalizing structured facts into markdown or natural language prompts for Llama 3 / Mistral.
|
|
19
|
+
3. Projecting GNN (RGCN) or KGE (TransE) embeddings into dense prefix vectors for soft prompt augmentation.
|
|
20
|
+
|
|
21
|
+
Example:
|
|
22
|
+
```python
|
|
23
|
+
import k3_node as k3
|
|
24
|
+
from k3_node.rag import GraphRAG
|
|
25
|
+
|
|
26
|
+
# Setup GraphRAG with your knowledge graph
|
|
27
|
+
rag = GraphRAG(
|
|
28
|
+
edge_index=edge_index,
|
|
29
|
+
edge_type=edge_type,
|
|
30
|
+
entity_to_id={"Aspirin": 0, "Headache": 1, "COX-1": 2},
|
|
31
|
+
relation_to_id={"treats": 0, "inhibits": 1},
|
|
32
|
+
llm_dim=4096, # Llama 3 / Mistral embedding dimension
|
|
33
|
+
num_prefix_tokens=4,
|
|
34
|
+
)
|
|
35
|
+
|
|
36
|
+
# 1. Text-based GraphRAG prompt generation
|
|
37
|
+
prompt = rag.build_prompt(
|
|
38
|
+
query="How does Aspirin alleviate headache?",
|
|
39
|
+
entities=["Aspirin", "Headache"],
|
|
40
|
+
model_family="llama3",
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
# 2. Dense Prefix Vector encoding
|
|
44
|
+
subgraph = rag.retrieve(["Aspirin"], num_hops=2)
|
|
45
|
+
prefix_embeds = rag.encode_prefix(subgraph) # [1, 4, 4096]
|
|
46
|
+
```
|
|
47
|
+
|
|
48
|
+
Args:
|
|
49
|
+
edge_index: Graph connectivity tensor of shape `(2, num_edges)`.
|
|
50
|
+
edge_type: Optional relation IDs tensor of shape `(num_edges,)`.
|
|
51
|
+
entity_to_id: Dictionary mapping entity strings to node IDs.
|
|
52
|
+
relation_to_id: Optional dictionary mapping relation strings to relation IDs.
|
|
53
|
+
x: Optional node feature tensor.
|
|
54
|
+
num_relations: Optional total number of relation types.
|
|
55
|
+
encoder_type: Subgraph encoder architecture (`"rgcn"`, `"transe"`, or a custom encoder).
|
|
56
|
+
hidden_dim: Hidden dimension for GNN/KGE encoder. (default: 64)
|
|
57
|
+
encoder_out_dim: Output dimension of graph encoder before LLM projection. (default: 128)
|
|
58
|
+
llm_dim: Target LLM embedding dimension (e.g. 4096 for Llama 3 / Mistral). (default: 4096)
|
|
59
|
+
num_prefix_tokens: Number of virtual prefix tokens to generate. (default: 8)
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
def __init__(
|
|
63
|
+
self,
|
|
64
|
+
edge_index: any,
|
|
65
|
+
edge_type: Optional[any] = None,
|
|
66
|
+
entity_to_id: Optional[Dict[str, int]] = None,
|
|
67
|
+
relation_to_id: Optional[Dict[str, int]] = None,
|
|
68
|
+
x: Optional[any] = None,
|
|
69
|
+
num_relations: Optional[int] = None,
|
|
70
|
+
encoder_type: Literal["rgcn", "transe", "none"] = "rgcn",
|
|
71
|
+
hidden_dim: int = 64,
|
|
72
|
+
encoder_out_dim: int = 128,
|
|
73
|
+
llm_dim: int = 4096,
|
|
74
|
+
num_prefix_tokens: int = 8,
|
|
75
|
+
connector: Optional[KGLLMConnector] = None,
|
|
76
|
+
):
|
|
77
|
+
self.edge_index = edge_index
|
|
78
|
+
self.edge_type = edge_type
|
|
79
|
+
self.x = x
|
|
80
|
+
self.entity_to_id = entity_to_id or {}
|
|
81
|
+
self.relation_to_id = relation_to_id or {}
|
|
82
|
+
self.id_to_entity = {v: k for k, v in self.entity_to_id.items()}
|
|
83
|
+
self.id_to_relation = {v: k for k, v in self.relation_to_id.items()}
|
|
84
|
+
|
|
85
|
+
self.retriever = KGEntityRetriever(
|
|
86
|
+
entity_to_id=self.entity_to_id,
|
|
87
|
+
relation_to_id=self.relation_to_id,
|
|
88
|
+
edge_index=self.edge_index,
|
|
89
|
+
edge_type=self.edge_type,
|
|
90
|
+
x=self.x,
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
if connector is not None:
|
|
94
|
+
self.connector = connector
|
|
95
|
+
elif encoder_type == "none":
|
|
96
|
+
self.connector = None
|
|
97
|
+
else:
|
|
98
|
+
n_rel = num_relations
|
|
99
|
+
if n_rel is None and self.relation_to_id:
|
|
100
|
+
n_rel = len(self.relation_to_id)
|
|
101
|
+
elif n_rel is None and edge_type is not None:
|
|
102
|
+
n_rel = int(np.max(ops.convert_to_numpy(edge_type))) + 1
|
|
103
|
+
else:
|
|
104
|
+
n_rel = n_rel or 1
|
|
105
|
+
|
|
106
|
+
in_ch = ops.shape(x)[-1] if x is not None else hidden_dim
|
|
107
|
+
if encoder_type == "rgcn":
|
|
108
|
+
encoder = RGCNSubGraphEncoder(
|
|
109
|
+
in_channels=in_ch,
|
|
110
|
+
hidden_channels=hidden_dim,
|
|
111
|
+
out_channels=encoder_out_dim,
|
|
112
|
+
num_relations=n_rel,
|
|
113
|
+
)
|
|
114
|
+
elif encoder_type == "transe":
|
|
115
|
+
num_nodes = len(self.entity_to_id) if self.entity_to_id else int(ops.max(edge_index)) + 1
|
|
116
|
+
encoder = TransEPrefixEncoder(
|
|
117
|
+
num_nodes=num_nodes,
|
|
118
|
+
num_relations=n_rel,
|
|
119
|
+
embedding_dim=hidden_dim,
|
|
120
|
+
out_channels=encoder_out_dim,
|
|
121
|
+
)
|
|
122
|
+
else:
|
|
123
|
+
raise ValueError(f"Unknown encoder_type: {encoder_type}")
|
|
124
|
+
|
|
125
|
+
projector = GraphPrefixProjector(
|
|
126
|
+
in_channels=encoder_out_dim,
|
|
127
|
+
llm_dim=llm_dim,
|
|
128
|
+
num_prefix_tokens=num_prefix_tokens,
|
|
129
|
+
)
|
|
130
|
+
self.connector = KGLLMConnector(encoder=encoder, projector=projector)
|
|
131
|
+
|
|
132
|
+
def retrieve(
|
|
133
|
+
self,
|
|
134
|
+
entities: Union[Sequence[Union[str, int]], str, int],
|
|
135
|
+
num_hops: int = 2,
|
|
136
|
+
max_nodes_per_hop: Optional[int] = None,
|
|
137
|
+
directed: bool = False,
|
|
138
|
+
) -> SubgraphResult:
|
|
139
|
+
"""Retrieve multi-hop subgraph around specified entities."""
|
|
140
|
+
return self.retriever.retrieve_subgraph(
|
|
141
|
+
entities=entities,
|
|
142
|
+
num_hops=num_hops,
|
|
143
|
+
max_nodes_per_hop=max_nodes_per_hop,
|
|
144
|
+
directed=directed,
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
def verbalize(
|
|
148
|
+
self,
|
|
149
|
+
subgraph: SubgraphResult,
|
|
150
|
+
format_style: Literal["triples", "markdown", "natural"] = "markdown",
|
|
151
|
+
max_triples: Optional[int] = 50,
|
|
152
|
+
) -> str:
|
|
153
|
+
"""Verbalize extracted subgraph into text knowledge for LLM prompt context."""
|
|
154
|
+
return verbalize_subgraph(
|
|
155
|
+
subgraph=subgraph,
|
|
156
|
+
id_to_entity=self.id_to_entity,
|
|
157
|
+
id_to_relation=self.id_to_relation,
|
|
158
|
+
format_style=format_style,
|
|
159
|
+
max_triples=max_triples,
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
def build_prompt(
|
|
163
|
+
self,
|
|
164
|
+
query: str,
|
|
165
|
+
entities: Optional[Sequence[Union[str, int]]] = None,
|
|
166
|
+
num_hops: int = 2,
|
|
167
|
+
format_style: Literal["triples", "markdown", "natural"] = "markdown",
|
|
168
|
+
model_family: Literal["llama3", "mistral", "chatml", "standard"] = "llama3",
|
|
169
|
+
system_prompt: Optional[str] = None,
|
|
170
|
+
) -> str:
|
|
171
|
+
"""End-to-end prompt builder: extracts subgraph and formats full LLM prompt."""
|
|
172
|
+
if entities is None or len(entities) == 0:
|
|
173
|
+
entities = self.retriever.find_entities_in_text(query)
|
|
174
|
+
|
|
175
|
+
subgraph = self.retrieve(entities=entities, num_hops=num_hops)
|
|
176
|
+
context = self.verbalize(subgraph=subgraph, format_style=format_style)
|
|
177
|
+
return format_llm_prompt(
|
|
178
|
+
query=query,
|
|
179
|
+
context=context,
|
|
180
|
+
system_prompt=system_prompt,
|
|
181
|
+
model_family=model_family,
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
def encode_prefix(self, subgraph: SubgraphResult) -> any:
|
|
185
|
+
"""Encode extracted subgraph into LLM prefix embeddings.
|
|
186
|
+
|
|
187
|
+
Returns:
|
|
188
|
+
Tensor of shape `(1, num_prefix_tokens, llm_dim)`.
|
|
189
|
+
"""
|
|
190
|
+
if self.connector is None:
|
|
191
|
+
raise ValueError("No KGLLMConnector configured on this GraphRAG instance.")
|
|
192
|
+
return self.connector.encode_subgraph(subgraph)
|
k3_node/rag/projector.py
ADDED
|
@@ -0,0 +1,184 @@
|
|
|
1
|
+
"""Prefix Projectors and KG-LLM Connectors for prompt augmentation."""
|
|
2
|
+
|
|
3
|
+
from typing import Literal, Optional, Tuple, Union
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import layers, ops
|
|
7
|
+
|
|
8
|
+
from k3_node.rag.subgraph import SubgraphResult
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class GraphPrefixProjector(keras.layers.Layer):
|
|
12
|
+
r"""Projects graph/KGE representations into virtual prefix vectors for LLM prompt augmentation.
|
|
13
|
+
|
|
14
|
+
Maps a graph embedding vector of shape `(batch_size, in_channels)` into `num_prefix_tokens`
|
|
15
|
+
virtual token embeddings of dimension `llm_dim` (e.g. 4096 for Llama 3 / Mistral),
|
|
16
|
+
suitable for prepending to text token embeddings in LLM forward passes.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
in_channels: Input graph feature dimension from the GNN/KGE encoder.
|
|
20
|
+
llm_dim: Embedding dimension of the target LLM (e.g., 4096 for Llama-3-8B / Mistral-7B,
|
|
21
|
+
2048 for Gemma-2B). (default: 4096)
|
|
22
|
+
num_prefix_tokens: Number of virtual prefix tokens to produce. (default: 8)
|
|
23
|
+
projector_type: Architecture of the projection head:
|
|
24
|
+
- `"mlp"`: 2-layer MLP with LayerNorm and GELU activation.
|
|
25
|
+
- `"linear"`: Single linear transformation.
|
|
26
|
+
hidden_dim: Optional hidden dimension for MLP projector. (default: `2 * in_channels`)
|
|
27
|
+
dropout: Dropout probability. (default: 0.0)
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
def __init__(
|
|
31
|
+
self,
|
|
32
|
+
in_channels: int,
|
|
33
|
+
llm_dim: int = 4096,
|
|
34
|
+
num_prefix_tokens: int = 8,
|
|
35
|
+
projector_type: Literal["mlp", "linear"] = "mlp",
|
|
36
|
+
hidden_dim: Optional[int] = None,
|
|
37
|
+
dropout: float = 0.0,
|
|
38
|
+
**kwargs,
|
|
39
|
+
):
|
|
40
|
+
super().__init__(**kwargs)
|
|
41
|
+
self.in_channels = in_channels
|
|
42
|
+
self.llm_dim = llm_dim
|
|
43
|
+
self.num_prefix_tokens = num_prefix_tokens
|
|
44
|
+
self.projector_type = projector_type
|
|
45
|
+
self.dropout_rate = dropout
|
|
46
|
+
|
|
47
|
+
out_total = num_prefix_tokens * llm_dim
|
|
48
|
+
if hidden_dim is None:
|
|
49
|
+
hidden_dim = max(in_channels * 2, 512)
|
|
50
|
+
|
|
51
|
+
if projector_type == "linear":
|
|
52
|
+
self.net = layers.Dense(out_total)
|
|
53
|
+
else: # "mlp"
|
|
54
|
+
self.net = keras.Sequential(
|
|
55
|
+
[
|
|
56
|
+
layers.Dense(hidden_dim),
|
|
57
|
+
layers.LayerNormalization(),
|
|
58
|
+
layers.Activation("gelu"),
|
|
59
|
+
layers.Dropout(dropout) if dropout > 0.0 else layers.Identity(),
|
|
60
|
+
layers.Dense(out_total),
|
|
61
|
+
layers.LayerNormalization(),
|
|
62
|
+
]
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
def call(self, graph_embedding: any, training: Optional[bool] = None) -> any:
|
|
66
|
+
"""Project graph embedding into soft prefix token embeddings.
|
|
67
|
+
|
|
68
|
+
Args:
|
|
69
|
+
graph_embedding: Tensor of shape `(batch_size, in_channels)` or `(in_channels,)`.
|
|
70
|
+
|
|
71
|
+
Returns:
|
|
72
|
+
Prefix tensor of shape `(batch_size, num_prefix_tokens, llm_dim)`.
|
|
73
|
+
"""
|
|
74
|
+
graph_embedding = ops.convert_to_tensor(graph_embedding)
|
|
75
|
+
# Ensure 2D (batch_size, in_channels)
|
|
76
|
+
if len(ops.shape(graph_embedding)) == 1:
|
|
77
|
+
graph_embedding = ops.expand_dims(graph_embedding, axis=0)
|
|
78
|
+
|
|
79
|
+
batch_size = ops.shape(graph_embedding)[0]
|
|
80
|
+
flat_proj = self.net(graph_embedding, training=training)
|
|
81
|
+
# Reshape to (batch_size, num_prefix_tokens, llm_dim)
|
|
82
|
+
prefix_tokens = ops.reshape(flat_proj, (batch_size, self.num_prefix_tokens, self.llm_dim))
|
|
83
|
+
return prefix_tokens
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class KGLLMConnector(keras.Model):
|
|
87
|
+
r"""High-level Knowledge Graph to Large Language Model (KG-LLM) Connector.
|
|
88
|
+
|
|
89
|
+
Connects a Knowledge Graph encoder (e.g. `RGCNSubGraphEncoder` or `TransEPrefixEncoder`)
|
|
90
|
+
with a `GraphPrefixProjector` to produce soft prompt prefix embeddings and inject them
|
|
91
|
+
into Llama 3, Mistral, or other LLMs.
|
|
92
|
+
|
|
93
|
+
Example:
|
|
94
|
+
```python
|
|
95
|
+
from k3_node.rag import RGCNSubGraphEncoder, GraphPrefixProjector, KGLLMConnector
|
|
96
|
+
|
|
97
|
+
encoder = RGCNSubGraphEncoder(in_channels=16, hidden_channels=32, out_channels=64, num_relations=5)
|
|
98
|
+
projector = GraphPrefixProjector(in_channels=64, llm_dim=4096, num_prefix_tokens=4)
|
|
99
|
+
connector = KGLLMConnector(encoder=encoder, projector=projector)
|
|
100
|
+
|
|
101
|
+
# Generate prefix embeddings for LLM prompt augmentation
|
|
102
|
+
prefix_embeds = connector.encode_subgraph(subgraph) # [1, 4, 4096]
|
|
103
|
+
|
|
104
|
+
# Inject into LLM input embeddings [batch, seq_len, 4096]
|
|
105
|
+
augmented_inputs = connector.inject_prefix(text_token_embeds, prefix_embeds)
|
|
106
|
+
```
|
|
107
|
+
|
|
108
|
+
Args:
|
|
109
|
+
encoder: Graph or KGE encoder (e.g., `RGCNSubGraphEncoder`, `TransEPrefixEncoder`).
|
|
110
|
+
projector: `GraphPrefixProjector` instance.
|
|
111
|
+
"""
|
|
112
|
+
|
|
113
|
+
def __init__(
|
|
114
|
+
self,
|
|
115
|
+
encoder: keras.layers.Layer,
|
|
116
|
+
projector: GraphPrefixProjector,
|
|
117
|
+
**kwargs,
|
|
118
|
+
):
|
|
119
|
+
super().__init__(**kwargs)
|
|
120
|
+
self.encoder = encoder
|
|
121
|
+
self.projector = projector
|
|
122
|
+
|
|
123
|
+
@property
|
|
124
|
+
def num_prefix_tokens(self) -> int:
|
|
125
|
+
return self.projector.num_prefix_tokens
|
|
126
|
+
|
|
127
|
+
@property
|
|
128
|
+
def llm_dim(self) -> int:
|
|
129
|
+
return self.projector.llm_dim
|
|
130
|
+
|
|
131
|
+
def call(self, *args, **kwargs) -> any:
|
|
132
|
+
"""Encode graph inputs and project directly into LLM prefix tokens."""
|
|
133
|
+
graph_emb = self.encoder(*args, **kwargs)
|
|
134
|
+
return self.projector(graph_emb)
|
|
135
|
+
|
|
136
|
+
def encode_subgraph(self, subgraph: SubgraphResult) -> any:
|
|
137
|
+
"""Encode an extracted SubgraphResult into LLM prefix tokens.
|
|
138
|
+
|
|
139
|
+
Returns:
|
|
140
|
+
Tensor of shape `(1, num_prefix_tokens, llm_dim)`.
|
|
141
|
+
"""
|
|
142
|
+
if hasattr(self.encoder, "encode_subgraph"):
|
|
143
|
+
graph_emb = self.encoder.encode_subgraph(subgraph)
|
|
144
|
+
else:
|
|
145
|
+
graph_emb = self.encoder(subgraph.x, subgraph.edge_index)
|
|
146
|
+
return self.projector(graph_emb)
|
|
147
|
+
|
|
148
|
+
@staticmethod
|
|
149
|
+
def inject_prefix(text_embeddings: any, prefix_embeddings: any) -> any:
|
|
150
|
+
"""Prepend graph prefix embeddings to text token embeddings along sequence dimension.
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
text_embeddings: Tensor of shape `(batch, seq_len, llm_dim)`.
|
|
154
|
+
prefix_embeddings: Tensor of shape `(batch, num_prefix_tokens, llm_dim)`.
|
|
155
|
+
|
|
156
|
+
Returns:
|
|
157
|
+
Concatenated tensor of shape `(batch, num_prefix_tokens + seq_len, llm_dim)`.
|
|
158
|
+
"""
|
|
159
|
+
text_embeddings = ops.convert_to_tensor(text_embeddings)
|
|
160
|
+
prefix_embeddings = ops.convert_to_tensor(prefix_embeddings)
|
|
161
|
+
|
|
162
|
+
# Match batch size if prefix was computed for single batch
|
|
163
|
+
b_text = ops.shape(text_embeddings)[0]
|
|
164
|
+
b_prefix = ops.shape(prefix_embeddings)[0]
|
|
165
|
+
if b_prefix == 1 and b_text > 1:
|
|
166
|
+
prefix_embeddings = ops.repeat(prefix_embeddings, b_text, axis=0)
|
|
167
|
+
|
|
168
|
+
return ops.concatenate([prefix_embeddings, text_embeddings], axis=1)
|
|
169
|
+
|
|
170
|
+
@staticmethod
|
|
171
|
+
def extend_attention_mask(attention_mask: any, num_prefix_tokens: int) -> any:
|
|
172
|
+
"""Extend LLM binary attention mask with 1s for the prepended prefix tokens.
|
|
173
|
+
|
|
174
|
+
Args:
|
|
175
|
+
attention_mask: Tensor of shape `(batch, seq_len)`.
|
|
176
|
+
num_prefix_tokens: Number of prefix tokens prepended.
|
|
177
|
+
|
|
178
|
+
Returns:
|
|
179
|
+
Extended mask of shape `(batch, num_prefix_tokens + seq_len)`.
|
|
180
|
+
"""
|
|
181
|
+
attention_mask = ops.convert_to_tensor(attention_mask)
|
|
182
|
+
batch_size = ops.shape(attention_mask)[0]
|
|
183
|
+
prefix_mask = ops.ones((batch_size, num_prefix_tokens), dtype=attention_mask.dtype)
|
|
184
|
+
return ops.concatenate([prefix_mask, attention_mask], axis=1)
|