Qwen3.8-27B-MTP-GGUF / bench /test_ctx_limits.py
Davidmg0815's picture
Upload bench/test_ctx_limits.py with huggingface_hub
59c7b70 verified
Raw
History Blame
6.02 kB
#!/usr/bin/env python3
"""Maximaler stabiler Kontext je KV-Cache-Quantisierung — und ob er auch
schnell bleibt.
Der Reddit-Beitrag von u/Opening-Broccoli9190 misst dasselbe auf einer
RTX 5090 (32 GB). Hier laeuft es auf 3x RTX 3080 20 GB (60 GB gesamt,
Layer-Split), was die interessantere Frage beantwortet: skaliert der
KV-Cache ueber mehrere Karten so, wie die Rechnung es verspricht?
Zwei Kriterien, nicht eins:
1. Der Server startet ueberhaupt (passt in den Speicher)
2. Der Durchsatz bricht nicht ein — genau die Falle aus dem Beitrag,
wo 180k Kontext zwar liefen, aber mit 25 statt 100 Token/s
Binaere Suche, weil jeder Serverstart eine knappe Minute kostet.
"""
import json
import os
import signal
import subprocess
import sys
import time
import urllib.request
BIN = "/mnt/models/bin/llama-server"
MODELL = "/mnt/models/PropellerA-models/qwen3.8-27b-Q8_0.gguf"
PORT = 8291
LOGDIR = "/mnt/models/qwen38-prep/ctx-test"
os.makedirs(LOGDIR, exist_ok=True)
# Referenzdurchsatz bei kleinem Kontext, gegen den der Einbruch gemessen wird.
BASIS_CTX = 16384
def stoppe():
subprocess.run(["pkill", "-x", "llama-server"], capture_output=True)
for _ in range(25):
p = subprocess.run(["pgrep", "-x", "llama-server"], capture_output=True)
if p.returncode != 0:
break
time.sleep(1)
subprocess.run(["pkill", "-9", "-x", "llama-server"], capture_output=True)
time.sleep(3)
def starte(ctx, kv, logname):
log = open(os.path.join(LOGDIR, logname), "w")
p = subprocess.Popen(
[BIN, "-m", MODELL, "--host", "127.0.0.1", "--port", str(PORT),
"--ctx-size", str(ctx), "--parallel", "1", "--n-gpu-layers", "99",
"--tensor-split", "1,1,1", "--cache-type-k", kv, "--cache-type-v", kv,
"--flash-attn", "on", "--jinja"],
stdout=log, stderr=subprocess.STDOUT, preexec_fn=os.setsid)
for _ in range(240):
if p.poll() is not None:
return None # abgestuerzt oder sauber abgelehnt
try:
with urllib.request.urlopen(
f"http://127.0.0.1:{PORT}/health", timeout=3) as r:
if b"ok" in r.read():
return p
except Exception:
pass
time.sleep(2)
return None
def durchsatz(tokens=60):
"""Kurzer Erzeugungslauf. Liefert Token/s oder 0."""
body = json.dumps({
"messages": [{"role": "user",
"content": "Zaehle von 1 bis 40, nur die Zahlen."}],
"max_tokens": tokens, "cache_prompt": False, "temperature": 0.2,
"chat_template_kwargs": {"enable_thinking": False},
}).encode()
req = urllib.request.Request(
f"http://127.0.0.1:{PORT}/v1/chat/completions", data=body,
headers={"Content-Type": "application/json"})
try:
with urllib.request.urlopen(req, timeout=600) as r:
d = json.load(r)
return d.get("timings", {}).get("predicted_per_second") or 0.0
except Exception:
return 0.0
def vram():
p = subprocess.run(
["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader,nounits"],
capture_output=True, text=True)
return sum(int(z) for z in p.stdout.split() if z.isdigit())
def versuch(ctx, kv):
"""Startet mit diesem Kontext. Liefert (laeuft, t/s, belegtes VRAM)."""
stoppe()
name = f"{kv}-{ctx}.log"
p = starte(ctx, kv, name)
if p is None:
stoppe()
return False, 0.0, 0
v = vram()
t = durchsatz()
try:
os.killpg(os.getpgid(p.pid), signal.SIGTERM)
except Exception:
pass
stoppe()
return True, t, v
def main():
kvs = sys.argv[1:] or ["q8_0", "q5_1", "q4_0"]
print(f"Modell: {os.path.basename(MODELL)} | 3x RTX 3080 20 GB\n")
# Referenz bei kleinem Kontext, damit ein Einbruch erkennbar ist.
ok, basis, v0 = versuch(BASIS_CTX, "q8_0")
print(f"Referenz q8_0 @ {BASIS_CTX} Token: {basis:.1f} t/s, {v0} MiB\n")
ergebnisse = {}
for kv in kvs:
print(f"=== KV-Cache {kv} ===")
lo, hi = 16384, 400000 # hi ist bewusst zu hoch angesetzt
bestes = (0, 0.0, 0)
# Erst grob verdoppeln, bis es kippt — spart Runden gegenueber
# einer Suche ueber die volle Spanne.
probe = 64000
while probe <= hi:
ok, t, v = versuch(probe, kv)
print(f" {probe:7d} Token: {'laeuft' if ok else 'passt nicht':12s}"
f" {t:5.1f} t/s {v:6d} MiB")
if not ok:
hi = probe
break
bestes = (probe, t, v)
lo = probe
probe *= 2
else:
hi = probe
# Dann halbieren, bis die Luecke klein genug ist.
while hi - lo > 8000:
mitte = (lo + hi) // 2 // 1024 * 1024
ok, t, v = versuch(mitte, kv)
print(f" {mitte:7d} Token: {'laeuft' if ok else 'passt nicht':12s}"
f" {t:5.1f} t/s {v:6d} MiB")
if ok:
lo, bestes = mitte, (mitte, t, v)
else:
hi = mitte
ergebnisse[kv] = bestes
einbruch = (1 - bestes[1] / basis) * 100 if basis else 0
print(f" -> Maximum {bestes[0]} Token, {bestes[1]:.1f} t/s "
f"({einbruch:+.0f}% gegen Referenz), {bestes[2]} MiB\n")
print("=== Ergebnis ===")
print(f"{'KV-Cache':10s} {'max. Kontext':>14s} {'Token/s':>9s} "
f"{'gg. Referenz':>13s} {'VRAM':>9s}")
for kv, (ctx, t, v) in ergebnisse.items():
d = (t / basis - 1) * 100 if basis else 0
print(f"{kv:10s} {ctx:14d} {t:9.1f} {d:+12.0f}% {v:8d} MiB")
with open("/mnt/models/qwen38-prep/ctx-test/ergebnis.json", "w") as f:
json.dump({"basis_tps": basis, "basis_ctx": BASIS_CTX,
"ergebnisse": {k: dict(ctx=c, tps=t, vram=v)
for k, (c, t, v) in ergebnisse.items()}},
f, indent=1)
if __name__ == "__main__":
main()