diff --git a/docs/ops/distributed-prefill-kv-network.md b/docs/ops/distributed-prefill-kv-network.md index e4363858..e601f4df 100644 --- a/docs/ops/distributed-prefill-kv-network.md +++ b/docs/ops/distributed-prefill-kv-network.md @@ -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. diff --git a/scripts/agent_gan_inference_demo.py b/scripts/agent_gan_inference_demo.py index 20d9172c..ed0f247a 100644 --- a/scripts/agent_gan_inference_demo.py +++ b/scripts/agent_gan_inference_demo.py @@ -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: @@ -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() diff --git a/scripts/agent_gan_repl.py b/scripts/agent_gan_repl.py new file mode 100644 index 00000000..b1307582 --- /dev/null +++ b/scripts/agent_gan_repl.py @@ -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()) diff --git a/scripts/run_agent_gan_repl.sh b/scripts/run_agent_gan_repl.sh new file mode 100644 index 00000000..49c09caa --- /dev/null +++ b/scripts/run_agent_gan_repl.sh @@ -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" \ + "$@" diff --git a/tests/inference_engine/bridge/test_agent_gan_repl.py b/tests/inference_engine/bridge/test_agent_gan_repl.py new file mode 100644 index 00000000..19d80fb5 --- /dev/null +++ b/tests/inference_engine/bridge/test_agent_gan_repl.py @@ -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