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/subgraph.py
ADDED
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
"""Subgraph extraction utilities for GraphRAG and Knowledge Graph LLM integration."""
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from typing import Dict, List, Optional, Sequence, Tuple, Union
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
from keras import ops
|
|
8
|
+
|
|
9
|
+
from k3_node.data import Data
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@dataclass
|
|
13
|
+
class SubgraphResult:
|
|
14
|
+
"""Structured result of a subgraph extraction around retrieved entities.
|
|
15
|
+
|
|
16
|
+
Attributes:
|
|
17
|
+
edge_index: Tensor of shape `(2, num_edges)` with subgraph edges.
|
|
18
|
+
edge_type: Optional tensor of shape `(num_edges,)` with relation types.
|
|
19
|
+
edge_attr: Optional tensor of shape `(num_edges, edge_dim)` with edge features.
|
|
20
|
+
x: Optional tensor of shape `(num_nodes, in_channels)` with node features.
|
|
21
|
+
nodes: 1D numpy array of original node indices in the full graph.
|
|
22
|
+
center_nodes: 1D numpy array of center entity indices in the relabeled subgraph.
|
|
23
|
+
mapping: Dictionary mapping original node index to subgraph node index.
|
|
24
|
+
edge_mask: 1D boolean array indicating edges kept from original graph.
|
|
25
|
+
num_nodes: Number of nodes in the extracted subgraph.
|
|
26
|
+
num_edges: Number of edges in the extracted subgraph.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
edge_index: any
|
|
30
|
+
edge_type: Optional[any] = None
|
|
31
|
+
edge_attr: Optional[any] = None
|
|
32
|
+
x: Optional[any] = None
|
|
33
|
+
nodes: Optional[np.ndarray] = None
|
|
34
|
+
center_nodes: Optional[np.ndarray] = None
|
|
35
|
+
mapping: Optional[Dict[int, int]] = None
|
|
36
|
+
edge_mask: Optional[np.ndarray] = None
|
|
37
|
+
num_nodes: int = 0
|
|
38
|
+
num_edges: int = 0
|
|
39
|
+
|
|
40
|
+
def to_data(self) -> Data:
|
|
41
|
+
"""Convert extracted subgraph into a K3-Node Data object."""
|
|
42
|
+
return Data(
|
|
43
|
+
x=self.x,
|
|
44
|
+
edge_index=self.edge_index,
|
|
45
|
+
edge_type=self.edge_type,
|
|
46
|
+
edge_attr=self.edge_attr,
|
|
47
|
+
center_nodes=ops.convert_to_tensor(self.center_nodes, dtype="int32")
|
|
48
|
+
if self.center_nodes is not None
|
|
49
|
+
else None,
|
|
50
|
+
original_nodes=ops.convert_to_tensor(self.nodes, dtype="int32")
|
|
51
|
+
if self.nodes is not None
|
|
52
|
+
else None,
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def extract_subgraph(
|
|
57
|
+
entities: Union[int, Sequence[int], np.ndarray, any],
|
|
58
|
+
edge_index: Union[np.ndarray, any],
|
|
59
|
+
edge_type: Optional[Union[np.ndarray, any]] = None,
|
|
60
|
+
edge_attr: Optional[Union[np.ndarray, any]] = None,
|
|
61
|
+
x: Optional[Union[np.ndarray, any]] = None,
|
|
62
|
+
num_hops: int = 2,
|
|
63
|
+
max_nodes_per_hop: Optional[int] = None,
|
|
64
|
+
directed: bool = False,
|
|
65
|
+
relabel_nodes: bool = True,
|
|
66
|
+
num_nodes: Optional[int] = None,
|
|
67
|
+
) -> SubgraphResult:
|
|
68
|
+
"""Extract multi-hop enclosing or ego-subgraphs around retrieved entities.
|
|
69
|
+
|
|
70
|
+
Given a knowledge graph or relational graph, this function expands `num_hops`
|
|
71
|
+
around seed `entities`, keeping all induced edges, relation types, and node/edge
|
|
72
|
+
attributes.
|
|
73
|
+
|
|
74
|
+
Args:
|
|
75
|
+
entities: Single node index or list/array of seed entity indices.
|
|
76
|
+
edge_index: Graph connectivity tensor of shape `(2, num_edges)`.
|
|
77
|
+
edge_type: Optional 1D relation type tensor of shape `(num_edges,)`.
|
|
78
|
+
edge_attr: Optional edge attribute tensor of shape `(num_edges, edge_dim)`.
|
|
79
|
+
x: Optional node feature tensor of shape `(total_nodes, in_channels)`.
|
|
80
|
+
num_hops: Number of hops to expand around retrieved entities. (default: 2)
|
|
81
|
+
max_nodes_per_hop: Optional maximum number of neighboring nodes to keep
|
|
82
|
+
per hop (useful to restrict explosion on hub entities).
|
|
83
|
+
directed: If True, only follows outgoing edges. If False, follows edges
|
|
84
|
+
in both directions (standard for KG context expansion). (default: False)
|
|
85
|
+
relabel_nodes: If True, relabels subgraph node IDs to `0..num_subgraph_nodes-1`.
|
|
86
|
+
(default: True)
|
|
87
|
+
num_nodes: Optional total number of nodes in graph. Inferred if not given.
|
|
88
|
+
|
|
89
|
+
Returns:
|
|
90
|
+
`SubgraphResult` containing relabeled edge_index, edge_type, features,
|
|
91
|
+
and entity mappings.
|
|
92
|
+
"""
|
|
93
|
+
edge_index_np = np.asarray(ops.convert_to_numpy(edge_index)).astype(np.int64)
|
|
94
|
+
if num_nodes is None:
|
|
95
|
+
num_nodes = int(edge_index_np.max()) + 1 if edge_index_np.size > 0 else 0
|
|
96
|
+
|
|
97
|
+
entities_arr = np.atleast_1d(np.asarray(ops.convert_to_numpy(entities))).astype(np.int64)
|
|
98
|
+
if entities_arr.size == 0:
|
|
99
|
+
empty_ei = ops.convert_to_tensor(np.zeros((2, 0), dtype=np.int64), dtype=edge_index.dtype)
|
|
100
|
+
return SubgraphResult(
|
|
101
|
+
edge_index=empty_ei,
|
|
102
|
+
nodes=np.array([], dtype=np.int64),
|
|
103
|
+
center_nodes=np.array([], dtype=np.int64),
|
|
104
|
+
mapping={},
|
|
105
|
+
num_nodes=0,
|
|
106
|
+
num_edges=0,
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
row, col = edge_index_np[0], edge_index_np[1]
|
|
110
|
+
subsets = [entities_arr]
|
|
111
|
+
visited = set(entities_arr.tolist())
|
|
112
|
+
|
|
113
|
+
for _ in range(num_hops):
|
|
114
|
+
current_frontier = subsets[-1]
|
|
115
|
+
if len(current_frontier) == 0:
|
|
116
|
+
break
|
|
117
|
+
|
|
118
|
+
mask_frontier = np.zeros(num_nodes, dtype=bool)
|
|
119
|
+
mask_frontier[current_frontier] = True
|
|
120
|
+
|
|
121
|
+
# If undirected (or GraphRAG bidirectional expansion), collect neighbors in both directions
|
|
122
|
+
if not directed:
|
|
123
|
+
neighbors_out = col[mask_frontier[row]]
|
|
124
|
+
neighbors_in = row[mask_frontier[col]]
|
|
125
|
+
new_neighbors = np.concatenate([neighbors_out, neighbors_in])
|
|
126
|
+
else:
|
|
127
|
+
new_neighbors = col[mask_frontier[row]]
|
|
128
|
+
|
|
129
|
+
if new_neighbors.size > 0:
|
|
130
|
+
unique_new = np.unique(new_neighbors)
|
|
131
|
+
unseen = [n for n in unique_new if n not in visited]
|
|
132
|
+
if max_nodes_per_hop is not None and len(unseen) > max_nodes_per_hop:
|
|
133
|
+
unseen = unseen[:max_nodes_per_hop]
|
|
134
|
+
visited.update(unseen)
|
|
135
|
+
subsets.append(np.array(unseen, dtype=np.int64))
|
|
136
|
+
else:
|
|
137
|
+
break
|
|
138
|
+
|
|
139
|
+
subgraph_nodes = np.unique(np.concatenate([s for s in subsets if len(s) > 0]))
|
|
140
|
+
|
|
141
|
+
# Induced edge mask: both endpoints must be in subgraph_nodes
|
|
142
|
+
node_mask = np.zeros(num_nodes, dtype=bool)
|
|
143
|
+
node_mask[subgraph_nodes] = True
|
|
144
|
+
edge_mask = node_mask[row] & node_mask[col]
|
|
145
|
+
|
|
146
|
+
sub_edge_index = edge_index_np[:, edge_mask]
|
|
147
|
+
|
|
148
|
+
mapping = {int(orig): int(new_idx) for new_idx, orig in enumerate(subgraph_nodes)}
|
|
149
|
+
center_indices = np.array([mapping[int(e)] for e in entities_arr if int(e) in mapping], dtype=np.int64)
|
|
150
|
+
|
|
151
|
+
if relabel_nodes:
|
|
152
|
+
new_id_table = np.full(num_nodes, -1, dtype=np.int64)
|
|
153
|
+
new_id_table[subgraph_nodes] = np.arange(len(subgraph_nodes), dtype=np.int64)
|
|
154
|
+
sub_edge_index = new_id_table[sub_edge_index]
|
|
155
|
+
|
|
156
|
+
# Preserve tensors in their original backend / dtype
|
|
157
|
+
sub_edge_index_tensor = ops.convert_to_tensor(sub_edge_index, dtype=edge_index.dtype)
|
|
158
|
+
|
|
159
|
+
sub_edge_type_tensor = None
|
|
160
|
+
if edge_type is not None:
|
|
161
|
+
et_np = np.asarray(ops.convert_to_numpy(edge_type))[edge_mask]
|
|
162
|
+
sub_edge_type_tensor = ops.convert_to_tensor(et_np, dtype=edge_type.dtype)
|
|
163
|
+
|
|
164
|
+
sub_edge_attr_tensor = None
|
|
165
|
+
if edge_attr is not None:
|
|
166
|
+
ea_np = np.asarray(ops.convert_to_numpy(edge_attr))[edge_mask]
|
|
167
|
+
sub_edge_attr_tensor = ops.convert_to_tensor(ea_np, dtype=edge_attr.dtype)
|
|
168
|
+
|
|
169
|
+
sub_x_tensor = None
|
|
170
|
+
if x is not None:
|
|
171
|
+
x_np = np.asarray(ops.convert_to_numpy(x))[subgraph_nodes]
|
|
172
|
+
sub_x_tensor = ops.convert_to_tensor(x_np, dtype=x.dtype)
|
|
173
|
+
|
|
174
|
+
return SubgraphResult(
|
|
175
|
+
edge_index=sub_edge_index_tensor,
|
|
176
|
+
edge_type=sub_edge_type_tensor,
|
|
177
|
+
edge_attr=sub_edge_attr_tensor,
|
|
178
|
+
x=sub_x_tensor,
|
|
179
|
+
nodes=subgraph_nodes,
|
|
180
|
+
center_nodes=center_indices,
|
|
181
|
+
mapping=mapping,
|
|
182
|
+
edge_mask=edge_mask,
|
|
183
|
+
num_nodes=len(subgraph_nodes),
|
|
184
|
+
num_edges=int(sub_edge_index.shape[1]),
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
class KGEntityRetriever:
|
|
189
|
+
"""Knowledge Graph Entity and Subgraph Retriever for GraphRAG.
|
|
190
|
+
|
|
191
|
+
Maintains entity name dictionaries and relation mappings, extracts seed entities
|
|
192
|
+
from text queries, and retrieves enclosing multi-hop subgraphs.
|
|
193
|
+
|
|
194
|
+
Args:
|
|
195
|
+
entity_to_id: Dictionary mapping entity strings to node integer IDs.
|
|
196
|
+
relation_to_id: Dictionary mapping relation strings to edge_type integer IDs.
|
|
197
|
+
edge_index: Graph connectivity tensor of shape `(2, num_edges)`.
|
|
198
|
+
edge_type: Optional relation type tensor of shape `(num_edges,)`.
|
|
199
|
+
edge_attr: Optional edge feature tensor.
|
|
200
|
+
x: Optional node feature tensor.
|
|
201
|
+
"""
|
|
202
|
+
|
|
203
|
+
def __init__(
|
|
204
|
+
self,
|
|
205
|
+
entity_to_id: Dict[str, int],
|
|
206
|
+
relation_to_id: Dict[str, int],
|
|
207
|
+
edge_index: Union[np.ndarray, any],
|
|
208
|
+
edge_type: Optional[Union[np.ndarray, any]] = None,
|
|
209
|
+
edge_attr: Optional[Union[np.ndarray, any]] = None,
|
|
210
|
+
x: Optional[Union[np.ndarray, any]] = None,
|
|
211
|
+
):
|
|
212
|
+
self.entity_to_id = entity_to_id
|
|
213
|
+
self.relation_to_id = relation_to_id
|
|
214
|
+
self.id_to_entity = {v: k for k, v in entity_to_id.items()}
|
|
215
|
+
self.id_to_relation = {v: k for k, v in relation_to_id.items()}
|
|
216
|
+
|
|
217
|
+
self.edge_index = edge_index
|
|
218
|
+
self.edge_type = edge_type
|
|
219
|
+
self.edge_attr = edge_attr
|
|
220
|
+
self.x = x
|
|
221
|
+
|
|
222
|
+
def get_entity_id(self, name: str) -> Optional[int]:
|
|
223
|
+
"""Look up entity ID by exact name (case-insensitive fallback)."""
|
|
224
|
+
if name in self.entity_to_id:
|
|
225
|
+
return self.entity_to_id[name]
|
|
226
|
+
# Case-insensitive fallback
|
|
227
|
+
lower_map = {k.lower(): v for k, v in self.entity_to_id.items()}
|
|
228
|
+
return lower_map.get(name.lower(), None)
|
|
229
|
+
|
|
230
|
+
def find_entities_in_text(self, text: str) -> List[str]:
|
|
231
|
+
"""Find matching known entity names mentioned in a query text."""
|
|
232
|
+
text_lower = text.lower()
|
|
233
|
+
matched = []
|
|
234
|
+
# Sort by length descending to match longest phrases first
|
|
235
|
+
for name in sorted(self.entity_to_id.keys(), key=lambda s: len(s), reverse=True):
|
|
236
|
+
if name.lower() in text_lower:
|
|
237
|
+
matched.append(name)
|
|
238
|
+
return matched
|
|
239
|
+
|
|
240
|
+
def retrieve_subgraph(
|
|
241
|
+
self,
|
|
242
|
+
entities: Union[Sequence[Union[str, int]], str, int],
|
|
243
|
+
num_hops: int = 2,
|
|
244
|
+
max_nodes_per_hop: Optional[int] = None,
|
|
245
|
+
directed: bool = False,
|
|
246
|
+
) -> SubgraphResult:
|
|
247
|
+
"""Extract multi-hop subgraph around the specified entity names or IDs."""
|
|
248
|
+
if isinstance(entities, (str, int)):
|
|
249
|
+
entities = [entities]
|
|
250
|
+
|
|
251
|
+
entity_ids = []
|
|
252
|
+
for e in entities:
|
|
253
|
+
if isinstance(e, str):
|
|
254
|
+
eid = self.get_entity_id(e)
|
|
255
|
+
if eid is not None:
|
|
256
|
+
entity_ids.append(eid)
|
|
257
|
+
else:
|
|
258
|
+
entity_ids.append(int(e))
|
|
259
|
+
|
|
260
|
+
return extract_subgraph(
|
|
261
|
+
entities=entity_ids,
|
|
262
|
+
edge_index=self.edge_index,
|
|
263
|
+
edge_type=self.edge_type,
|
|
264
|
+
edge_attr=self.edge_attr,
|
|
265
|
+
x=self.x,
|
|
266
|
+
num_hops=num_hops,
|
|
267
|
+
max_nodes_per_hop=max_nodes_per_hop,
|
|
268
|
+
directed=directed,
|
|
269
|
+
relabel_nodes=True,
|
|
270
|
+
)
|
k3_node/rag/test_rag.py
ADDED
|
@@ -0,0 +1,347 @@
|
|
|
1
|
+
"""Tests for k3_node.rag GraphRAG and KG-LLM connectors."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pytest
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
import k3_node as k3
|
|
8
|
+
from k3_node.layers.kge import TransE
|
|
9
|
+
from k3_node.rag import (
|
|
10
|
+
GraphPrefixProjector,
|
|
11
|
+
GraphRAG,
|
|
12
|
+
KGEntityRetriever,
|
|
13
|
+
KGLLMConnector,
|
|
14
|
+
RGCNSubGraphEncoder,
|
|
15
|
+
SubgraphResult,
|
|
16
|
+
TransEPrefixEncoder,
|
|
17
|
+
extract_subgraph,
|
|
18
|
+
format_llm_prompt,
|
|
19
|
+
subgraph_to_triples,
|
|
20
|
+
verbalize_subgraph,
|
|
21
|
+
)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@pytest.fixture
|
|
25
|
+
def sample_kg():
|
|
26
|
+
"""Create a sample knowledge graph for testing.
|
|
27
|
+
|
|
28
|
+
Graph:
|
|
29
|
+
0 (Aspirin) --[0: treats]--> 1 (Headache)
|
|
30
|
+
0 (Aspirin) --[1: inhibits]--> 2 (COX-1)
|
|
31
|
+
2 (COX-1) --[2: produces]--> 3 (Prostaglandin)
|
|
32
|
+
3 (Prostaglandin) --[3: causes]--> 1 (Headache)
|
|
33
|
+
4 (Ibuprofen) --[0: treats]--> 1 (Headache)
|
|
34
|
+
4 (Ibuprofen) --[1: inhibits]--> 2 (COX-1)
|
|
35
|
+
"""
|
|
36
|
+
edges = np.array(
|
|
37
|
+
[
|
|
38
|
+
[0, 0, 2, 3, 4, 4],
|
|
39
|
+
[1, 2, 3, 1, 1, 2],
|
|
40
|
+
],
|
|
41
|
+
dtype=np.int64,
|
|
42
|
+
)
|
|
43
|
+
edge_type = np.array([0, 1, 2, 3, 0, 1], dtype=np.int64)
|
|
44
|
+
x = np.random.randn(5, 16).astype(np.float32)
|
|
45
|
+
|
|
46
|
+
entity_to_id = {
|
|
47
|
+
"Aspirin": 0,
|
|
48
|
+
"Headache": 1,
|
|
49
|
+
"COX-1": 2,
|
|
50
|
+
"Prostaglandin": 3,
|
|
51
|
+
"Ibuprofen": 4,
|
|
52
|
+
}
|
|
53
|
+
relation_to_id = {
|
|
54
|
+
"treats": 0,
|
|
55
|
+
"inhibits": 1,
|
|
56
|
+
"produces": 2,
|
|
57
|
+
"causes": 3,
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
return {
|
|
61
|
+
"edge_index": ops.convert_to_tensor(edges, dtype="int32"),
|
|
62
|
+
"edge_type": ops.convert_to_tensor(edge_type, dtype="int32"),
|
|
63
|
+
"x": ops.convert_to_tensor(x, dtype="float32"),
|
|
64
|
+
"entity_to_id": entity_to_id,
|
|
65
|
+
"relation_to_id": relation_to_id,
|
|
66
|
+
"num_nodes": 5,
|
|
67
|
+
"num_relations": 4,
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def test_extract_subgraph_1hop(sample_kg):
|
|
72
|
+
sub = extract_subgraph(
|
|
73
|
+
entities=[0], # Aspirin
|
|
74
|
+
edge_index=sample_kg["edge_index"],
|
|
75
|
+
edge_type=sample_kg["edge_type"],
|
|
76
|
+
x=sample_kg["x"],
|
|
77
|
+
num_hops=1,
|
|
78
|
+
directed=False,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
assert isinstance(sub, SubgraphResult)
|
|
82
|
+
# Neighbors of 0 within 1 hop: 0, 1, 2
|
|
83
|
+
assert set(sub.nodes.tolist()) == {0, 1, 2}
|
|
84
|
+
assert sub.center_nodes.tolist() == [sub.mapping[0]]
|
|
85
|
+
assert sub.num_nodes == 3
|
|
86
|
+
assert sub.num_edges > 0
|
|
87
|
+
assert sub.x is not None
|
|
88
|
+
assert ops.shape(sub.x)[0] == 3
|
|
89
|
+
|
|
90
|
+
# Check to_data()
|
|
91
|
+
data = sub.to_data()
|
|
92
|
+
assert isinstance(data, k3.data.Data)
|
|
93
|
+
assert ops.shape(data.edge_index)[0] == 2
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def test_extract_subgraph_2hop_multiseed(sample_kg):
|
|
97
|
+
sub = extract_subgraph(
|
|
98
|
+
entities=[0, 3], # Aspirin & Prostaglandin
|
|
99
|
+
edge_index=sample_kg["edge_index"],
|
|
100
|
+
edge_type=sample_kg["edge_type"],
|
|
101
|
+
num_hops=2,
|
|
102
|
+
)
|
|
103
|
+
assert sub.num_nodes == 5
|
|
104
|
+
assert len(sub.center_nodes) == 2
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def test_kg_entity_retriever(sample_kg):
|
|
108
|
+
retriever = KGEntityRetriever(
|
|
109
|
+
entity_to_id=sample_kg["entity_to_id"],
|
|
110
|
+
relation_to_id=sample_kg["relation_to_id"],
|
|
111
|
+
edge_index=sample_kg["edge_index"],
|
|
112
|
+
edge_type=sample_kg["edge_type"],
|
|
113
|
+
x=sample_kg["x"],
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
# Lookup
|
|
117
|
+
assert retriever.get_entity_id("Aspirin") == 0
|
|
118
|
+
assert retriever.get_entity_id("aspirin") == 0 # case-insensitive
|
|
119
|
+
assert retriever.get_entity_id("Unknown") is None
|
|
120
|
+
|
|
121
|
+
# Text entity extraction
|
|
122
|
+
query = "Does Aspirin or Ibuprofen treat headache?"
|
|
123
|
+
found = retriever.find_entities_in_text(query)
|
|
124
|
+
assert "Aspirin" in found
|
|
125
|
+
assert "Ibuprofen" in found
|
|
126
|
+
assert "Headache" in found
|
|
127
|
+
|
|
128
|
+
# Subgraph retrieval
|
|
129
|
+
sub = retriever.retrieve_subgraph(["Aspirin"], num_hops=1)
|
|
130
|
+
assert sub.num_nodes >= 2
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def test_verbalization(sample_kg):
|
|
134
|
+
sub = extract_subgraph(
|
|
135
|
+
entities=[0],
|
|
136
|
+
edge_index=sample_kg["edge_index"],
|
|
137
|
+
edge_type=sample_kg["edge_type"],
|
|
138
|
+
num_hops=1,
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
id_to_e = {v: k for k, v in sample_kg["entity_to_id"].items()}
|
|
142
|
+
id_to_r = {v: k for k, v in sample_kg["relation_to_id"].items()}
|
|
143
|
+
|
|
144
|
+
# Triples list
|
|
145
|
+
triples = subgraph_to_triples(sub, id_to_e, id_to_r)
|
|
146
|
+
assert len(triples) > 0
|
|
147
|
+
assert ("Aspirin", "treats", "Headache") in triples
|
|
148
|
+
|
|
149
|
+
# Markdown format
|
|
150
|
+
md_text = verbalize_subgraph(sub, id_to_e, id_to_r, format_style="markdown")
|
|
151
|
+
assert "**Aspirin**" in md_text
|
|
152
|
+
assert "*treats*" in md_text
|
|
153
|
+
|
|
154
|
+
# Natural language format
|
|
155
|
+
natural_text = verbalize_subgraph(sub, id_to_e, id_to_r, format_style="natural")
|
|
156
|
+
assert "Aspirin treats Headache." in natural_text
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def test_prompt_formatting():
|
|
160
|
+
context = "- **Aspirin** — *treats* -> **Headache**"
|
|
161
|
+
query = "What treats headache?"
|
|
162
|
+
|
|
163
|
+
# Llama 3
|
|
164
|
+
llama3_prompt = format_llm_prompt(query, context, model_family="llama3")
|
|
165
|
+
assert "<|start_header_id|>system<|end_header_id|>" in llama3_prompt
|
|
166
|
+
assert "<|start_header_id|>user<|end_header_id|>" in llama3_prompt
|
|
167
|
+
assert "Knowledge Graph Context:" in llama3_prompt
|
|
168
|
+
assert query in llama3_prompt
|
|
169
|
+
|
|
170
|
+
# Mistral
|
|
171
|
+
mistral_prompt = format_llm_prompt(query, context, model_family="mistral")
|
|
172
|
+
assert "<s>[INST]" in mistral_prompt
|
|
173
|
+
assert "[/INST]" in mistral_prompt
|
|
174
|
+
|
|
175
|
+
# ChatML
|
|
176
|
+
chatml_prompt = format_llm_prompt(query, context, model_family="chatml")
|
|
177
|
+
assert "<|im_start|>system" in chatml_prompt
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
def test_rgcn_subgraph_encoder(sample_kg):
|
|
181
|
+
encoder = RGCNSubGraphEncoder(
|
|
182
|
+
in_channels=16,
|
|
183
|
+
hidden_channels=32,
|
|
184
|
+
out_channels=64,
|
|
185
|
+
num_relations=sample_kg["num_relations"],
|
|
186
|
+
num_layers=2,
|
|
187
|
+
pooling="center",
|
|
188
|
+
)
|
|
189
|
+
|
|
190
|
+
sub = extract_subgraph(
|
|
191
|
+
entities=[0],
|
|
192
|
+
edge_index=sample_kg["edge_index"],
|
|
193
|
+
edge_type=sample_kg["edge_type"],
|
|
194
|
+
x=sample_kg["x"],
|
|
195
|
+
num_hops=1,
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
emb = encoder.encode_subgraph(sub)
|
|
199
|
+
assert ops.shape(emb) == (1, 64)
|
|
200
|
+
|
|
201
|
+
# Test other pooling options
|
|
202
|
+
encoder_mean = RGCNSubGraphEncoder(
|
|
203
|
+
in_channels=16,
|
|
204
|
+
hidden_channels=32,
|
|
205
|
+
out_channels=64,
|
|
206
|
+
num_relations=sample_kg["num_relations"],
|
|
207
|
+
pooling="mean",
|
|
208
|
+
)
|
|
209
|
+
emb_mean = encoder_mean.encode_subgraph(sub)
|
|
210
|
+
assert ops.shape(emb_mean) == (1, 64)
|
|
211
|
+
|
|
212
|
+
|
|
213
|
+
def test_transe_prefix_encoder(sample_kg):
|
|
214
|
+
transe = TransE(
|
|
215
|
+
num_nodes=sample_kg["num_nodes"],
|
|
216
|
+
num_relations=sample_kg["num_relations"],
|
|
217
|
+
hidden_channels=32,
|
|
218
|
+
)
|
|
219
|
+
|
|
220
|
+
kge_encoder = TransEPrefixEncoder(
|
|
221
|
+
kge_model=transe,
|
|
222
|
+
out_channels=64,
|
|
223
|
+
pooling="center",
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
sub = extract_subgraph(
|
|
227
|
+
entities=[0, 1],
|
|
228
|
+
edge_index=sample_kg["edge_index"],
|
|
229
|
+
edge_type=sample_kg["edge_type"],
|
|
230
|
+
num_hops=1,
|
|
231
|
+
)
|
|
232
|
+
|
|
233
|
+
emb = kge_encoder.encode_subgraph(sub)
|
|
234
|
+
assert ops.shape(emb) == (1, 64)
|
|
235
|
+
|
|
236
|
+
# Test standalone initialization without pretrained model
|
|
237
|
+
standalone_encoder = TransEPrefixEncoder(
|
|
238
|
+
num_nodes=10,
|
|
239
|
+
num_relations=4,
|
|
240
|
+
embedding_dim=32,
|
|
241
|
+
out_channels=64,
|
|
242
|
+
)
|
|
243
|
+
emb_standalone = standalone_encoder.encode_subgraph(sub)
|
|
244
|
+
assert ops.shape(emb_standalone) == (1, 64)
|
|
245
|
+
|
|
246
|
+
|
|
247
|
+
def test_graph_prefix_projector():
|
|
248
|
+
projector = GraphPrefixProjector(
|
|
249
|
+
in_channels=64,
|
|
250
|
+
llm_dim=256, # test dimension
|
|
251
|
+
num_prefix_tokens=4,
|
|
252
|
+
projector_type="mlp",
|
|
253
|
+
)
|
|
254
|
+
|
|
255
|
+
graph_emb = ops.convert_to_tensor(np.random.randn(1, 64).astype(np.float32))
|
|
256
|
+
prefix = projector(graph_emb)
|
|
257
|
+
assert ops.shape(prefix) == (1, 4, 256)
|
|
258
|
+
|
|
259
|
+
# Linear projector
|
|
260
|
+
lin_projector = GraphPrefixProjector(
|
|
261
|
+
in_channels=64,
|
|
262
|
+
llm_dim=256,
|
|
263
|
+
num_prefix_tokens=4,
|
|
264
|
+
projector_type="linear",
|
|
265
|
+
)
|
|
266
|
+
prefix_lin = lin_projector(graph_emb)
|
|
267
|
+
assert ops.shape(prefix_lin) == (1, 4, 256)
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
def test_kg_llm_connector(sample_kg):
|
|
271
|
+
encoder = RGCNSubGraphEncoder(
|
|
272
|
+
in_channels=16,
|
|
273
|
+
hidden_channels=32,
|
|
274
|
+
out_channels=64,
|
|
275
|
+
num_relations=sample_kg["num_relations"],
|
|
276
|
+
)
|
|
277
|
+
projector = GraphPrefixProjector(
|
|
278
|
+
in_channels=64,
|
|
279
|
+
llm_dim=512,
|
|
280
|
+
num_prefix_tokens=4,
|
|
281
|
+
)
|
|
282
|
+
connector = KGLLMConnector(encoder=encoder, projector=projector)
|
|
283
|
+
|
|
284
|
+
sub = extract_subgraph(
|
|
285
|
+
entities=[0],
|
|
286
|
+
edge_index=sample_kg["edge_index"],
|
|
287
|
+
edge_type=sample_kg["edge_type"],
|
|
288
|
+
x=sample_kg["x"],
|
|
289
|
+
num_hops=1,
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
prefix_tokens = connector.encode_subgraph(sub)
|
|
293
|
+
assert ops.shape(prefix_tokens) == (1, 4, 512)
|
|
294
|
+
|
|
295
|
+
# Test prefix injection into LLM text tokens
|
|
296
|
+
text_tokens = ops.convert_to_tensor(np.random.randn(1, 10, 512).astype(np.float32))
|
|
297
|
+
augmented = connector.inject_prefix(text_tokens, prefix_tokens)
|
|
298
|
+
assert ops.shape(augmented) == (1, 14, 512)
|
|
299
|
+
|
|
300
|
+
# Test attention mask extension
|
|
301
|
+
attn_mask = ops.ones((1, 10), dtype="int32")
|
|
302
|
+
extended_mask = connector.extend_attention_mask(attn_mask, num_prefix_tokens=4)
|
|
303
|
+
assert ops.shape(extended_mask) == (1, 14)
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
def test_graph_rag_pipeline(sample_kg):
|
|
307
|
+
# Pipeline with RGCN
|
|
308
|
+
rag_rgcn = GraphRAG(
|
|
309
|
+
edge_index=sample_kg["edge_index"],
|
|
310
|
+
edge_type=sample_kg["edge_type"],
|
|
311
|
+
entity_to_id=sample_kg["entity_to_id"],
|
|
312
|
+
relation_to_id=sample_kg["relation_to_id"],
|
|
313
|
+
x=sample_kg["x"],
|
|
314
|
+
encoder_type="rgcn",
|
|
315
|
+
hidden_dim=32,
|
|
316
|
+
encoder_out_dim=64,
|
|
317
|
+
llm_dim=256,
|
|
318
|
+
num_prefix_tokens=4,
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
# Retrieval
|
|
322
|
+
sub = rag_rgcn.retrieve(["Aspirin"], num_hops=1)
|
|
323
|
+
assert sub.num_nodes >= 2
|
|
324
|
+
|
|
325
|
+
# Prompt builder
|
|
326
|
+
prompt = rag_rgcn.build_prompt("What does Aspirin treat?", model_family="llama3")
|
|
327
|
+
assert "Aspirin" in prompt
|
|
328
|
+
assert "<|begin_of_text|>" in prompt
|
|
329
|
+
|
|
330
|
+
# Prefix encoding
|
|
331
|
+
prefix = rag_rgcn.encode_prefix(sub)
|
|
332
|
+
assert ops.shape(prefix) == (1, 4, 256)
|
|
333
|
+
|
|
334
|
+
# Pipeline with TransE
|
|
335
|
+
rag_transe = GraphRAG(
|
|
336
|
+
edge_index=sample_kg["edge_index"],
|
|
337
|
+
edge_type=sample_kg["edge_type"],
|
|
338
|
+
entity_to_id=sample_kg["entity_to_id"],
|
|
339
|
+
relation_to_id=sample_kg["relation_to_id"],
|
|
340
|
+
encoder_type="transe",
|
|
341
|
+
hidden_dim=32,
|
|
342
|
+
encoder_out_dim=64,
|
|
343
|
+
llm_dim=256,
|
|
344
|
+
num_prefix_tokens=4,
|
|
345
|
+
)
|
|
346
|
+
prefix_transe = rag_transe.encode_prefix(sub)
|
|
347
|
+
assert ops.shape(prefix_transe) == (1, 4, 256)
|