k3-a40-bootstrap / banc_flashnext.py
patdev's picture
banc Qwen3.8-Flash-Next avec dechargement PLE
e94c0a0 verified
Raw History Blame Contribute Delete
8.51 kB
"""Qwen3.8-Flash-Next sur une seule carte, avec la table n-gramme en RAM.
CE QU'ON CHERCHE A SAVOIR. Ce modele bat Claude Opus 4.6 sur SWE-bench Pro
(62,5 contre 53,4) et Multilingual (81,0 contre 77,5), mesures avec le harnais
de Claude Code. Il est donc le bon modele pour cet usage. La seule question est
s'il tourne sur du materiel accessible.
125 B de parametres dont 6 B actifs, PLUS une table d'embeddings n-grammes de
51 B et 4 B de MTP -- 180 B au total. En NVFP4 : 125,9 Go, dont 47,7 pour la
seule table n-gramme.
`VLLM_PLE_CPU_OFFLOAD=1` place cette table en RAM hote. Il reste alors 78,2 Go
sur GPU, ce qui tient sur une RTX PRO 6000 de 96 Go SANS parallelisme de
tenseurs. C'est le point : toute la campagne a montre que les penalites
multi-cartes sont severes, et une carte unique les supprime toutes.
CE QUI PEUT ECHOUER, dans l'ordre de vraisemblance :
1. la RAM hote est insuffisante -- il faut 51 Go plus une marge ;
2. le dechargement fait traverser le PCIe a chaque consultation de n-gramme,
et le debit s'effondre ;
3. l'architecture `Qwen4Exp` n'a jamais ete servie sur sm_120 -- la recette
officielle vise GB300 et H200 ;
4. les parseurs different de Nemotron : `qwen3` et non `nemotron_v3`.
Reperes sur cette meme carte, meme banc, contexte 131 072 :
Nemotron NVFP4 (3,70 B actifs) : 275,3 solo | 910 agrege@8
Socle actif calcule pour Flash-Next : 4,02 Go/jeton -> plafond ~446 tok/s ici.
"""
import json
import os
import re
import statistics
import subprocess
import threading
import time
import urllib.request
MODEL = os.environ.get("BANC_MODEL", "RadixArk/Qwen3.8-Flash-Next-NVFP4")
PORT = 8000
URL = "http://127.0.0.1:%d" % PORT
CTX = os.environ.get("BANC_CTX", "131072")
SEQS = os.environ.get("BANC_SEQS", "32")
def dire(*a):
print(*a, flush=True)
dire("=" * 76)
subprocess.run(["nvidia-smi", "--query-gpu=name,memory.total,compute_cap",
"--format=csv,noheader"], check=False)
subprocess.run(["bash", "-lc",
"echo -n 'RAM hote : '; free -g | awk '/^Mem:/{print $2\" Go total, \"$7\" Go dispo\"}'"],
check=False)
subprocess.run(["python3", "-c",
"import vllm;print('vllm',vllm.__version__)"], check=False)
dire("PLE_CPU_OFFLOAD =", os.environ.get("VLLM_PLE_CPU_OFFLOAD", "(non pose)"))
dire("=" * 76)
BASE = ["vllm", "serve", MODEL,
"--served-model-name", "m",
"--host", "127.0.0.1", "--port", str(PORT),
"--trust-remote-code",
"--max-model-len", CTX,
"--max-num-seqs", SEQS,
"--enable-prefix-caching",
"--gpu-memory-utilization", "0.90",
# Parseurs de CE modele, differents de Nemotron.
"--reasoning-parser", "qwen3",
"--tool-call-parser", "qwen3_coder",
"--enable-auto-tool-choice",
# La recette officielle les recommande explicitement.
"--no-enable-flashinfer-autotune",
# Pas de vision : on sert du texte, et la tour visuelle coute de la
# memoire pour rien dans cet usage.
"--limit-mm-per-prompt", '{"image":0,"video":0}']
SUJETS = ["un cache LRU avec dict et liste doublement chainee",
"un pool de connexions avec expiration et sante des sockets",
"un analyseur d'expressions arithmetiques par descente recursive",
"une file de priorite par tas binaire avec decrease-key",
"un limiteur de debit par seau a jetons",
"un index inverse pour recherche plein texte",
"un ordonnanceur de taches avec dependances et cycles detectes",
"un serialiseur binaire avec versionnement de schema"]
journal = "/tmp/vllm.log"
with open(journal, "w") as f:
proc = subprocess.Popen(BASE, stdout=f, stderr=subprocess.STDOUT)
JALONS = ("Starting to load model", "Loading weights took", "Using ", "PLE",
"offload", "GPU KV cache size", "Capturing CUDA graphs",
"init engine", "Application startup complete", "Error", "Traceback")
pret = None
vus = set()
# Le telechargement fait 125,9 Go : compter large, et publier l'avancement
# sinon on reste aveugle une demi-heure.
for i in range(400):
try:
urllib.request.urlopen(URL + "/v1/models", timeout=5).read()
pret = i * 10
break
except Exception:
pass
if proc.poll() is not None:
break
if i and i % 6 == 0:
t = open(journal, errors="replace").read().splitlines()
recent = [l for l in t[-80:] if any(j in l for j in JALONS)]
msg = recent[-1].split("] ")[-1][:120] if recent else "(rien de neuf, %d lignes)" % len(t)
if msg not in vus:
vus.add(msg)
dire(" [%4d s] %s" % (i * 10, msg))
time.sleep(10)
texte = open(journal, errors="replace").read()
for motif in ("attention backend", "MoE backend", "GEMM", "GPU KV cache size",
"Maximum concurrency", "PLE", "offload"):
for ligne in texte.splitlines():
if motif in ligne and "ERROR" not in ligne:
dire(" " + ligne.split("] ")[-1][:150])
break
if pret is None:
dire("NE DEMARRE PAS")
vu = []
for ligne in texte.splitlines():
m = re.search(r"((?:Value|Runtime|NotImplemented|Assertion|Type|Memory|OS)Error"
r"|out of memory|Unsupported|is not supported|CUDA error)[:\s].{0,170}", ligne)
if m and m.group(0) not in vu:
vu.append(m.group(0))
for t in vu[:6]:
dire(" > " + t)
if not vu:
for ligne in texte.splitlines()[-15:]:
dire(" | " + ligne[:150])
raise SystemExit(1)
dire(" PRET en %d s (contexte %s, %s sequences)" % (pret, CTX, SEQS))
subprocess.run(["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader"], check=False)
def appel(contenu, sortie=400):
corps = json.dumps({"model": "m",
"messages": [{"role": "user", "content": contenu}],
"max_tokens": sortie, "temperature": 1.0, "top_p": 0.95,
"stream": True,
"stream_options": {"include_usage": True}}).encode()
r = urllib.request.Request(URL + "/v1/chat/completions", data=corps,
headers={"Content-Type": "application/json"})
t0 = time.time()
t1 = None
n = 0
bouts = []
with urllib.request.urlopen(r, timeout=900) as rep:
for l in rep:
l = l.strip()
if not l.startswith(b"data: ") or l[6:] == b"[DONE]":
continue
d = json.loads(l[6:])
ch = d.get("choices") or []
if not ch:
continue
de = ch[0].get("delta", {}) or {}
x = de.get("content") or de.get("reasoning_content") or de.get("reasoning")
if x:
if t1 is None:
t1 = time.time()
n += 1
bouts.append(x)
return {"ttft": (t1 - t0) if t1 else None, "n": n, "t1": t1,
"t2": time.time(), "txt": "".join(bouts)}
def div4(t):
m = t.split()
if len(m) < 40:
return 1.0
g = [tuple(m[i:i + 4]) for i in range(len(m) - 3)]
return len(set(g)) / len(g)
dire("conc | agrege | par flux | ttft med | 4-gr | etat")
for conc in (1, 4, 8):
res = [None] * conc
def un(i):
try:
res[i] = appel("Ecris en Python %s, avec trois tests unittest."
% SUJETS[i % len(SUJETS)])
except Exception as e:
res[i] = {"err": "%s: %s" % (type(e).__name__, str(e)[:70])}
d0 = time.time()
fils = [threading.Thread(target=un, args=(i,)) for i in range(conc)]
for f in fils:
f.start()
for f in fils:
f.join()
d1 = time.time()
bons = [r for r in res if r and not r.get("err") and r.get("t1")]
if not bons:
dire("%4d | ECHEC %s" % (conc, (res[0] or {}).get("err")))
continue
par = statistics.median([(r["n"] - 1) / (r["t2"] - r["t1"])
for r in bons if r["t2"] > r["t1"]])
ttft = statistics.median([r["t1"] - d0 for r in bons])
dv = statistics.median([div4(r["txt"]) for r in bons])
dire("%4d | %8.1f | %8.1f | %7.2fs | %.3f | %s"
% (conc, sum(r["n"] for r in bons) / (d1 - d0), par, ttft, dv,
"SAIN" if dv > 0.6 else "DEGENERE"))
dire("")
dire("repere meme carte, meme banc : Nemotron NVFP4 275,3 solo | 910 agrege@8")
dire("plafond memoire theorique de Flash-Next ici : ~446 tok/s (socle 4,02 Go)")
proc.terminate()