Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 48 additions & 60 deletions src/axon/core/embeddings/embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,21 +11,28 @@

from __future__ import annotations

import logging
import threading
from typing import TYPE_CHECKING

from axon.core.embeddings.text import build_class_method_index, generate_text
from axon.core.graph.graph import KnowledgeGraph
from axon.core.graph.model import NodeLabel
from axon.core.graph.model import GraphNode, NodeLabel
from axon.core.storage.base import NodeEmbedding

if TYPE_CHECKING:
from fastembed import TextEmbedding

logger = logging.getLogger(__name__)

_model_cache: dict[str, "TextEmbedding"] = {}
_model_lock = threading.Lock()

# BGE-small max sequence is 512 tokens (~2000 chars). Truncating long
# descriptions avoids wasting tokenisation and padding time on text that
# the model would discard anyway.
_MAX_TEXT_CHARS = 2000


def _get_model(model_name: str) -> "TextEmbedding":
cached = _model_cache.get(model_name)
Expand All @@ -36,7 +43,9 @@ def _get_model(model_name: str) -> "TextEmbedding":
if cached is not None:
return cached
from fastembed import TextEmbedding
model = TextEmbedding(model_name=model_name)
# threads=0 lets ONNX Runtime use all CPU cores for intra-op
# parallelism, which significantly speeds up batch inference.
model = TextEmbedding(model_name=model_name, threads=0)
_model_cache[model_name] = model
return model

Expand Down Expand Up @@ -76,10 +85,42 @@ def embed_query(query: str, model_name: str = _DEFAULT_MODEL) -> list[float] | N
return None


def _embed_node_list(
nodes: list[GraphNode],
graph: KnowledgeGraph,
model_name: str,
batch_size: int,
) -> list[NodeEmbedding]:
"""Shared implementation for embedding a list of graph nodes.

Generates text descriptions, truncates to model context window,
encodes via fastembed, and returns :class:`NodeEmbedding` objects.
"""
class_method_idx = build_class_method_index(graph)

texts: list[str] = []
valid_nodes: list[GraphNode] = []
for node in nodes:
text = generate_text(node, graph, class_method_idx)
if text and text.strip():
texts.append(text[:_MAX_TEXT_CHARS])
valid_nodes.append(node)

if not texts:
return []

logger.info("Embedding %d texts (batch_size=%d) …", len(texts), batch_size)
model = _get_model(model_name)
return [
NodeEmbedding(node_id=node.id, embedding=vector.tolist())
for node, vector in zip(valid_nodes, model.embed(texts, batch_size=batch_size))
]


def embed_graph(
graph: KnowledgeGraph,
model_name: str = "BAAI/bge-small-en-v1.5",
batch_size: int = 64,
batch_size: int = 256,
) -> list[NodeEmbedding]:
"""Generate embeddings for all embeddable nodes in the graph.

Expand All @@ -91,83 +132,30 @@ def embed_graph(
graph: The knowledge graph whose nodes should be embedded.
model_name: The fastembed model identifier. Defaults to
``"BAAI/bge-small-en-v1.5"``.
batch_size: Number of texts to encode per batch. Defaults to 64.
batch_size: Number of texts to encode per batch. Defaults to 256.

Returns:
A list of :class:`NodeEmbedding` instances, one per embeddable node,
each carrying the node's ID and its embedding vector as a plain
Python ``list[float]``.
"""
all_nodes = [n for n in graph.iter_nodes() if n.label in EMBEDDABLE_LABELS]

if not all_nodes:
return []

class_method_idx = build_class_method_index(graph)

texts: list[str] = []
nodes = []
for node in all_nodes:
text = generate_text(node, graph, class_method_idx)
if text and text.strip():
texts.append(text)
nodes.append(node)

if not texts:
return []

model = _get_model(model_name)
vectors = list(model.embed(texts, batch_size=batch_size))

results: list[NodeEmbedding] = []
for node, vector in zip(nodes, vectors):
results.append(
NodeEmbedding(
node_id=node.id,
embedding=vector.tolist(),
)
)

return results
return _embed_node_list(all_nodes, graph, model_name, batch_size)


