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
11 changes: 11 additions & 0 deletions docs/ops/distributed-prefill-kv-network.md
Original file line number Diff line number Diff line change
Expand Up @@ -324,6 +324,17 @@ One Generator/Critic cycle is the safe default for the current 16GB allens
Gemma worker. Additional rounds grow agent histories and may exceed the strict
remote-prefill timeout; increase worker memory/timeout before enabling them.

For an interactive terminal REPL with real-time token output:

```bash
bash scripts/run_agent_gan_repl.sh
```

Type a prompt at `prompt>`. The terminal prints allens Prefill progress, then
streams `generator>` and `critic>` tokens as they arrive, followed by KV hit
rate, decode tok/s, generation latency and E2E tok/s. `/quit` exits; services
remain running. Reports remain redacted and appear in the Benchmarks tab.

## Rollback

The cache is an optimization; inference correctness does not depend on it.
Expand Down
11 changes: 10 additions & 1 deletion scripts/agent_gan_inference_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,14 @@ def _output_metadata(text: str) -> dict:
}


def _infer(client, eos_ids, token_ids, output_tokens: int, get_stats):
def _infer(
client,
eos_ids,
token_ids,
output_tokens: int,
get_stats,
on_token=None,
):
before = get_stats()
started = time.perf_counter()
with client.create_session(eos_token_ids=eos_ids, client_label="agent-gan") as s:
Expand All @@ -45,6 +52,8 @@ def _infer(client, eos_ids, token_ids, output_tokens: int, get_stats):
generated = []
for token in s.generate(max_tokens=output_tokens):
generated.append(int(token))
if on_token is not None:
on_token(generated)
if first_at is None:
first_at = time.perf_counter()
done = time.perf_counter()
Expand Down
259 changes: 259 additions & 0 deletions scripts/agent_gan_repl.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,259 @@
#!/usr/bin/env python3
"""Interactive Generator/Critic REPL with real-time token streaming."""
from __future__ import annotations

import argparse
import hashlib
import json
import time
import uuid
from pathlib import Path

from scripts.agent_gan_inference_demo import (
_agent_cache_gate,
_infer,
)
from scripts.benchmark_prefill_architecture import (
_ensure_services,
_json_request,
)


class TokenPrinter:
def __init__(self, tokenizer, label: str) -> None:
self.tokenizer = tokenizer
self.last = ""
print(f"{label}> ", end="", flush=True)

def __call__(self, token_ids) -> None:
text = self.tokenizer.decode(token_ids, skip_special_tokens=True)
print(text[len(self.last):], end="", flush=True)
self.last = text

def finish(self) -> None:
print(flush=True)


def _stage(name: str, warm: dict, actual: dict, text: str) -> dict:
delta = actual["delta"]
return {
**actual,
"name": f"agent_{name}",
"agent": name,
"round": 1,
"hit_source": "primary_hot" if delta["local_hits"] else "unknown",
"ok": _agent_cache_gate(warm["delta"], delta),
"warmup_prefix_tokens": warm["prefix_tokens"],
"warmup_tokens_reused": (
warm["delta"]["tokens_reused"]
if warm["delta"]["remote_jobs"] == 0 else 0
),
"warmup_wall_s": warm["e2e_s"],
"warmup_remote_jobs": warm["delta"]["remote_jobs"],
"output_chars": len(text),
"output_hash": hashlib.sha256(text.encode()).hexdigest(),
}


def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--worker-ssh", default="allens")
parser.add_argument("--address", default="127.0.0.1:51051")
parser.add_argument("--dashboard", default="http://127.0.0.1:8090")
parser.add_argument("--api-key-file", default="~/.kakeya/network_api_key")
parser.add_argument("--tokenizer-id", required=True)
parser.add_argument("--output-tokens", type=int, default=64)
parser.add_argument("--skip-ensure", action="store_true")
args = parser.parse_args()

from kakeya import Client
from transformers import AutoTokenizer
from scripts.chat_grpc import _resolve_eos_token_ids

if not args.skip_ensure:
print("[startup] ensuring Primary and allens services...", flush=True)
_ensure_services(args.worker_ssh)
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_id)
eos_ids = _resolve_eos_token_ids(tokenizer)
api_key = Path(args.api_key_file).expanduser().read_text().strip()

def get_stats():
return _json_request(f"{args.dashboard}/v1/network/prefill")

