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
5 changes: 4 additions & 1 deletion scripts/agent_gan_inference_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,10 @@

def _agent_cache_gate(warm_delta: dict, actual_delta: dict) -> bool:
return (
warm_delta["remote_hits"] >= 1
(
warm_delta.get("remote_hits", 0) >= 1
or warm_delta.get("local_hits", 0) >= 1
)
and warm_delta["tokens_reused"] >= 1
and warm_delta["tokens_computed"] == 0
and warm_delta["fallbacks"] == 0
Expand Down
25 changes: 23 additions & 2 deletions scripts/agent_gan_repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,23 @@ def _stage(
return stage


def _gate_failure(name: str, warm: dict, actual: dict) -> RuntimeError:
keys = (
"local_hits",
"remote_hits",
"remote_jobs",
"tokens_reused",
"tokens_computed",
"fallbacks",
"remote_job_failures",
)
compact = lambda delta: {key: delta.get(key, 0) for key in keys}
return RuntimeError(
f"{name} KV gate failed: "
f"warm={compact(warm['delta'])} actual={compact(actual['delta'])}",
)


def main() -> int:
install_signal_protection()
parser = argparse.ArgumentParser()
Expand Down Expand Up @@ -274,7 +291,11 @@ def get_stats():
generator_text,
)
if not generator_stage["ok"] and not telemetry_state["degraded"]:
raise RuntimeError("Generator KV gate failed")
raise _gate_failure(
"Generator",
generator_warm,
generator_actual,
)
if remote_run:
_telemetry_request(
f"{args.dashboard}/v1/network/benchmarks/{run_id}",
Expand Down Expand Up @@ -337,7 +358,7 @@ def get_stats():
extra_metrics=context_metrics,
)
if not critic_stage["ok"] and not telemetry_state["degraded"]:
raise RuntimeError("Critic KV gate failed")
raise _gate_failure("Critic", critic_warm, critic_actual)
completed = None
if remote_run:
completed = _telemetry_request(
Expand Down
4 changes: 4 additions & 0 deletions tests/inference_engine/bridge/test_agent_gan_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,10 @@ def test_agent_gate_accepts_remote_compute_or_exact_remote_cache_hit():
}
assert _agent_cache_gate(warm, actual)
assert _agent_cache_gate({**warm, "remote_jobs": 0}, actual)
assert _agent_cache_gate(
{**warm, "remote_hits": 0, "local_hits": 1, "remote_jobs": 0},
actual,
)
assert not _agent_cache_gate({**warm, "remote_hits": 0}, actual)
assert not _agent_cache_gate({**warm, "tokens_reused": 0}, actual)
assert not _agent_cache_gate({**warm, "fallbacks": 1}, actual)
Expand Down
13 changes: 13 additions & 0 deletions tests/inference_engine/bridge/test_agent_gan_repl.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from scripts.agent_gan_repl import (
PrefillHeartbeat,
TokenPrinter,
_gate_failure,
_stage,
_telemetry_request,
build_critic_messages,
Expand Down Expand Up @@ -160,6 +161,18 @@ def timeout(*_args, **_kwargs):
assert "inference will continue" in output


def test_gate_failure_exposes_reuse_counters():
error = _gate_failure(
"Generator",
{"delta": {"local_hits": 0, "remote_hits": 0}},
{"delta": {"local_hits": 0, "fallbacks": 1}},
)
message = str(error)
assert "Generator KV gate failed" in message
assert "'remote_hits': 0" in message
assert "'fallbacks': 1" in message


def test_interactive_prompts_are_deterministic_for_kv_reuse():
generator_a = build_generator_messages("prove RH")
generator_b = build_generator_messages("prove RH")
Expand Down
Loading