File size: 8,443 Bytes
6a55750 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 | """
Inference Script: RAG-Augmented Code Generation with Tool Calling
Combines:
- Fine-tuned Qwen2.5-Coder-7B (code gen + tool-calling)
- RAG pipeline (AST-aware code search)
- ReAct-style reasoning loop for multi-step tasks
Usage:
python inference.py --model your-username/qwen25-coder-7b-code-toolcall --repo /path/to/codebase
python inference.py --model your-username/qwen25-coder-7b-code-toolcall --repo /path/to/codebase \
--query "Implement a caching layer for the user service"
"""
import os, json, re, subprocess
from typing import Optional
from rag_pipeline import CodebaseIndexer, CodeRetriever, ContextBuilder
TOOLS = [
{"type": "function", "function": {"name": "search_codebase",
"description": "Search the internal codebase using semantic search.",
"parameters": {"type": "object", "properties": {
"query": {"type": "string"}, "file_pattern": {"type": "string"},
"max_results": {"type": "integer", "default": 5}}, "required": ["query"]}}},
{"type": "function", "function": {"name": "execute_python",
"description": "Execute Python code in a sandboxed environment.",
"parameters": {"type": "object", "properties": {
"code": {"type": "string"}, "timeout": {"type": "integer", "default": 30}},
"required": ["code"]}}},
{"type": "function", "function": {"name": "read_file",
"description": "Read contents of a file from the codebase.",
"parameters": {"type": "object", "properties": {
"path": {"type": "string"}, "start_line": {"type": "integer"},
"end_line": {"type": "integer"}}, "required": ["path"]}}},
{"type": "function", "function": {"name": "run_tests",
"description": "Run unit tests for a module or file.",
"parameters": {"type": "object", "properties": {
"test_path": {"type": "string"}, "verbose": {"type": "boolean", "default": True}},
"required": ["test_path"]}}}
]
SYSTEM_PROMPT = '''You are an expert Python programmer with access to our internal codebase via tools.
When helping with code tasks:
1. ALWAYS search the codebase first to understand existing patterns
2. Reference and integrate with existing code
3. Follow the codebase's conventions
4. Write comprehensive docstrings and type hints
5. Consider edge cases and error handling
6. For complex tasks, break them into steps and use tools iteratively
Think step by step.'''
class ToolExecutor:
def __init__(self, retriever: CodeRetriever, repo_path: str):
self.retriever = retriever
self.repo_path = repo_path
def execute(self, tool_name: str, arguments: dict) -> str:
handlers = {"search_codebase": self._search, "execute_python": self._execute,
"read_file": self._read, "run_tests": self._test}
return handlers.get(tool_name, lambda **kw: f"Unknown tool: {tool_name}")(**arguments)
def _search(self, query, file_pattern=None, max_results=5):
import fnmatch
results = self.retriever.search(query, top_k=max_results)
output = [{"file": c.file_path, "name": c.name, "type": c.chunk_type,
"score": round(s, 3), "signature": c.signature,
"content": c.content[:500]} for c, s in results
if not file_pattern or fnmatch.fnmatch(c.file_path, file_pattern)]
return json.dumps({"results": output}, indent=2)
def _execute(self, code, timeout=30):
try:
r = subprocess.run(["python", "-c", code], capture_output=True, text=True,
timeout=timeout, cwd=self.repo_path)
return (r.stdout or "") + (r.stderr or "") or "Success (no output)"
except subprocess.TimeoutExpired:
return f"Timed out after {timeout}s"
def _read(self, path, start_line=None, end_line=None):
fp = os.path.join(self.repo_path, path)
if not os.path.exists(fp): return f"File not found: {path}"
with open(fp) as f: lines = f.readlines()
if start_line: lines = lines[start_line-1:end_line]
return "".join(lines)
def _test(self, test_path, verbose=True):
cmd = ["python", "-m", "pytest", test_path] + (["-v"] if verbose else [])
try:
r = subprocess.run(cmd, capture_output=True, text=True, timeout=120, cwd=self.repo_path)
return r.stdout + r.stderr
except subprocess.TimeoutExpired:
return "Tests timed out"
class CodeAgent:
"""ReAct-style agent combining fine-tuned LLM with tool execution."""
def __init__(self, model_id, retriever, repo_path, max_turns=10, device="auto"):
self.tool_executor = ToolExecutor(retriever, repo_path)
self.context_builder = ContextBuilder(retriever)
self.max_turns = max_turns
from transformers import AutoModelForCausalLM, AutoTokenizer
self.tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
self.model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype="auto",
device_map=device, trust_remote_code=True)
def run(self, user_query, current_file=None):
rag_context = self.context_builder.build_context(user_query)
user_content = user_query
if rag_context:
user_content += f"\\n\\n--- Code context ---\\n{rag_context}\\n--- End ---"
messages = [{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": user_content}]
for turn in range(self.max_turns):
response = self._generate(messages)
tool_calls = self._extract_tool_calls(response)
if not tool_calls:
return response
messages.append({"role": "assistant", "content": response})
for tc in tool_calls:
result = self.tool_executor.execute(tc["name"], tc["arguments"])
messages.append({"role": "tool", "name": tc["name"], "content": result[:3000]})
return response
def _generate(self, messages):
import torch
text = self.tokenizer.apply_chat_template(messages, tools=TOOLS, tokenize=False,
add_generation_prompt=True)
inputs = self.tokenizer(text, return_tensors="pt").to(self.model.device)
with torch.no_grad():
outputs = self.model.generate(**inputs, max_new_tokens=4096, temperature=0.7,
top_p=0.9, do_sample=True)
return self.tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:],
skip_special_tokens=True).strip()
def _extract_tool_calls(self, response):
tool_calls = []
for pattern in [r'\\{"name":\\s*"(\\w+)",\\s*"arguments":\\s*(\\{[^}]+\\})\\}',
r'<tool_call>\\n?(.*?)\\n?</tool_call>']:
for m in re.finditer(pattern, response, re.DOTALL):
try:
if m.lastindex == 2:
tool_calls.append({"name": m.group(1), "arguments": json.loads(m.group(2))})
else:
data = json.loads(m.group(1))
if "name" in data: tool_calls.append(data)
except: pass
return tool_calls
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True)
parser.add_argument("--repo", required=True)
parser.add_argument("--query", "-q", default=None)
parser.add_argument("--index-dir", default=None)
parser.add_argument("--embedding-model", default="jinaai/jina-embeddings-v2-base-code")
args = parser.parse_args()
if args.index_dir and os.path.exists(os.path.join(args.index_dir, "chunks.json")):
retriever = CodeRetriever(args.embedding_model)
retriever.load_index(args.index_dir)
else:
retriever = CodebaseIndexer(args.repo, embedding_model=args.embedding_model).index()
if args.index_dir: retriever.save_index(args.index_dir)
agent = CodeAgent(args.model, retriever, args.repo)
if args.query:
print(agent.run(args.query))
else:
print("Code Assistant (type 'quit' to exit)")
while True:
q = input("> ").strip()
if q.lower() in ("quit", "exit", "q"): break
if q: print(agent.run(q))
|