print(
"Kakeya Agent GAN REPL ready. Type a prompt; /quit exits.\n"
"Each turn runs allens Prefill → Primary hot Generator → "
"allens Prefill → Primary hot Critic.",
flush=True,
)
with Client(args.address) as client:
while True:
try:
prompt = input("\nprompt> ").strip()
except EOFError:
print("\n[bye]")
break
if not prompt:
continue
if prompt.lower() in {"/quit", "/exit"}:
print("[bye]")
break
run_nonce = uuid.uuid4().hex
run = _json_request(
f"{args.dashboard}/v1/network/benchmarks",
api_key=api_key,
method="POST",
body={
"kind": "agent_gan_interactive",
"config": {
"model_id": "gemma-4-26B-A4B-it-mlx-4bit",
"topology": "primary-decode-allens-prefill",
"agents": ["generator", "critic"],
"rounds": 1,
"output_tokens": args.output_tokens,
},
},
)
run_id = run["id"]
try:
generator_messages = [
{
"role": "system",
"content": (
"You are the Generator agent. Produce a concrete, "
"technically rigorous answer. Internal run "
f"{run_nonce}."
),
},
{"role": "user", "content": prompt},
]
generator_ids = tokenizer.apply_chat_template(
generator_messages,
add_generation_prompt=True,
tokenize=True,
return_dict=False,
enable_thinking=False,
)
print(
f"[allens] Generator Prefill: {len(generator_ids)} tokens...",
flush=True,
)
_, generator_warm = _infer(
client, eos_ids, generator_ids, 1, get_stats,
)
generator_printer = TokenPrinter(tokenizer, "generator")
generator_tokens, generator_actual = _infer(
client,
eos_ids,
generator_ids,
args.output_tokens,
get_stats,
on_token=generator_printer,
)
generator_printer.finish()
generator_text = tokenizer.decode(
generator_tokens,
skip_special_tokens=True,
)
generator_stage = _stage(
"generator",
generator_warm,
generator_actual,
generator_text,
)
if not generator_stage["ok"]:
raise RuntimeError("Generator KV gate failed")
_json_request(
f"{args.dashboard}/v1/network/benchmarks/{run_id}",
api_key=api_key,
method="PATCH",
body={"stages": [generator_stage]},
)

critic_messages = [
{
"role": "system",
"content": (
"You are the Critic/Discriminator. Score the answer "
"0-10, identify false assumptions, and propose "
"specific corrections. Internal run "
f"{run_nonce}."
),
},
{
"role": "user",
"content": (
f"Original task:\n{prompt}\n\n"
f"Generator answer:\n{generator_text}"
),
},
]
critic_ids = tokenizer.apply_chat_template(
critic_messages,
add_generation_prompt=True,
tokenize=True,
return_dict=False,
enable_thinking=False,
)
print(
f"[allens] Critic Prefill: {len(critic_ids)} tokens...",
flush=True,
)
_, critic_warm = _infer(
client, eos_ids, critic_ids, 1, get_stats,
)
critic_printer = TokenPrinter(tokenizer, "critic")
critic_tokens, critic_actual = _infer(
client,
eos_ids,
critic_ids,
args.output_tokens,
get_stats,
on_token=critic_printer,
)
critic_printer.finish()
critic_text = tokenizer.decode(
critic_tokens,
skip_special_tokens=True,
)
critic_stage = _stage(
"critic",
critic_warm,
critic_actual,
critic_text,
)
if not critic_stage["ok"]:
raise RuntimeError("Critic KV gate failed")
completed = _json_request(
f"{args.dashboard}/v1/network/benchmarks/{run_id}",
api_key=api_key,
method="PATCH",
body={
"stages": [critic_stage],
"status": "completed",
"finished_at": time.time(),
},
)
summary = completed["summary"]
print(
"[metrics] "
f"KV hit={summary['workload_kv_token_hit_rate']:.1%} "
f"decode={summary['aggregate_decode_tok_s']:.2f} tok/s "
f"latency={summary['generation_latency_ms_p50']:.2f} ms/token "
f"e2e={summary['aggregate_e2e_tok_s']:.2f} tok/s "
f"run={run_id}",
flush=True,
)
except Exception as exc:
_json_request(
f"{args.dashboard}/v1/network/benchmarks/{run_id}",
api_key=api_key,
method="PATCH",
body={"status": "failed", "finished_at": time.time()},
)
print(f"[error] {type(exc).__name__}: {exc}", flush=True)
return 0


if __name__ == "__main__":
raise SystemExit(main())
11 changes: 11 additions & 0 deletions scripts/run_agent_gan_repl.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
#!/usr/bin/env bash
set -euo pipefail

REPO_ROOT="$(cd "$(dirname "$0")/.." && pwd)"
PYTHON="${KAKEYA_BENCH_PYTHON:-$HOME/.venv-distwan/bin/python}"
MODEL="${KAKEYA_BENCH_MODEL:-$HOME/kakeya-models/gemma-4-26B-A4B-it-mlx-4bit}"

exec env PYTHONPATH="$REPO_ROOT:$REPO_ROOT/sdks/python" \
"$PYTHON" "$REPO_ROOT/scripts/agent_gan_repl.py" \
--tokenizer-id "$MODEL" \
"$@"
45 changes: 45 additions & 0 deletions tests/inference_engine/bridge/test_agent_gan_repl.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
from scripts.agent_gan_repl import TokenPrinter, _stage


class Tokenizer:
def decode(self, token_ids, **_kwargs):
return "".join(chr(96 + token) for token in token_ids)


def test_token_printer_streams_only_new_suffix(capsys):
printer = TokenPrinter(Tokenizer(), "generator")
printer([1])
printer([1, 2])
printer.finish()
assert capsys.readouterr().out == "generator> ab\n"


def test_repl_stage_is_redacted_and_passes_cache_gate():
warm = {
"prefix_tokens": 10,
"e2e_s": 2,
"delta": {
"remote_jobs": 1,
"remote_hits": 1,
"tokens_reused": 10,
},
}
actual = {
"prefix_tokens": 10,
"output_tokens": 2,
"append_s": 0.1,
"ttft_s": 0.2,
"decode_s": 0.3,
"e2e_s": 0.4,
"delta": {
"local_hits": 1,
"remote_jobs": 0,
"tokens_computed": 0,
"fallbacks": 0,
},
}
stage = _stage("generator", warm, actual, "private output")
assert stage["ok"]
assert stage["output_chars"] == 14
assert len(stage["output_hash"]) == 64
assert "output" not in stage
Loading