Mooncake/mooncake-integration/ollama/benchmarks/smoke_e2e.py

129 lines
6.3 KiB
Python

#!/usr/bin/env python3
"""End-to-end smoke test: cross-process / cross-GPU KV reuse via Mooncake.
Agent A (on llama server #0, GPU X): prepare -> miss; /completion fully
prefills a long prompt; commit -> store the KV to Mooncake.
Agent B (on llama server #1, GPU Y): prepare -> the sidecar fetches the KV from
Mooncake and restores it into B's slot; /completion now prefills ~0 tokens.
We print the prompt-token count actually re-computed (timings.prompt_n) and the
prefill wall time for A vs B. B should be dramatically faster, proving the KV
crossed processes/GPUs through the store.
"""
import argparse, json, time, sys
import requests
def tokenize(llama, text, add_special=True):
r = requests.post(f"{llama}/tokenize", json={"content": text, "add_special": add_special})
r.raise_for_status()
return r.json()["tokens"]
def completion(llama, tokens, slot, n_predict=8):
r = requests.post(f"{llama}/completion", json={
"prompt": tokens, "id_slot": slot, "cache_prompt": True,
"n_predict": n_predict, "temperature": 0.0,
})
r.raise_for_status()
return r.json()
def bridge_call(bridge, ep, fp, policy, tokens, target):
r = requests.post(f"{bridge}/v1/{ep}", json={
"fp": fp, "policy": policy, "tokens": tokens, "target": target,
})
r.raise_for_status()
return r.json()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--bridge", default="http://127.0.0.1:52052")
ap.add_argument("--llama-a", default="http://127.0.0.1:52070")
ap.add_argument("--llama-b", default="http://127.0.0.1:52071")
ap.add_argument("--model-path", required=True)
ap.add_argument("--ctx-tokens", type=int, default=8000)
ap.add_argument("--namespace", default="smoke")
ap.add_argument("--block-size", type=int, default=256)
args = ap.parse_args()
fp = {"model_path": args.model_path, "kv_type": "f16", "block_size": args.block_size}
policy = {"enable": True, "namespace": args.namespace, "read": True, "write": True,
"block_size": args.block_size, "replica_num": 1, "min_prefix_blocks": 1}
# Build a long, code-like context prompt of ~ctx-tokens tokens.
unit = ("// module {i}: utility helpers for the data pipeline\n"
"func process_{i}(records []Record) (Result, error) {{\n"
" // validate, transform, and aggregate the {i}-th shard\n"
" return aggregate(transform(validate(records))), nil\n}}\n\n")
text = "You are reviewing a large Go repository. Here is the source:\n\n"
i = 0
while len(tokenize(args.llama_a, text)) < args.ctx_tokens:
text += unit.format(i=i)
i += 1
text += "\nSummarize the overall architecture of this repository."
tokens = tokenize(args.llama_a, text)
print(f"[setup] prompt = {len(tokens)} tokens, block_size={args.block_size} "
f"=> {len(tokens)//args.block_size} full blocks\n")
# ---- Agent A: cold on server #0 ----
ta = time.perf_counter()
pa = bridge_call(args.bridge, "prepare", fp, policy, tokens,
{"base_url": args.llama_a, "slot_id": 0})
pa_wall = (time.perf_counter() - ta) * 1e3
t0 = time.perf_counter()
ca = completion(args.llama_a, tokens, slot=0)
a_wall = (time.perf_counter() - t0) * 1e3
co = bridge_call(args.bridge, "commit", fp, policy, tokens,
{"base_url": args.llama_a, "slot_id": 0})
a_ttft = pa_wall + ca['timings']['prompt_ms']
print("AGENT A (cold, server #0):")
print(f" prepare : hit={pa['hit']} decision={pa['decision']} wall={pa_wall:.1f}ms ({pa.get('reason','')})")
print(f" prefill : prompt_n={ca['timings']['prompt_n']} toks, "
f"prompt_ms={ca['timings']['prompt_ms']:.1f}, e2e_wall={a_wall:.1f}ms")
print(f" commit : stored={co['stored']} blocks={co['stored_blocks']} "
f"bytes={co['bytes']} put_ms={co['store_put_ms']:.1f} key=...{co['key'][-24:]}")
print(f" >> honest TTFT(A) = prepare {pa_wall:.1f} + prefill {ca['timings']['prompt_ms']:.1f} = {a_ttft:.1f}ms\n")
# ---- Agent B: warm on server #1 (must restore from the store) ----
tb = time.perf_counter()
pb = bridge_call(args.bridge, "prepare", fp, policy, tokens,
{"base_url": args.llama_b, "slot_id": 0})
pb_wall = (time.perf_counter() - tb) * 1e3
t0 = time.perf_counter()
cb = completion(args.llama_b, tokens, slot=0)
b_wall = (time.perf_counter() - t0) * 1e3
b_ttft = pb_wall + cb['timings']['prompt_ms']
print("AGENT B (warm, server #1 -- different process & GPU):")
print(f" prepare : hit={pb['hit']} decision={pb['decision']} restored={pb['restored']} "
f"restored_tokens={pb['restored_tokens']} wall={pb_wall:.1f}ms")
print(f" store_get_ms={pb['store_get_ms']:.1f} bytes={pb['bytes']} ({pb.get('reason','')})")
print(f" prefill : prompt_n={cb['timings']['prompt_n']} toks, "
f"prompt_ms={cb['timings']['prompt_ms']:.1f}, e2e_wall={b_wall:.1f}ms")
print(f" >> honest TTFT(B) = prepare {pb_wall:.1f} + prefill {cb['timings']['prompt_ms']:.1f} = {b_ttft:.1f}ms\n")
# ---- verdict ----
saved = ca['timings']['prompt_n'] - cb['timings']['prompt_n']
ttft_red = 100.0 * (1 - b_ttft / max(a_ttft, 1e-9))
print("=" * 70)
print(f" prefill tokens recomputed: A={ca['timings']['prompt_n']} -> B={cb['timings']['prompt_n']} "
f"(saved {saved} tokens, {100.0*saved/max(ca['timings']['prompt_n'],1):.1f}%)")
print(f" HONEST end-to-end TTFT: A={a_ttft:.1f}ms -> B={b_ttft:.1f}ms "
f"(reduced {ttft_red:.1f}%)")
print(f" (B breakdown: store_get {pb['store_get_ms']:.1f}ms + restore/overhead "
f"{pb_wall - pb['store_get_ms']:.1f}ms + prefill {cb['timings']['prompt_ms']:.1f}ms)")
print("=" * 70)
stats = requests.get(f"{args.bridge}/stats").json()
print(f"sidecar stats: hits={stats['hits']} misses={stats['misses']} "
f"saved_prefill_tokens={stats['saved_prefill_tokens']} "
f"learned_get_gbps={stats['learned_get_gbps']:.1f} learned_prefill_tps={stats['learned_prefill_tps']:.0f}")
ok = cb['timings']['prompt_n'] < ca['timings']['prompt_n'] * 0.3 and pb['restored']
print("\nRESULT:", "PASS — KV reused across processes/GPUs via Mooncake" if ok else "FAIL")
sys.exit(0 if ok else 1)
if __name__ == "__main__":
main()