Skip to content

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 (2, num_edges).

required
edge_type Optional[Union[ndarray, any]]

Optional 1D relation type tensor of shape (num_edges,).

None
edge_attr Optional[Union[ndarray, any]]

Optional edge attribute tensor of shape (num_edges, edge_dim).

None
x Optional[Union[ndarray, any]]

Optional node feature tensor of shape (total_nodes, in_channels).

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 0..num_subgraph_nodes-1. (default: True)

True
num_nodes Optional[int]

Optional total number of nodes in graph. Inferred if not given.

None

Returns:

Type Description
SubgraphResult

SubgraphResult containing relabeled edge_index, edge_type, features,

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 (2, num_edges).

required
edge_type Optional[Union[ndarray, any]]

Optional relation type tensor of shape (num_edges,).

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']
  • "triples": List of (Head, Relation, Tail) text triples.
  • "markdown": Markdown bullet list with facts.
  • "natural": Natural language sentences ("Head relation Tail.").
'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": Global average over all subgraph nodes. - "sum": Global sum over all subgraph nodes. - "max": Global maximum over all subgraph nodes. - "center": Pool only the retrieved center seed entities. - "none": Return all node embeddings without pooling.

'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 (num_nodes, in_channels).

required
edge_index any

Graph edge indices of shape (2, num_edges).

required
edge_type Optional[any]

1D tensor of relation IDs for each edge (num_edges,).

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 (1, out_channels) (or (num_nodes, out_channels) if pooling="none").

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 KGEModel (e.g. TransE). If provided, its embeddings are reused.

None
pooling Literal['mean', 'sum', 'center']

Readout pooling strategy across entities ("mean", "center", "sum").

'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 (num_nodes,).

required
edge_index Optional[any]

Optional edge index tensor (2, num_edges).

None
edge_type Optional[any]

Optional relation type tensor (num_edges,).

None
center_nodes Optional[any]

Optional indices of seed entities within subgraph_nodes.

None

Returns:

Type Description
any

Dense embedding tensor of shape (1, out_channels).

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": 2-layer MLP with LayerNorm and GELU activation. - "linear": Single linear transformation.

'mlp'
hidden_dim Optional[int]

Optional hidden dimension for MLP projector. (default: 2 * in_channels)

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 (batch_size, in_channels) or (in_channels,).

required

Returns:

Type Description
any

Prefix tensor of shape (batch_size, num_prefix_tokens, llm_dim).

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., RGCNSubGraphEncoder, TransEPrefixEncoder).

required
projector GraphPrefixProjector

GraphPrefixProjector instance.

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 (1, num_prefix_tokens, llm_dim).

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 (batch, seq_len).

required
num_prefix_tokens int

Number of prefix tokens prepended.

required

Returns:

Type Description
any

Extended mask of shape (batch, num_prefix_tokens + seq_len).

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 (batch, seq_len, llm_dim).

required
prefix_embeddings any

Tensor of shape (batch, num_prefix_tokens, llm_dim).

required

Returns:

Type Description
any

Concatenated tensor of shape (batch, num_prefix_tokens + seq_len, llm_dim).

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 (2, num_edges).

required
edge_type Optional[any]

Optional relation IDs tensor of shape (num_edges,).

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", "transe", or a custom encoder).

'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 (1, num_prefix_tokens, llm_dim).

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)