def embed_nodes(
graph: KnowledgeGraph,
node_ids: set[str],
model_name: str = "BAAI/bge-small-en-v1.5",
batch_size: int = 64,
batch_size: int = 256,
) -> list[NodeEmbedding]:
"""Like :func:`embed_graph`, but only for the given *node_ids*."""
if not node_ids:
return []

nodes = [graph.get_node(nid) for nid in node_ids]
nodes = [n for n in nodes if n is not None and n.label in EMBEDDABLE_LABELS]

if not nodes:
return []

class_method_idx = build_class_method_index(graph)

texts: list[str] = []
valid_nodes = []
for node in nodes:
text = generate_text(node, graph, class_method_idx)
if text and text.strip():
texts.append(text)
valid_nodes.append(node)

if not texts:
return []

model = _get_model(model_name)
embeddings: list[NodeEmbedding] = []
for node, vector in zip(valid_nodes, model.embed(texts, batch_size=batch_size)):
embeddings.append(
NodeEmbedding(
node_id=node.id,
embedding=vector.tolist(),
)
)

return embeddings
return _embed_node_list(nodes, graph, model_name, batch_size)
163 changes: 89 additions & 74 deletions src/axon/core/ingestion/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
import time
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path

Expand Down Expand Up @@ -120,72 +121,84 @@ def run_pipeline(
"""
start = time.monotonic()
result = PipelineResult()
log = logging.getLogger(__name__)

phase_times: dict[str, float] = {}

def report(phase: str, pct: float) -> None:
if progress_callback is not None:
progress_callback(phase, pct)

report("Walking files", 0.0)
gitignore = load_gitignore(repo_path)
files = walk_repo(repo_path, gitignore)
result.files = len(files)
report("Walking files", 1.0)
@contextmanager
def _timed(phase_name: str):
"""Context manager that logs and records phase wall-clock time."""
report(phase_name, 0.0)
t0 = time.monotonic()
try:
yield
finally:
elapsed = time.monotonic() - t0
phase_times[phase_name] = elapsed
log.info("Phase %-30s %.2fs", phase_name, elapsed)
report(phase_name, 1.0)

with _timed("Walking files"):
gitignore = load_gitignore(repo_path)
files = walk_repo(repo_path, gitignore)
result.files = len(files)

graph = KnowledgeGraph()

report("Processing structure", 0.0)
process_structure(files, graph)
report("Processing structure", 1.0)
with _timed("Processing structure"):
process_structure(files, graph)

report("Parsing code", 0.0)
parse_data = process_parsing(files, graph)
report("Parsing code", 1.0)
with _timed("Parsing code"):
parse_data = process_parsing(files, graph)

report("Resolving imports", 0.0)
process_imports(parse_data, graph, parallel=True)
report("Resolving imports", 1.0)
with _timed("Resolving imports"):
process_imports(parse_data, graph, parallel=True)

shared_labels = (
NodeLabel.FUNCTION, NodeLabel.METHOD, NodeLabel.CLASS,
NodeLabel.INTERFACE, NodeLabel.TYPE_ALIAS,
)
shared_name_index = build_name_index(graph, shared_labels)
heritage_labels = {NodeLabel.CLASS, NodeLabel.INTERFACE}
heritage_name_index: dict[str, list[str]] = {}
for name, ids in shared_name_index.items():
filtered = [
nid for nid in ids
if (n := graph.get_node(nid)) is not None and n.label in heritage_labels
]
if filtered:
heritage_name_index[name] = filtered

report("Resolving relationships", 0.0)
with ThreadPoolExecutor(max_workers=3) as pool:
calls_f = pool.submit(
process_calls, parse_data, graph,
name_index=shared_name_index, parallel=False, collect=True,
)
heritage_f = pool.submit(
process_heritage, parse_data, graph,
name_index=heritage_name_index, parallel=False, collect=True,
)
types_f = pool.submit(
process_types, parse_data, graph,
name_index=shared_name_index, parallel=False, collect=True,
with _timed("Building indexes"):
shared_labels = (
NodeLabel.FUNCTION, NodeLabel.METHOD, NodeLabel.CLASS,
NodeLabel.INTERFACE, NodeLabel.TYPE_ALIAS,
)
shared_name_index = build_name_index(graph, shared_labels)
heritage_labels = {NodeLabel.CLASS, NodeLabel.INTERFACE}
heritage_name_index: dict[str, list[str]] = {}
for name, ids in shared_name_index.items():
filtered = [
nid for nid in ids
if (n := graph.get_node(nid)) is not None and n.label in heritage_labels
]
if filtered:
heritage_name_index[name] = filtered

with _timed("Resolving relationships"):
with ThreadPoolExecutor(max_workers=3) as pool:
calls_f = pool.submit(
process_calls, parse_data, graph,
name_index=shared_name_index, parallel=False, collect=True,
)
heritage_f = pool.submit(
process_heritage, parse_data, graph,
name_index=heritage_name_index, parallel=False, collect=True,
)
types_f = pool.submit(
process_types, parse_data, graph,
name_index=shared_name_index, parallel=False, collect=True,
)

_write_collected_edges(calls_f.result() or [], graph)
_write_collected_edges(calls_f.result() or [], graph)

heritage_edges, heritage_patches = heritage_f.result()
_write_collected_edges(heritage_edges, graph)
for patch in heritage_patches:
node = graph.get_node(patch.node_id)
if node is not None:
node.properties[patch.key] = patch.value
heritage_edges, heritage_patches = heritage_f.result()
_write_collected_edges(heritage_edges, graph)
for patch in heritage_patches:
node = graph.get_node(patch.node_id)
if node is not None:
node.properties[patch.key] = patch.value

_write_collected_edges(types_f.result() or [], graph)
report("Resolving relationships", 1.0)
_write_collected_edges(types_f.result() or [], graph)

coupling_file_nodes = graph.get_nodes_by_label(NodeLabel.FILE)

Expand All @@ -194,48 +207,50 @@ def report(phase: str, pct: float) -> None:
resolve_coupling, graph, repo_path, file_nodes=coupling_file_nodes,
)

report("Detecting communities", 0.0)
result.clusters = process_communities(graph)
report("Detecting communities", 1.0)
with _timed("Detecting communities"):
result.clusters = process_communities(graph)

report("Detecting execution flows", 0.0)
result.processes = process_processes(graph)
report("Detecting execution flows", 1.0)
with _timed("Detecting execution flows"):
result.processes = process_processes(graph)

report("Finding dead code", 0.0)
result.dead_code = process_dead_code(graph)
report("Finding dead code", 1.0)
with _timed("Finding dead code"):
result.dead_code = process_dead_code(graph)

report("Analyzing git history", 0.0)
coupling_edges = coupling_future.result()
_write_collected_edges(coupling_edges, graph)
result.coupled_pairs = len(coupling_edges)
report("Analyzing git history", 1.0)
with _timed("Analyzing git history"):
coupling_edges = coupling_future.result()
_write_collected_edges(coupling_edges, graph)
result.coupled_pairs = len(coupling_edges)

result.symbols = sum(1 for n in graph.iter_nodes() if n.label in _SYMBOL_LABELS)
result.relationships = graph.relationship_count

if storage is not None:
report("Loading to storage", 0.0)
storage.bulk_load(graph)
report("Loading to storage", 1.0)
with _timed("Loading to storage"):
storage.bulk_load(graph)

if embeddings:
try:
report("Generating embeddings", 0.0)
node_embeddings = embed_graph(graph)
storage.store_embeddings(node_embeddings)
result.embeddings = len(node_embeddings)
report("Generating embeddings", 1.0)
with _timed("Generating embeddings"):
node_embeddings = embed_graph(graph)
storage.store_embeddings(node_embeddings)
result.embeddings = len(node_embeddings)
except Exception:
logging.getLogger(__name__).warning(
log.warning(
"Embedding phase failed — search will use FTS only",
exc_info=True,
)
report("Generating embeddings", 1.0)

result.duration_seconds = time.monotonic() - start

# Log phase breakdown summary
if phase_times:
log.info("─── Phase timing breakdown ───")
for phase, elapsed in phase_times.items():
pct = (elapsed / result.duration_seconds * 100) if result.duration_seconds > 0 else 0
log.info(" %-30s %6.1fs (%4.1f%%)", phase, elapsed, pct)
log.info(" %-30s %6.1fs", "TOTAL", result.duration_seconds)

return graph, result

def reindex_files(
Expand Down
Loading