Team Ai
Apppublic

lablab-ai-amd-developer-hackathon/gpu-goblin

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
loop.py260 linesDownload Raw Back to agent
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