GraphRAG & KG-LLM Connectors
The k3_node.rag module provides a comprehensive suite of tools for connecting Knowledge Graphs to Large Language Models (LLMs) such as Meta Llama 3 and Mistral.
1. Subgraph Extraction
extract_subgraph
k3_node.rag.extract_subgraph(entities, edge_index, edge_type=None, edge_attr=None, x=None, num_hops=2, max_nodes_per_hop=None, directed=False, relabel_nodes=True, num_nodes=None)
Extract multi-hop enclosing or ego-subgraphs around retrieved entities.
Given a knowledge graph or relational graph, this function expands num_hops
around seed entities, keeping all induced edges, relation types, and node/edge
attributes.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
entities
|
Union[int, Sequence[int], ndarray, any]
|
Single node index or list/array of seed entity indices. |
required |
edge_index
|
Union[ndarray, any]
|
Graph connectivity tensor of shape |
required |
edge_type
|
Optional[Union[ndarray, any]]
|
Optional 1D relation type tensor of shape |
None
|
edge_attr
|
Optional[Union[ndarray, any]]
|
Optional edge attribute tensor of shape |
None
|
x
|
Optional[Union[ndarray, any]]
|
Optional node feature tensor of shape |
None
|
num_hops
|
int
|
Number of hops to expand around retrieved entities. (default: 2) |
2
|
max_nodes_per_hop
|
Optional[int]
|
Optional maximum number of neighboring nodes to keep per hop (useful to restrict explosion on hub entities). |
None
|
directed
|
bool
|
If True, only follows outgoing edges. If False, follows edges in both directions (standard for KG context expansion). (default: False) |
False
|
relabel_nodes
|
bool
|
If True, relabels subgraph node IDs to |
True
|
num_nodes
|
Optional[int]
|
Optional total number of nodes in graph. Inferred if not given. |
None
|
Returns:
| Type | Description |
|---|---|
SubgraphResult
|
|
SubgraphResult
|
and entity mappings. |
Usage Example:
import keras
from keras import ops
from k3_node.rag import extract_subgraph
# A small knowledge graph with 5 nodes and 6 directed edges
edge_index = ops.convert_to_tensor([[0, 0, 2, 3, 4, 4], [1, 2, 3, 1, 1, 2]], dtype="int32")
edge_type = ops.convert_to_tensor([0, 1, 2, 3, 0, 1], dtype="int32")
x = keras.random.normal((5, 16))
# Extract 1-hop enclosing subgraph around node 0
subgraph = extract_subgraph(entities=[0], edge_index=edge_index, edge_type=edge_type, x=x, num_hops=1)
print("Nodes in subgraph:", subgraph.nodes)
print("Subgraph edges:", subgraph.edge_index.shape)
KGEntityRetriever
k3_node.rag.KGEntityRetriever
Knowledge Graph Entity and Subgraph Retriever for GraphRAG.
Maintains entity name dictionaries and relation mappings, extracts seed entities from text queries, and retrieves enclosing multi-hop subgraphs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
entity_to_id
|
Dict[str, int]
|
Dictionary mapping entity strings to node integer IDs. |
required |
relation_to_id
|
Dict[str, int]
|
Dictionary mapping relation strings to edge_type integer IDs. |
required |
edge_index
|
Union[ndarray, any]
|
Graph connectivity tensor of shape |
required |
edge_type
|
Optional[Union[ndarray, any]]
|
Optional relation type tensor of shape |
None
|
edge_attr
|
Optional[Union[ndarray, any]]
|
Optional edge feature tensor. |
None
|
x
|
Optional[Union[ndarray, any]]
|
Optional node feature tensor. |
None
|
find_entities_in_text(text)
Find matching known entity names mentioned in a query text.
get_entity_id(name)
Look up entity ID by exact name (case-insensitive fallback).
retrieve_subgraph(entities, num_hops=2, max_nodes_per_hop=None, directed=False)
Extract multi-hop subgraph around the specified entity names or IDs.
Usage Example:
from k3_node.rag import KGEntityRetriever
entity_to_id = {"Aspirin": 0, "Headache": 1, "COX-1": 2}
relation_to_id = {"treats": 0, "inhibits": 1}
retriever = KGEntityRetriever(
entity_to_id=entity_to_id,
relation_to_id=relation_to_id,
edge_index=edge_index,
edge_type=edge_type,
)
# Extract entities mentioned in free-form text query
matched_entities = retriever.find_entities_in_text("What treats Headache?")
print(matched_entities) # ['Headache']
# Retrieve multi-hop neighborhood
subgraph = retriever.retrieve_subgraph(matched_entities, num_hops=2)
2. Graph Verbalization & Prompt Formatting
verbalize_subgraph
k3_node.rag.verbalize_subgraph(subgraph, id_to_entity=None, id_to_relation=None, format_style='markdown', max_triples=50)
Verbalize an extracted subgraph into textual knowledge for LLM prompt augmentation.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
subgraph
|
SubgraphResult
|
Extracted SubgraphResult around retrieved entities. |
required |
id_to_entity
|
Optional[Dict[int, str]]
|
Optional dictionary mapping node ID to entity name. |
None
|
id_to_relation
|
Optional[Dict[int, str]]
|
Optional dictionary mapping relation ID to relation name. |
None
|
format_style
|
Literal['triples', 'markdown', 'natural']
|
|
'markdown'
|
max_triples
|
Optional[int]
|
Maximum number of facts to include in the context. |
50
|
Returns:
| Type | Description |
|---|---|
str
|
Formatted textual context string ready to be injected into an LLM prompt. |
Usage Example:
from k3_node.rag import verbalize_subgraph
# Markdown format
markdown_context = verbalize_subgraph(
subgraph,
id_to_entity={0: "Aspirin", 1: "Headache", 2: "COX-1"},
id_to_relation={0: "treats", 1: "inhibits"},
format_style="markdown",
)
print(markdown_context)
format_llm_prompt
k3_node.rag.format_llm_prompt(query, context, system_prompt=None, model_family='llama3')
Format query and retrieved KG context into prompt templates for LLMs.
Supported model families include:
- "llama3": Meta Llama 3 / 3.1 instruct template.
- "mistral": Mistral / Mixtral instruct template.
- "chatml": OpenAI / Qwen ChatML template.
- "standard": General markdown system/user format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
query
|
str
|
The user's input question or instruction. |
required |
context
|
str
|
The verbalized knowledge graph context. |
required |
system_prompt
|
Optional[str]
|
Optional system prompt to instruct the LLM. |
None
|
model_family
|
Literal['llama3', 'mistral', 'chatml', 'standard']
|
Target model prompt format. (default: "llama3") |
'llama3'
|
Returns:
| Type | Description |
|---|---|
str
|
Formatted prompt string. |
Usage Example:
from k3_node.rag import format_llm_prompt
# Llama 3 prompt formatting
prompt = format_llm_prompt(
query="How does Aspirin treat Headache?",
context=markdown_context,
model_family="llama3",
)
print(prompt)
3. Subgraph GNN & KGE Encoders
RGCNSubGraphEncoder
k3_node.rag.RGCNSubGraphEncoder
Bases: Model
Relational Graph Convolutional Network (RGCN) encoder for multi-relational subgraphs.
Processes subgraphs extracted from Knowledge Graphs with multiple relation types, updating entity representations through relational message passing and pooling them into a dense graph embedding.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
in_channels
|
int
|
Dimensionality of input node features. |
required |
hidden_channels
|
int
|
Hidden representation dimension. |
required |
out_channels
|
int
|
Output graph embedding dimension. |
required |
num_relations
|
int
|
Total number of relation types in the Knowledge Graph. |
required |
num_layers
|
int
|
Number of RGCN message passing layers. (default: 2) |
2
|
num_bases
|
Optional[int]
|
Optional number of basis decomposition components for relation weights. |
None
|
pooling
|
Literal['mean', 'sum', 'max', 'center', 'none']
|
Readout pooling strategy across subgraph nodes:
- |
'mean'
|
dropout
|
float
|
Dropout rate applied between convolution layers. (default: 0.0) |
0.0
|
call(x, edge_index, edge_type=None, center_nodes=None, training=None)
Forward pass encoding the subgraph into a graph embedding.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
any
|
Node feature tensor of shape |
required |
edge_index
|
any
|
Graph edge indices of shape |
required |
edge_type
|
Optional[any]
|
1D tensor of relation IDs for each edge |
None
|
center_nodes
|
Optional[any]
|
Optional 1D tensor of seed entity indices in the subgraph. |
None
|
training
|
Optional[bool]
|
Whether running in training mode. |
None
|
Returns:
| Type | Description |
|---|---|
any
|
Tensor of shape |
encode_subgraph(subgraph, default_x_dim=None)
Helper to encode a SubgraphResult directly.
Usage Example:
from k3_node.rag import RGCNSubGraphEncoder
encoder = RGCNSubGraphEncoder(
in_channels=16,
hidden_channels=32,
out_channels=64,
num_relations=4,
num_layers=2,
pooling="center",
)
# Encode subgraph
emb = encoder.encode_subgraph(subgraph)
print("Graph representation shape:", emb.shape) # (1, 64)
TransEPrefixEncoder
k3_node.rag.TransEPrefixEncoder
Bases: Model
Knowledge Graph Embedding (KGE) prefix encoder using TransE representations.
Uses pretrained or end-to-end entity and relation embeddings from TransE (\(h + r \approx t\)) to encode extracted knowledge subgraphs into a unified dense embedding vector.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
num_nodes
|
Optional[int]
|
Total number of entities in the knowledge graph. |
None
|
num_relations
|
Optional[int]
|
Total number of relation types. |
None
|
embedding_dim
|
int
|
Dimension of entity and relation embeddings. (default: 64) |
64
|
out_channels
|
int
|
Output projected graph embedding dimension. (default: 128) |
128
|
kge_model
|
Optional[KGEModel]
|
Optional pretrained |
None
|
pooling
|
Literal['mean', 'sum', 'center']
|
Readout pooling strategy across entities ( |
'mean'
|
call(subgraph_nodes, edge_index=None, edge_type=None, center_nodes=None)
Encode a set of subgraph entities and relations into a dense embedding.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
subgraph_nodes
|
any
|
1D tensor of original entity IDs in the subgraph |
required |
edge_index
|
Optional[any]
|
Optional edge index tensor |
None
|
edge_type
|
Optional[any]
|
Optional relation type tensor |
None
|
center_nodes
|
Optional[any]
|
Optional indices of seed entities within |
None
|
Returns:
| Type | Description |
|---|---|
any
|
Dense embedding tensor of shape |
encode_subgraph(subgraph)
Helper to encode a SubgraphResult directly.
Usage Example:
from k3_node.rag import TransEPrefixEncoder
from k3_node.layers.kge import TransE
# Option A: reuse a pretrained TransE model
transe = TransE(num_nodes=100, num_relations=10, hidden_channels=32)
encoder = TransEPrefixEncoder(kge_model=transe, out_channels=64)
# Option B: standalone encoder
encoder = TransEPrefixEncoder(num_nodes=100, num_relations=10, embedding_dim=32, out_channels=64)
emb = encoder.encode_subgraph(subgraph)
4. Connectors & Virtual Prefix Projectors
GraphPrefixProjector
k3_node.rag.GraphPrefixProjector
Bases: Layer
Projects graph/KGE representations into virtual prefix vectors for LLM prompt augmentation.
Maps a graph embedding vector of shape (batch_size, in_channels) into num_prefix_tokens
virtual token embeddings of dimension llm_dim (e.g. 4096 for Llama 3 / Mistral),
suitable for prepending to text token embeddings in LLM forward passes.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
in_channels
|
int
|
Input graph feature dimension from the GNN/KGE encoder. |
required |
llm_dim
|
int
|
Embedding dimension of the target LLM (e.g., 4096 for Llama-3-8B / Mistral-7B, 2048 for Gemma-2B). (default: 4096) |
4096
|
num_prefix_tokens
|
int
|
Number of virtual prefix tokens to produce. (default: 8) |
8
|
projector_type
|
Literal['mlp', 'linear']
|
Architecture of the projection head:
- |
'mlp'
|
hidden_dim
|
Optional[int]
|
Optional hidden dimension for MLP projector. (default: |
None
|
dropout
|
float
|
Dropout probability. (default: 0.0) |
0.0
|
call(graph_embedding, training=None)
Project graph embedding into soft prefix token embeddings.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
graph_embedding
|
any
|
Tensor of shape |
required |
Returns:
| Type | Description |
|---|---|
any
|
Prefix tensor of shape |
Usage Example:
import keras
from k3_node.rag import GraphPrefixProjector
projector = GraphPrefixProjector(
in_channels=64,
llm_dim=4096, # Llama 3 / Mistral embedding dimension
num_prefix_tokens=8, # Number of soft prefix tokens
projector_type="mlp",
)
graph_embedding = keras.random.normal((1, 64))
prefix_tokens = projector(graph_embedding)
print(prefix_tokens.shape) # (1, 8, 4096)
KGLLMConnector
k3_node.rag.KGLLMConnector
Bases: Model
High-level Knowledge Graph to Large Language Model (KG-LLM) Connector.
Connects a Knowledge Graph encoder (e.g. RGCNSubGraphEncoder or TransEPrefixEncoder)
with a GraphPrefixProjector to produce soft prompt prefix embeddings and inject them
into Llama 3, Mistral, or other LLMs.
Example
from k3_node.rag import RGCNSubGraphEncoder, GraphPrefixProjector, KGLLMConnector
encoder = RGCNSubGraphEncoder(in_channels=16, hidden_channels=32, out_channels=64, num_relations=5)
projector = GraphPrefixProjector(in_channels=64, llm_dim=4096, num_prefix_tokens=4)
connector = KGLLMConnector(encoder=encoder, projector=projector)
# Generate prefix embeddings for LLM prompt augmentation
prefix_embeds = connector.encode_subgraph(subgraph) # [1, 4, 4096]
# Inject into LLM input embeddings [batch, seq_len, 4096]
augmented_inputs = connector.inject_prefix(text_token_embeds, prefix_embeds)
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
encoder
|
Layer
|
Graph or KGE encoder (e.g., |
required |
projector
|
GraphPrefixProjector
|
|
required |
call(*args, **kwargs)
Encode graph inputs and project directly into LLM prefix tokens.
encode_subgraph(subgraph)
Encode an extracted SubgraphResult into LLM prefix tokens.
Returns:
| Type | Description |
|---|---|
any
|
Tensor of shape |
extend_attention_mask(attention_mask, num_prefix_tokens)
staticmethod
Extend LLM binary attention mask with 1s for the prepended prefix tokens.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
attention_mask
|
any
|
Tensor of shape |
required |
num_prefix_tokens
|
int
|
Number of prefix tokens prepended. |
required |
Returns:
| Type | Description |
|---|---|
any
|
Extended mask of shape |
inject_prefix(text_embeddings, prefix_embeddings)
staticmethod
Prepend graph prefix embeddings to text token embeddings along sequence dimension.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
text_embeddings
|
any
|
Tensor of shape |
required |
prefix_embeddings
|
any
|
Tensor of shape |
required |
Returns:
| Type | Description |
|---|---|
any
|
Concatenated tensor of shape |
Usage Example:
import keras
from k3_node.rag import RGCNSubGraphEncoder, GraphPrefixProjector, KGLLMConnector
encoder = RGCNSubGraphEncoder(in_channels=16, hidden_channels=32, out_channels=64, num_relations=5)
projector = GraphPrefixProjector(in_channels=64, llm_dim=4096, num_prefix_tokens=4)
connector = KGLLMConnector(encoder=encoder, projector=projector)
# Encode extracted subgraph directly into LLM prefix tokens
prefix_embeds = connector.encode_subgraph(subgraph) # [1, 4, 4096]
# Inject prefix into LLM text token embeddings
text_token_embeds = keras.random.normal((1, 15, 4096))
augmented_embeds = connector.inject_prefix(text_token_embeds, prefix_embeds)
print(augmented_embeds.shape) # (1, 19, 4096)
5. Unified End-to-End Pipeline
GraphRAG
k3_node.rag.GraphRAG
End-to-End Graph-Augmented Generation (GraphRAG) Pipeline for Knowledge Graphs.
Provides a unified API for: 1. Extracting subgraphs around query entities. 2. Verbalizing structured facts into markdown or natural language prompts for Llama 3 / Mistral. 3. Projecting GNN (RGCN) or KGE (TransE) embeddings into dense prefix vectors for soft prompt augmentation.
Example
import k3_node as k3
from k3_node.rag import GraphRAG
# Setup GraphRAG with your knowledge graph
rag = GraphRAG(
edge_index=edge_index,
edge_type=edge_type,
entity_to_id={"Aspirin": 0, "Headache": 1, "COX-1": 2},
relation_to_id={"treats": 0, "inhibits": 1},
llm_dim=4096, # Llama 3 / Mistral embedding dimension
num_prefix_tokens=4,
)
# 1. Text-based GraphRAG prompt generation
prompt = rag.build_prompt(
query="How does Aspirin alleviate headache?",
entities=["Aspirin", "Headache"],
model_family="llama3",
)
# 2. Dense Prefix Vector encoding
subgraph = rag.retrieve(["Aspirin"], num_hops=2)
prefix_embeds = rag.encode_prefix(subgraph) # [1, 4, 4096]
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
edge_index
|
any
|
Graph connectivity tensor of shape |
required |
edge_type
|
Optional[any]
|
Optional relation IDs tensor of shape |
None
|
entity_to_id
|
Optional[Dict[str, int]]
|
Dictionary mapping entity strings to node IDs. |
None
|
relation_to_id
|
Optional[Dict[str, int]]
|
Optional dictionary mapping relation strings to relation IDs. |
None
|
x
|
Optional[any]
|
Optional node feature tensor. |
None
|
num_relations
|
Optional[int]
|
Optional total number of relation types. |
None
|
encoder_type
|
Literal['rgcn', 'transe', 'none']
|
Subgraph encoder architecture ( |
'rgcn'
|
hidden_dim
|
int
|
Hidden dimension for GNN/KGE encoder. (default: 64) |
64
|
encoder_out_dim
|
int
|
Output dimension of graph encoder before LLM projection. (default: 128) |
128
|
llm_dim
|
int
|
Target LLM embedding dimension (e.g. 4096 for Llama 3 / Mistral). (default: 4096) |
4096
|
num_prefix_tokens
|
int
|
Number of virtual prefix tokens to generate. (default: 8) |
8
|
build_prompt(query, entities=None, num_hops=2, format_style='markdown', model_family='llama3', system_prompt=None)
End-to-end prompt builder: extracts subgraph and formats full LLM prompt.
encode_prefix(subgraph)
Encode extracted subgraph into LLM prefix embeddings.
Returns:
| Type | Description |
|---|---|
any
|
Tensor of shape |
retrieve(entities, num_hops=2, max_nodes_per_hop=None, directed=False)
Retrieve multi-hop subgraph around specified entities.
verbalize(subgraph, format_style='markdown', max_triples=50)
Verbalize extracted subgraph into text knowledge for LLM prompt context.
Usage Example:
from k3_node.rag import GraphRAG
rag = GraphRAG(
edge_index=edge_index,
edge_type=edge_type,
entity_to_id={"Aspirin": 0, "Headache": 1, "COX-1": 2},
relation_to_id={"treats": 0, "inhibits": 1},
llm_dim=4096,
num_prefix_tokens=4,
)
# Text-based GraphRAG prompt
prompt = rag.build_prompt("What does Aspirin treat?", model_family="llama3")
# Soft prompt prefix encoding
subgraph = rag.retrieve(["Aspirin"], num_hops=1)
prefix_embeds = rag.encode_prefix(subgraph) # (1, 4, 4096)