lablab-ai-amd-developer-hackathon/gpu-goblin
0
1"""Agent loop driver — provider-agnostic tool-use loop for one audit.2 3`run_audit(file_path)` is an async generator that yields `SSEEvent` objects in4the order the UI should render them: thoughts, tool calls, tool results, and5finally either a `final_report` event (extracted from the most recent6successful `compare_runs` tool result) or an `error` event.7 8The loop itself doesn't know about Anthropic or Hugging Face — it talks to9whichever `Backend` `make_backend()` returns. The backend (Claude or Qwen-HF10today) handles all per-API translation. See `agent/backends/__init__.py`.11"""12 13from __future__ import annotations14 15import json16from collections.abc import AsyncIterator17from typing import Any18 19from agent import tools as tools_module20from agent.backends import Backend, ToolCall, make_backend21from agent.prompts import SYSTEM_PROMPT22from agent.schemas import SSEEvent23 24MAX_STEPS = 1025"""Hard cap on tool calls per audit. The canonical trajectory is six calls26(parse → profile → query_kb → patch → benchmark×2 → compare). The extra274 calls of headroom let the model recover from common mistakes (JSON28nesting glitches, retry on ToolResult(ok=False)) without exhausting the29budget before compare_runs. Was 8; bumped after a live run hit a wall when30two misnested-arg benchmark retries ate the slack meant for compare_runs.31"""32MAX_TOKENS = 204833 34 35def _extract_final_report(36 tool_results: list[dict[str, Any]],37) -> dict[str, Any] | None:38 """Walk tool results in reverse and return the most recent successful39 compare_runs payload, or None if there isn't one."""40 for entry in reversed(tool_results):41 if entry["name"] == "compare_runs" and entry["ok"]:42 return entry["result"]43 return None44 45 46def _auto_compare(47 tool_results: list[dict[str, Any]],48) -> dict[str, Any] | None:49 """Synthesize a Report from whatever the audit produced when the model50 didn't reach `compare_runs` cleanly. Three recovery tiers, in order of51 fidelity:52 53 Tier 1 — full data: ≥2 benchmarks + ≥1 propose_patch.54 Treat first benchmark as baseline, last as patched run. Highest55 fidelity since both numbers are real.56 57 Tier 2 — patch but only one benchmark: ≥1 patch + 1 benchmark.58 Use the single benchmark as baseline. For the "after" side, run59 FakeRunner on the patched config to get a deterministic projection.60 Marks the report as projected so the demo is honest about it.61 62 Tier 3 — no patch ran but we have rules from query_rocm_kb + ≥1 benchmark.63 We *could* deterministically apply propose_patch ourselves here, but64 that's over-reaching. Return None and let the caller surface a65 clean error instead.66 67 Returns the Report dict, or None when no tier applies.68 """69 benchmarks = [70 e for e in tool_results if e["name"] == "benchmark" and e["ok"]71 ]72 patches = [73 e for e in tool_results if e["name"] == "propose_patch" and e["ok"]74 ]75 76 # Tier 1: full data path.77 if len(benchmarks) >= 2 and patches:78 latest_patch = patches[-1]["result"]79 before = benchmarks[0]["result"]80 after = benchmarks[-1]["result"]81 return _call_compare_runs(latest_patch, before, after, " (auto-synthesized compare_runs)")82 83 # Tier 2: patch + 1 benchmark — fill in the patched-side metrics from84 # FakeRunner so the demo still produces a Report with a clear note.85 if patches and len(benchmarks) == 1:86 latest_patch = patches[-1]["result"]87 before = benchmarks[0]["result"]88 # Project the patched run via FakeRunner. The synthetic corpus has89 # a `02_optimized` scenario the patched config typically matches.90 from agent.schemas import WorkloadConfig91 from runner.protocol import FakeRunner92 93 try:94 patched_cfg = WorkloadConfig.model_validate(latest_patch["new_config"])95 after_metrics = FakeRunner().run(patched_cfg, steps=before.get("steps", 50))96 after = after_metrics.model_dump()97 except Exception:98 return None99 return _call_compare_runs(100 latest_patch,101 before,102 after,103 " (auto-synthesized; patched-side projected via FakeRunner)",104 )105 106 return None107 108 109def _call_compare_runs(110 patch: dict[str, Any],111 before: dict[str, Any],112 after: dict[str, Any],113 suffix: str,114) -> dict[str, Any] | None:115 workload_name = (116 patch.get("new_config", {}).get("model_name")117 or "Audited Workload"118 ) + suffix119 result = tools_module.call(120 "compare_runs",121 workload_name=workload_name,122 before=before,123 after=after,124 patch=patch,125 )126 return result.result if result.ok else None127 128 129def _safe_json(value: Any) -> str:130 """Serialize a tool result for inclusion in a tool_result content block.131 132 Falls back to ``str(value)`` if json can't represent the value (e.g. a133 Pydantic model already coerced upstream — shouldn't happen, but defensive).134 """135 try:136 return json.dumps(value, default=str)137 except Exception:138 return str(value)139 140 141async def _drive(backend: Backend) -> AsyncIterator[SSEEvent]:142 """Pure orchestration loop. Backend handles per-API state; we yield events."""143 tool_results_log: list[dict[str, Any]] = []144 145 for _step in range(MAX_STEPS):146 turn = await backend.next_turn(tools_module.tool_schemas())147 148 for text in turn.text_blocks:149 if text:150 yield SSEEvent(type="thought", data={"text": text})151 152 for tc in turn.tool_calls:153 async for ev in _execute_tool_call(backend, tc, tool_results_log):154 yield ev155 156 if turn.stop_reason == "end_turn":157 break158 159 report = _extract_final_report(tool_results_log)160 if report is not None:161 yield SSEEvent(type="final_report", data={"report": report})162 return163 164 # Fallback: the model didn't call compare_runs (or its tool_call landed165 # inside a thinking block where the parser couldn't extract it).166 # Synthesize the report deterministically from the tool log if we have167 # enough material. See _auto_compare for the prerequisites.168 auto = _auto_compare(tool_results_log)169 if auto is not None:170 yield SSEEvent(171 type="thought",172 data={173 "text": (174 "Note: model did not emit a compare_runs tool call (likely "175 "left it inside a <think> block). Synthesizing the final "176 "report from the latest propose_patch + two benchmarks."177 )178 },179 )180 yield SSEEvent(type="final_report", data={"report": auto})181 return182 183 yield SSEEvent(184 type="error",185 data={186 "message": (187 "Audit completed without producing a final report (and "188 "auto-synthesis fallback couldn't run — need at least one "189 "successful propose_patch and two successful benchmarks)."190 )191 },192 )193 194 195async def _execute_tool_call(196 backend: Backend,197 tc: ToolCall,198 tool_results_log: list[dict[str, Any]],199) -> AsyncIterator[SSEEvent]:200 """Yield the tool_call/tool_result event pair and record the outcome."""201 yield SSEEvent(202 type="tool_call",203 data={"id": tc.id, "name": tc.name, "input": tc.input},204 )205 206 result = tools_module.call(tc.name, **tc.input)207 208 yield SSEEvent(209 type="tool_result",210 data={211 "id": tc.id,212 "name": tc.name,213 "ok": result.ok,214 "result": result.result,215 "error": result.error,216 },217 )218 219 tool_results_log.append(220 {221 "id": tc.id,222 "name": tc.name,223 "ok": result.ok,224 "result": result.result,225 "error": result.error,226 }227 )228 229 content = (230 _safe_json(result.result) if result.ok else (result.error or "tool failed")231 )232 backend.add_tool_result(233 tool_call_id=tc.id,234 name=tc.name,235 content=content,236 is_error=not result.ok,237 )238 239 240async def run_audit(file_path: str) -> AsyncIterator[SSEEvent]:241 """Run one audit and yield SSE events as they happen.242 243 Selects the LLM backend from the `GOBLIN_AGENT_BACKEND` env var (defaults244 to `claude`; `qwen` routes through HF Inference Providers). On any245 backend or loop exception, yields a single `error` SSE event and stops.246 """247 try:248 backend = make_backend(system_prompt=SYSTEM_PROMPT, max_tokens=MAX_TOKENS)249 except Exception as exc:250 yield SSEEvent(type="error", data={"message": str(exc)})251 return252 253 backend.add_user_message(f"Audit this fine-tuning workload: {file_path}")254 255 try:256 async for ev in _drive(backend):257 yield ev258 except Exception as exc:259 yield SSEEvent(type="error", data={"message": str(exc)})260 