makeitwork / src /search_agent.py
Reizxn's picture
Upload folder using huggingface_hub
3738348 verified
Raw
History Blame Contribute Delete
14.6 kB
"""
Inference-time search agent harness.
This is the runtime that executes the search loop at inference time.
It:
1. Takes a query from the calling (bigger) model
2. Formats it as a chat message for the agent model
3. Generates the agent's response (which contains <|search|> actions)
4. Parses the search query from the response
5. Retrieves code chunks from the index (simple keyword/TF-IDF search)
6. Feeds the results back to the agent as <|result|> messages
7. The agent generates more reasoning or <|evidence|>/<|finish|>
8. Returns the evidence package to the caller
The retrieval backend is a simple in-memory keyword search over the
chunked corpus. This can be replaced with any retrieval backend
(embedding search, BM25, etc.) β€” the agent interface is the same.
Usage:
python src/search_agent.py --query "How does nginx handle connections?"
python src/search_agent.py --interactive
"""
import argparse
import json
import os
import re
import sys
from collections import Counter
import torch
import torch.nn.functional as F
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from model import ModelConfig, Retriever500M
from tokenizers import Tokenizer
# ─── Paths ───────────────────────────────────────────────────────────────────
PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CHECKPOINT_DIR = os.path.join(PROJECT_DIR, "checkpoints")
TOKENIZER_PATH = os.path.join(PROJECT_DIR, "tokenizer", "tokenizer_agent.json")
CHUNKS_PATH = os.path.join(PROJECT_DIR, "data", "chunks.jsonl")
# ─── Special tokens ──────────────────────────────────────────────────────────
SYSTEM_PROMPT = (
"You are a code search agent. Given a query from a reasoning model, "
"decompose it into subqueries, search the codebase, inspect results, "
"and return curated evidence. Use <|search|> to issue searches, "
"<|reasoning|> to analyze, and <|evidence|> to return findings. "
"Be concise. Extract only the relevant facts. End with <|finish|>."
)
# ─── Simple retrieval backend ────────────────────────────────────────────────
class KeywordRetriever:
"""Simple keyword-based retrieval over code chunks.
For production, replace this with an embedding-based retriever
(e.g., mxbai-embed-large or similar). The agent interface stays the same.
"""
def __init__(self, chunks_path: str):
print(f"Loading chunks from {chunks_path}...")
self.chunks = []
with open(chunks_path, "r", encoding="utf-8") as f:
for line in f:
self.chunks.append(json.loads(line))
print(f" Loaded {len(self.chunks):,} chunks")
# Build simple term frequency index
self.chunk_tokens = []
for chunk in self.chunks:
code = chunk["code"].lower()
# Simple tokenization: split on non-alphanumeric
tokens = re.findall(r"[a-z_][a-z0-9_]*", code)
self.chunk_tokens.append(Counter(tokens))
def search(self, query: str, top_k: int = 3) -> list[dict]:
"""Search for chunks matching the query. Returns top_k results."""
query_tokens = re.findall(r"[a-z_][a-z0-9_]*", query.lower())
if not query_tokens:
return []
scores = []
for i, chunk_tf in enumerate(self.chunk_tokens):
score = sum(chunk_tf.get(t, 0) for t in query_tokens)
# Normalize by chunk length to avoid bias toward long chunks
if sum(chunk_tf.values()) > 0:
score = score / (1 + sum(chunk_tf.values()) * 0.001)
scores.append((score, i))
scores.sort(reverse=True)
results = []
for score, idx in scores[:top_k]:
if score > 0:
chunk = self.chunks[idx]
results.append({
"code": chunk["code"],
"filepath": chunk["filepath"],
"name": chunk["name"],
"type": chunk["type"],
"language": chunk["language"],
"score": score,
})
return results
# ─── Agent harness ───────────────────────────────────────────────────────────
class SearchAgent:
"""The search agent harness that runs the search loop."""
def __init__(
self,
checkpoint_path: str,
tokenizer_path: str,
chunks_path: str,
device: torch.device = None,
max_search_rounds: int = 5,
max_new_tokens: int = 256,
):
self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.max_search_rounds = max_search_rounds
self.max_new_tokens = max_new_tokens
# Load tokenizer
self.tokenizer = Tokenizer.from_file(tokenizer_path)
# Load model
print(f"Loading model from {checkpoint_path}...")
ckpt = torch.load(checkpoint_path, map_location=self.device, weights_only=False)
config = ModelConfig(**ckpt["config"])
self.model = Retriever500M(config).to(self.device)
self.model.load_state_dict(ckpt["model_state_dict"])
self.model.eval()
print(f" Loaded (step {ckpt.get('step', '?')}, loss {ckpt.get('loss', '?')})")
# Load retriever
self.retriever = KeywordRetriever(chunks_path)
# Special token IDs
vocab = self.tokenizer.get_vocab()
self.system_id = vocab.get("<tool_call>", 32000)
self.user_id = vocab.get("<tool_call>", 32001)
self.assistant_id = vocab.get("<tool_call>", 32002)
self.search_id = vocab.get("<|search|>", 32003)
self.result_id = vocab.get("<|result|>", 32004)
self.evidence_id = vocab.get("<|evidence|>", 32005)
self.reasoning_id = vocab.get("<|reasoning|>", 32006)
self.finish_id = vocab.get("<|finish|>", 32007)
self.end_id = vocab.get("<|end|>", 32008)
def _encode(self, text: str) -> list[int]:
"""Encode text to token IDs."""
return self.tokenizer.encode(text).ids
def _decode(self, ids: list[int]) -> str:
"""Decode token IDs to text."""
return self.tokenizer.decode(ids)
def _generate(self, input_ids: torch.Tensor, max_new_tokens: int) -> str:
"""Generate text from the model, stopping at <|end|> or <|finish|>."""
with torch.no_grad():
for _ in range(max_new_tokens):
# Crop context if too long
if input_ids.size(1) > self.model.config.max_seq_len:
input_ids = input_ids[:, -self.model.config.max_seq_len:]
logits = self.model(input_ids)["logits"]
next_logits = logits[:, -1, :]
# Apply temperature and sample
probs = F.softmax(next_logits / 0.8, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
input_ids = torch.cat([input_ids, next_token], dim=1)
# Stop on <|end|> or <|finish|>
if next_token.item() == self.end_id or next_token.item() == self.finish_id:
break
# Decode the generated part (after the input)
generated_ids = input_ids[0, -max_new_tokens:].tolist()
return self._decode(generated_ids)
def _parse_search_query(self, text: str) -> str | None:
"""Extract the search query from the agent's response."""
# Look for <|search|>query<|end|>
match = re.search(r"<\|search\|>(.*?)<\|end\|>", text, re.DOTALL)
if match:
return match.group(1).strip()
return None
def _parse_evidence(self, text: str) -> str | None:
"""Extract the evidence from the agent's response."""
match = re.search(r"<\|evidence\|>(.*?)(?:<\|end\|>|<\|finish\|>|$)", text, re.DOTALL)
if match:
return match.group(1).strip()
return None
def _has_finish(self, text: str) -> bool:
"""Check if the agent has signaled completion."""
return "<|finish|>" in text
def search(self, query: str) -> dict:
"""Run the full search loop for a query.
Returns:
{
"query": the original query,
"evidence": the curated evidence (or None if not found),
"searches": list of search queries issued,
"results": list of all results retrieved,
"trace": the full conversation trace,
}
"""
print(f"\n{'='*60}")
print(f"QUERY: {query}")
print(f"{'='*60}")
# Build initial context
trace = []
# System prompt
system_tokens = [self.system_id] + self._encode(SYSTEM_PROMPT) + [self.end_id]
trace.append({"role": "system", "tokens": system_tokens})
# User query
user_tokens = [self.user_id] + self._encode(query) + [self.end_id]
trace.append({"role": "user", "tokens": user_tokens})
all_searches = []
all_results = []
evidence = None
for round_num in range(self.max_search_rounds):
# Build input from trace
all_tokens = []
for entry in trace:
all_tokens.extend(entry["tokens"])
input_ids = torch.tensor([all_tokens], dtype=torch.long, device=self.device)
# Generate agent response
print(f"\n--- Round {round_num + 1} ---")
response = self._generate(input_ids, self.max_new_tokens)
print(f"Agent: {response[:200]}...")
# Add assistant tokens to trace
assistant_tokens = [self.assistant_id] + self._encode(response)
if not response.endswith("<|end|>"):
assistant_tokens.append(self.end_id)
trace.append({"role": "assistant", "tokens": assistant_tokens})
# Check for finish
if self._has_finish(response):
evidence = self._parse_evidence(response)
print(f"\n[EVIDENCE]: {evidence}")
break
# Parse search query
search_query = self._parse_search_query(response)
if search_query:
print(f"[SEARCH]: {search_query}")
all_searches.append(search_query)
# Retrieve results
results = self.retriever.search(search_query, top_k=3)
if results:
for result in results:
print(f" [RESULT]: {result['name']} ({result['language']}, score={result['score']:.2f})")
result_tokens = [self.result_id] + self._encode(result["code"]) + [self.end_id]
trace.append({"role": "result", "tokens": result_tokens, "data": result})
all_results.append(result)
else:
print(" [NO RESULTS]")
result_tokens = [self.result_id, self.end_id]
trace.append({"role": "result", "tokens": result_tokens})
else:
# No search query found β€” try to extract evidence directly
evidence = self._parse_evidence(response)
if evidence:
print(f"\n[EVIDENCE]: {evidence}")
break
else:
print("[WARNING] No search or evidence found, continuing...")
return {
"query": query,
"evidence": evidence,
"searches": all_searches,
"results": all_results,
"trace": trace,
}
# ─── CLI ─────────────────────────────────────────────────────────────────────
def main():
parser = argparse.ArgumentParser(description="Run the search agent")
parser.add_argument("--checkpoint", type=str, default=os.path.join(CHECKPOINT_DIR, "sft_latest.pt"))
parser.add_argument("--query", type=str, default=None, help="Query to search for")
parser.add_argument("--interactive", action="store_true", help="Interactive mode")
parser.add_argument("--max_rounds", type=int, default=5, help="Max search rounds")
args = parser.parse_args()
agent = SearchAgent(
checkpoint_path=args.checkpoint,
tokenizer_path=TOKENIZER_PATH,
chunks_path=CHUNKS_PATH,
max_search_rounds=args.max_rounds,
)
if args.interactive:
print("\nInteractive mode. Type 'quit' to exit.")
while True:
query = input("\nQuery> ").strip()
if query.lower() in ("quit", "exit", "q"):
break
if query:
result = agent.search(query)
print(f"\n{'='*60}")
print(f"FINAL EVIDENCE:")
print(f"{'='*60}")
print(result["evidence"] or "No evidence found.")
elif args.query:
result = agent.search(args.query)
print(f"\n{'='*60}")
print(f"FINAL EVIDENCE:")
print(f"{'='*60}")
print(result["evidence"] or "No evidence found.")
else:
# Run sample queries
sample_queries = [
"How does nginx handle reusable connections?",
"What does the with_params_help decorator do?",
"What fields does the ngx_listening_s struct have?",
"How does the concatenate function work?",
"Where is the database connection pool implemented?",
]
for query in sample_queries:
result = agent.search(query)
print(f"\n{'='*60}")
print(f"FINAL EVIDENCE:")
print(f"{'='*60}")
print(result["evidence"] or "No evidence found.")
if __name__ == "__main__":
main()