GLM-5.3-Flash on two desk-side DGX Sparks — 989,727 tokens of KV, every script inline
What you'll build: a 320B-parameter multimodal MoE served from two desk-side boxes, with speculative decoding tuned to the point where the hardware, not the software, is the limit. Every script, patch and flag is on this page, in the order you run them, with the exact command for each machine. Copy nothing but what's here and you end up with a serving stack.
Everything below is also available as a single bundle: download the bundle — SHA-256 2d5931db561d58c22bb8dca60a1a45e41905e1e4d11d7cafd44cd950ef41eb94 (launcher, Dockerfile, overlay, strip script, ordered README — laid out as the working directory, so the build command works as printed).
Before you start
Hardware. Two NVIDIA DGX Spark machines (GB10, 128 GiB unified memory each), headless Linux, joined by a direct CX7 cable with the link UP on both sides — the launcher fails fast if not.
Software. Docker installed and working on both nodes. Passwordless SSH from the head node to the worker node. python3 on the head node with torch (or numpy + ml_dtypes) for the strip script.
Disk. ~400 GiB free on the head — it holds the original checkpoint (~180 GiB) plus the stripped copy (~176 GiB) plus the drafter, transiently. ≥200 GiB on the worker (stripped copy + drafter only).
Names. The first box is head (serves the API, TP rank 0); the second is worker (TP rank 1). Every step below tells you which machine you're on. Total time: roughly an hour of downloads and ~10 min of build — once.
Step 1 — download the weights (head node, ~1 h)
Download the NVFP4 checkpoint, then clone its HF cache directory — the clone becomes the stripped working copy; the original stays pristine:
hf download LibertAIDAI/GLM-5.3-Flash-NVFP4
cp -a ~/.cache/huggingface/hub/models--LibertAIDAI--GLM-5.3-Flash-NVFP4 \
~/.cache/huggingface/hub/models--LibertAIDAI--GLM-5.3-Flash-NVFP4-no-MTP
Step 2 — strip the MTP layer (head node, ~15 min)
DFlash2 replaces the model's native MTP layer as the speculation method, so we remove it (~4.1 GiB back). Create strip_mtp.py on the head node — it edits config.json + the safetensors index and, with --strip-shards, physically removes the MTP tensors, verified per shard, dtype-preserving.
Install its dependencies first:
pip install safetensors torch
(or pip install safetensors numpy ml_dtypes for the lighter path — the script needs safetensors even for a dry-run):
#!/usr/bin/env python3
"""
strip_glm53_mtp.py -- drop the MTP stick from a GLM-5.3 Flash (or similar
GLM/DeepSeek "Next") checkpoint so vLLM stops loading those weights.
Always run against a COPY. Two modes:
default / --index-only : edits config.json + model.safetensors.index.json only.
vLLM (and HF) load tensors through the index's weight_map, so MTP tensors
vanish from its view without touching a single weight byte. Safe,
seconds, but does not reclaim disk.
--strip-shards : physically deletes the MTP tensors from the .safetensors
shards so the folder shrinks. Each shard is rewritten to a temp file,
fsynced, then atomically replaced; after every rewrite the new file is
verified (kept tensors present with identical dtype+shape, dropped gone).
Weight bytes are preserved exactly: U8 / F8_E4M3 / F32 / BF16 etc. round-trip
as-is. Backend is auto-selected -- PyTorch if importable (native fp8/fp4), else
numpy + ml_dtypes with explicit aliasing. If a dtype the numpy path cannot
round-trip is ever found, the run aborts loudly instead of corrupting.
Usage:
python3 strip_glm53_mtp.py /path/to/COPY # index-only (recommended)
python3 strip_glm53_mtp.py /path/to/COPY --dry-run
python3 strip_glm53_mtp.py /path/to/COPY --strip-shards
"""
import argparse
import json
import os
import re
import shutil
import struct
import sys
from pathlib import Path
# ---------------------------------------------------------------------------
# Backend selection
# ---------------------------------------------------------------------------
try:
import torch # noqa: F401
BACKEND = "pt"
from safetensors import safe_open, save_file
except ImportError:
BACKEND = "np"
import numpy as np
import ml_dtypes # noqa: F401
from safetensors import safe_open
from safetensors.numpy import save_file
# safetensors' numpy backend resolves dtypes as ATTRIBUTES on the numpy
# module; importing ml_dtypes does not add them -- alias explicitly.
np.bfloat16 = ml_dtypes.bfloat16
np.float8_e4m3fn = ml_dtypes.float8_e4m3fn
if hasattr(ml_dtypes, "float8_e5m2"):
np.float8_e5m2 = ml_dtypes.float8_e5m2
if hasattr(ml_dtypes, "float4_e2m1fn"):
np.float4_e2m1fn = ml_dtypes.float4_e2m1fn
DTYPE_BYTES = {
"F64": 8, "F32": 4, "F16": 2, "BF16": 2,
"F8_E4M3": 1, "F8_E5M2": 1, "F4_E2M1": 1,
"I64": 8, "I32": 4, "I16": 2, "I8": 1,
"U64": 8, "U32": 4, "U16": 2, "U8": 1, "BOOL": 1,
}
# Only real general-purpose dtypes; numpy backend round-trips these exactly.
SAFE_NP = set(DTYPE_BYTES) - {"F4_E2M1"}
def dtype_basename(d) -> str:
"""'<Dtype.F8_E4M3: 19>' / 'Dtype.F8_E4M3' -> 'F8_E4M3'."""
s = str(d).split(":")[0].split(".")[-1].strip()
return s.strip("<> ")
def numel(shape) -> int:
n = 1
for x in shape:
n *= int(x)
return n
def fmt_bytes(n: int) -> str:
for unit in ("B", "KiB", "MiB", "GiB", "TiB"):
if abs(n) < 1024 or unit == "TiB":
return f"{n:.2f} {unit}" if unit != "B" else f"{n} B"
n /= 1024
return f"{n:.2f} TiB"
# ---------------------------------------------------------------------------
# MTP key detection
# ---------------------------------------------------------------------------
def is_mtp_key(key: str, num_hidden_layers: int) -> bool:
# layer index >= num_hidden_layers (main stack is 0..N-1, MTP sits at N)
m = re.search(r"\.layers\.(\d+)\.", key)
if m and int(m.group(1)) >= num_hidden_layers:
return True
# explicit markers used by some checkpoints
if re.search(r"(^|\.)(mtp|mtp_layers|nextn|shared_head)(\.|$)", key):
return True
return False
# ---------------------------------------------------------------------------
# JSON + file helpers
# ---------------------------------------------------------------------------
def find_model_dir(base: Path):
"""Locate the directory holding config.json; auto-resolve an HF hub cache root.
Returns (model_dir, from_cache). A cache root is models--NAME/{refs,snapshots,blobs};
the files themselves live under snapshots/<revision> (symlinks to blobs/).
"""
if (base / "config.json").exists() and (base / "model.safetensors.index.json").exists():
return base, False
refs = base / "refs" / "main"
if refs.exists():
rev = refs.read_text().strip() or refs.resolve().name
cand = base / "snapshots" / rev
if cand.is_dir() and (cand / "config.json").exists():
print(f" HF cache root detected -> using snapshot {rev}")
return cand, True
snaps = base / "snapshots"
if snaps.is_dir():
cands = [p for p in sorted(snaps.iterdir())
if p.is_dir() and (p / "config.json").exists()]
if len(cands) == 1:
print(f" HF cache root detected -> using snapshot {cands[0].name}")
return cands[0], True
if len(cands) > 1:
raise SystemExit(
"multiple snapshots found: " + ", ".join(p.name for p in cands) +
"\npass the exact snapshot dir instead of the cache root")
raise SystemExit(
f"no config.json / model.safetensors.index.json under {base}\n"
"pass either the flat model dir or an HF cache root (models--NAME)")
def load_json(path: Path):
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
def save_json_atomic(path: Path, data) -> None:
tmp = path.with_suffix(path.suffix + ".tmp")
with open(tmp, "w", encoding="utf-8") as f:
json.dump(data, f, indent=2, ensure_ascii=False)
f.write("\n")
f.flush()
os.fsync(f.fileno())
os.replace(tmp, path)
def backup(path: Path) -> None:
shutil.copy2(path, path.with_suffix(path.suffix + ".bak"))
def peek(path: Path) -> dict:
"""header-only read: key -> (dtype_basename, shape)."""
with safe_open(str(path), framework=BACKEND) as f:
return {
k: (dtype_basename(f.get_slice(k).get_dtype()), tuple(f.get_slice(k).get_shape()))
for k in f.keys()
}
def write_empty_safetensors(path: Path) -> None:
h = json.dumps({"__metadata__": {"format": "pt"}}).encode()
h += b" " * ((8 - len(h) % 8) % 8)
with open(path, "wb") as f:
f.write(struct.pack("<Q", len(h)))
f.write(h)
f.flush()
os.fsync(f.fileno())
# ---------------------------------------------------------------------------
# Shard rewriting
# ---------------------------------------------------------------------------
def strip_shard(path: Path, drop: set, backend: str) -> int:
real = path.resolve()
orig = peek(real)
unknown = [k for k in drop if k not in orig]
if unknown:
print(f" ! keys listed in index but missing here (will still drop from index): {unknown}")
kept = {k: v for k, v in orig.items() if k not in drop}
freed = 0
for k, (dn, sh) in orig.items():
if k in drop:
freed += numel(sh) * DTYPE_BYTES.get(dn, 0)
# dtype safety gate: refuse rather than risk corrupting weights
if backend == "np":
bad = {dn for _, (dn, _) in orig.items() if dn not in SAFE_NP}
if bad:
raise SystemExit(
f" shard {real.name} contains dtypes {bad} that the numpy backend "
"cannot round-trip safely. Install torch (`pip install torch`) and rerun."
)
if not kept:
write_empty_safetensors(real)
with safe_open(str(real), framework=BACKEND) as f:
assert list(f.keys()) == [], "empty shard did not stay empty"
print(f" → shard now empty (all tensors were MTP): {real.name} ({fmt_bytes(freed)} freed)")
return freed
arrays = {}
with safe_open(str(real), framework=BACKEND) as f:
for k in f.keys():
if k not in drop:
arrays[k] = f.get_tensor(k)
tmp = real.with_name(real.name + ".tmp")
mode = os.stat(real).st_mode & 0o7777
try:
save_file(arrays, str(tmp))
os.chmod(tmp, mode)
os.replace(tmp, real)
finally:
tmp.unlink(missing_ok=True)
# ---- verify -----------------------------------------------------------
new = peek(real)
assert set(new) == set(kept), (
f"verify FAILED: {len(new)} tensors written, expected {len(kept)}")
mism = [k for k in kept if new[k] != kept[k]]
if mism:
raise SystemExit(f"verify FAILED on {real.name}: dtype/shape changed for {mism[:5]}")
print(f" → {real.name}: {len(orig)} → {len(kept)} tensors, {fmt_bytes(freed)} freed, verified")
return freed
# ---------------------------------------------------------------------------
# main
# ---------------------------------------------------------------------------
def main():
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("model_dir", help="path to a COPY of the model")
ap.add_argument("--index-only", action="store_true",
help="only fix config.json + index.json (default; no shard rewrite)")
ap.add_argument("--strip-shards", action="store_true",
help="physically delete MTP tensors from the safetensors shards")
ap.add_argument("--dry-run", action="store_true",
help="show the plan, change nothing")
ap.add_argument("--force", action="store_true",
help="proceed even if the index already has no MTP keys")
args = ap.parse_args()
model_dir = Path(args.model_dir).resolve()
if not model_dir.is_dir():
raise SystemExit(f"Directory does not exist: {model_dir}")
model_dir, from_cache = find_model_dir(model_dir)
cfg_path = model_dir / "config.json"
idx_path = model_dir / "model.safetensors.index.json"
for p in (cfg_path, idx_path):
if not p.exists():
raise SystemExit(f"missing {p.name} under {model_dir}")
strip = args.strip_shards
mode = "STRIP-SHARDS" if strip else "INDEX-ONLY"
# ------------------------------------------------------------------
# 1. config.json -> num_nextn_predict_layers = 0
# ------------------------------------------------------------------
print(f">>> [{mode}] Reading {cfg_path.name} …")
config = load_json(cfg_path)
def zero_mtp_flags(cfg: dict) -> bool:
changed = False
if cfg.get("num_nextn_predict_layers", 0) != 0:
cfg["num_nextn_predict_layers"] = 0
changed = True
if cfg.get("index_share_for_mtp_iteration", False) is not False:
cfg["index_share_for_mtp_iteration"] = False
changed = True
return changed
changed = zero_mtp_flags(config.get("text_config", config))
if "text_config" in config:
changed |= zero_mtp_flags(config) # catch any top-level duplicate keys
if changed:
print(" → num_nextn_predict_layers = 0, index_share_for_mtp_iteration = false")
if not args.dry_run:
backup(cfg_path)
save_json_atomic(cfg_path, config)
print(" config.json updated (backup → config.json.bak)")
else:
print(" config already has MTP disabled.")
num_hidden = config.get("text_config", config).get("num_hidden_layers", 45)
# ------------------------------------------------------------------
# 2. find MTP keys in the index
# ------------------------------------------------------------------
print(f"\n>>> Scanning {idx_path.name} for MTP tensors (layers >= {num_hidden}) …")
index = load_json(idx_path)
wm = index.get("weight_map", {})
mtp_keys = [k for k in wm if is_mtp_key(k, num_hidden)]
print(f" found {len(mtp_keys)} MTP tensors of {len(wm)} total (index)")
for k in sorted(mtp_keys)[:6]:
print(f" • {k}")
# ------------------------------------------------------------------
# 2b. physical tensors: this vLLM loader enumerates the shard FILES,
# so MTP tensors physically present MUST go, even if the index
# was already cleaned. Scan every shard header (cheap).
# ------------------------------------------------------------------
shard_to_keys: dict = {}
for k in mtp_keys:
shard_to_keys.setdefault(wm[k], []).append(k)
if strip:
print(f"\n>>> Scanning shard files for MTP tensors (loader enumerates files, not the index) …")
for shard in sorted(p.name for p in model_dir.glob("*.safetensors")):
p = model_dir / shard
try:
spec = peek(p)
except Exception:
if args.dry_run:
continue
raise
phys = [k for k in spec if is_mtp_key(k, num_hidden)]
if phys:
known = set(shard_to_keys.get(shard, []))
extra = [k for k in phys if k not in known]
if extra:
print(f" {shard}: {len(phys)} physical MTP tensors "
f"({len(known)} already indexed, {len(extra)} extra)")
shard_to_keys[shard] = known | set(phys)
phys_total = sum(len(v) for v in shard_to_keys.values())
print(f" total MTP tensors to remove physically: {phys_total} across "
f"{len(shard_to_keys)} shard(s)")
if not shard_to_keys:
print(" nothing to do. (use --force to re-run anyway)")
return
# compute removed bytes from shard headers (header-only, cheap)
freed = 0
for shard, keys in shard_to_keys.items():
p = model_dir / shard
if not p.exists():
print(f" ! shard absent, skipping size calc: {shard}")
continue
try:
spec = peek(p)
except Exception as e:
if args.dry_run:
continue
raise
for k in keys:
if k in spec:
freed += numel(spec[k][1]) * DTYPE_BYTES.get(spec[k][0], 0)
print(f" estimated savings: {fmt_bytes(freed)} in {len(shard_to_keys)} shard(s)")
if args.dry_run:
print("\n[DRY-RUN] stopping -- no files changed.")
return
# ------------------------------------------------------------------
# 3. optionally physically strip the shards
# ------------------------------------------------------------------
if strip:
print(f"\n>>> Stripping tensors from safetensors shards (backend: {BACKEND}) …")
if from_cache:
print(" NOTE: this looks like an HF cache tree. Rewriting goes through the")
print(" snapshot symlinks into blobs/ -- fine on a COPY, dangerous on the")
print(" ORIGINAL (mutates the shared blob files). Double-check you are")
print(" pointed at a copy.")
for shard, keys in shard_to_keys.items():
p = model_dir / shard
if not p.exists():
print(f" WARNING: shard not found, skipping → {shard}")
continue
if p.is_symlink():
tgt = p.resolve()
print(f" {shard} is a symlink → {tgt.parent.name}/{tgt.name}")
if tgt.parent != (model_dir / "blobs") and str(tgt).startswith(str(model_dir)):
pass
print(f" processing {shard} ({len(keys)} tensors) …")
strip_shard(p, set(keys), BACKEND)
# ------------------------------------------------------------------
# 4. rewrite the index + total_size
# ------------------------------------------------------------------
print(f"\n>>> Updating {idx_path.name} …")
new_wm = {k: v for k, v in wm.items() if k not in mtp_keys}
index["weight_map"] = new_wm
meta = index.setdefault("metadata", {})
old_total = meta.get("total_size")
if isinstance(old_total, int):
meta["total_size"] = max(0, old_total - freed)
print(f" total_size {old_total} → {meta['total_size']} (freed {fmt_bytes(freed)})")
else:
meta.pop("total_size", None)
backup(idx_path)
save_json_atomic(idx_path, index)
print(f" index updated: {len(wm)} → {len(new_wm)} entries (backup → {idx_path.name}.bak)")
print("\n✅ Done.")
if strip:
print(f" Removed {len(mtp_keys)} MTP tensors physically; folder shrank by {fmt_bytes(freed)}.")
else:
print(" MTP tensors are now unreferenced in the index; vLLM will not load them.")
print(" (folder size unchanged -- rerun with --strip-shards to reclaim disk)")
print(" Run on the COPY dir with vLLM; keep the original untouched.")
if __name__ == "__main__":
main()
Dry-run first — it prints exactly what it would do and touches nothing (~1 min, index-only):
hf download incoai/GLM-5.3-Flash-DFlash2
Happy with the plan? Run it for real — this rewrites ~180 GiB of shards, so expect 5–15 min on typical NVMe:
rsync -a --info=progress2 ~/.cache/huggingface/hub/models--LibertAIDAI--GLM-5.3-Flash-NVFP4-no-MTP/ user@worker-node:~/.cache/huggingface/hub/
rsync -a --info=progress2 ~/.cache/huggingface/hub/models--incoai--GLM-5.3-Flash-DFlash2/ user@worker-node:~/.cache/huggingface/hub/
Step 3 — pull the drafter (head node, ~2 min)
The DFlash2 drafter weights, ~2.2 GiB:
hf download incoai/GLM-5.3-Flash-DFlash2
Step 4 — sync both caches to the worker (head node, ~30 min)
The worker needs the same weights. Copy the whole models-- directories — snapshots contain relative symlinks into blobs/, so copy the directory, not loose files:
rsync -a --info=progress2 ~/.cache/huggingface/hub/models--LibertAIDAI--GLM-5.3-Flash-NVFP4-no-MTP/ user@worker-node:~/.cache/huggingface/hub/
rsync -a --info=progress2 ~/.cache/huggingface/hub/models--incoai--GLM-5.3-Flash-DFlash2/ user@worker-node:~/.cache/huggingface/hub/
Step 5 — build the serving image (head node, ~10 min)
Create the working directory (this is the build context) and its tree:
mkdir -p glm5.3-flash-spark/files/overlay-dflash2/dflash2 glm5.3-flash-spark/tony-lane
cd glm5.3-flash-spark
Create tony-lane/Dockerfile:
# DFLASH2 tony-lane image — built directly on the proven day-0 stack that
# tonyd2wild served 46.9 tok/s (C1 warm) / 0.741 acceptance with on 2x DGX
# Spark, fp8 KV, DFLASH2. Base = their published sm121-v8; overlay = the
# vendored overlay-dflash2/ (their exact patch set: registry+select,
# GLM aux capture, drafter KV group). No Ray, no Mia kernel layer.
#
# Build context = repo ROOT so `files/overlay-dflash2` is reachable:
# docker build --network=host -t glm53-dflash2-tony:v1 -f tony-lane/Dockerfile .
# (accepts arm64 on the Spark; do NOT add --platform.)
FROM radixark/vllm-glm53-flash:sm121-v8
# Optional InstantTensor direct-I/O loader (v9's trick, 15x faster loads).
# CONFIRMED WORKING on this kit (2026-08-29: boot + serve). tonyd2wild's md
# reports 4/4 silent TP2 deaths ~1 min after load on their stack at every KV
# budget — the same-layer nccl re-pin to 2.30.7 here may be why we survive
# (unproven). Off by default only because safetensors remains the proven
# default: build with `--build-arg INSTANTTENSOR=1` and launch with
# `LOAD_FORMAT=instanttensor`. Same-layer nccl re-pin is REQUIRED
# (instanttensor downgrades nccl->2.29.7, fabric-fatal).
ARG INSTANTTENSOR=0
RUN if [ "$INSTANTTENSOR" = "1" ]; then \
pip install -q instanttensor \
&& pip install -q nvidia-nccl-cu13==2.30.7 \
&& pip show nvidia-nccl-cu13 | grep -q "Version: 2.30.7" \
&& python3 -c "import instanttensor" \
&& echo "[glm53-tony] instanttensor ready (nccl re-pinned 2.30.7)"; \
fi
COPY files/overlay-dflash2/ /opt/dflash2/
RUN cp /opt/dflash2/qwen3_dflash2.py \
/usr/local/lib/python3.12/dist-packages/vllm/model_executor/models/ \
&& mkdir -p /usr/local/lib/python3.12/dist-packages/vllm/v1/worker/gpu/spec_decode/dflash2 \
&& cp /opt/dflash2/dflash2/__init__.py \
/opt/dflash2/dflash2/speculator.py \
/usr/local/lib/python3.12/dist-packages/vllm/v1/worker/gpu/spec_decode/dflash2/ \
&& python3 /opt/dflash2/patch_registry_and_select.py \
/usr/local/lib/python3.12/dist-packages/vllm \
&& python3 /opt/dflash2/patch_glm_aux_capture.py \
&& python3 /opt/dflash2/patch_glm5_drafter_group.py \
&& python3 -c "from vllm.model_executor.models.registry import ModelRegistry; \
assert 'DFlash2DraftModel' in ModelRegistry.get_supported_archs(); \
print('[glm53-tony] registry: DFlash2 OK')" \
&& python3 -c "import vllm.model_executor.models.qwen3_dflash2, \
vllm.v1.worker.gpu.spec_decode.dflash2.speculator; \
print('[glm53-tony] dflash2 modules import OK')"
The Dockerfile builds on tonyd2wild's day-0 image and applies his overlay — three of his four patch scripts; we omit patch_kv_page_lcm2.py, because on our config --block-size 2304 plus the engine's automatic mamba page-size padding covers what it does (visible in the boot log). The overlay files are his, republished verbatim (byte-identical to commit 9642d4f) so this build needs nothing else.
Create files/overlay-dflash2/qwen3_dflash2.py:
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import torch
import torch.nn.functional as F
from torch import nn
from vllm.compilation.backends import set_model_tag
from vllm.compilation.decorators import support_torch_compile
from vllm.config import CacheConfig, VllmConfig
from vllm.model_executor.layers.linear import ReplicatedLinear
from vllm.model_executor.layers.logits_processor import LogitsProcessor
from vllm.model_executor.layers.quantization.base_config import QuantizationConfig
from .qwen3_dflash import (
DFlashQwen3DecoderLayer,
DFlashQwen3ForCausalLM,
DFlashQwen3Model,
)
from .utils import maybe_prefix
def _grouped_conv(
hidden_states: torch.Tensor,
delta: torch.Tensor,
base: torch.Tensor,
block_size: int,
num_groups: int,
group_size: int,
taps: int,
) -> torch.Tensor:
blocks = hidden_states.unflatten(-1, (num_groups, group_size))
coefficients = base.view(1, taps, num_groups, group_size) + delta.unsqueeze(-1)
output = coefficients[:, 0] * blocks
position = torch.arange(hidden_states.shape[0], device=hidden_states.device)
if block_size & (block_size - 1) == 0:
position = position & (block_size - 1)
else:
position = position % block_size
for tap in range(1, taps):
shifted = F.pad(blocks[:-tap], (0, 0, 0, 0, tap, 0))
output += coefficients[:, tap] * shifted * (position >= tap).view(-1, 1, 1)
return output.flatten(-2)
class DFlashGroupedConv(nn.Module):
def __init__(
self,
hidden_size: int,
taps: int,
group_size: int,
block_size: int,
params_dtype: torch.dtype,
prefix: str,
) -> None:
super().__init__()
if hidden_size % group_size:
raise ValueError(
f"conv_group_size={group_size} must divide hidden_size={hidden_size}."
)
self.block_size = block_size
self.taps = taps
self.group_size = group_size
self.num_groups = hidden_size // group_size
self.base_kernel = nn.Parameter(
torch.empty(2, taps, hidden_size, dtype=params_dtype),
requires_grad=False,
)
self.kernel_projection = ReplicatedLinear(
hidden_size,
2 * taps * self.num_groups,
bias=False,
params_dtype=params_dtype,
quant_config=None,
prefix=maybe_prefix(prefix, "kernel_projection"),
return_bias=False,
)
def _convolve(
self, hidden_states: torch.Tensor, delta: torch.Tensor, side: int
) -> torch.Tensor:
return _grouped_conv(
hidden_states,
delta,
self.base_kernel[side],
self.block_size,
self.num_groups,
self.group_size,
self.taps,
)
def prepare(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
coefficients = self.kernel_projection(hidden_states).reshape(
hidden_states.shape[0], 2, self.taps, self.num_groups
)
return self._convolve(hidden_states, coefficients[:, 0], 0), coefficients[:, 1]
def finish(
self, hidden_states: torch.Tensor, coefficients: torch.Tensor
) -> torch.Tensor:
return self._convolve(hidden_states, coefficients, 1)
class DFlash2Qwen3DecoderLayer(DFlashQwen3DecoderLayer):
def __init__(
self,
vllm_config: VllmConfig,
*,
config,
layer_idx: int,
cache_config: CacheConfig | None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__(
vllm_config,
config=config,
layer_idx=layer_idx,
cache_config=cache_config,
quant_config=quant_config,
prefix=prefix,
)
draft_config = config.dflash_config
speculative_config = vllm_config.speculative_config
assert speculative_config is not None
conv_args = dict(
hidden_size=config.hidden_size,
taps=int(draft_config["conv_kernel_size"]),
group_size=int(draft_config["conv_group_size"]),
# Query tokens per request: the bonus token plus the mask tokens.
block_size=1 + speculative_config.num_speculative_tokens,
params_dtype=vllm_config.model_config.dtype,
)
self.attention_conv = DFlashGroupedConv(
**conv_args, prefix=maybe_prefix(prefix, "attention_conv")
)
self.mlp_conv = DFlashGroupedConv(
**conv_args, prefix=maybe_prefix(prefix, "mlp_conv")
)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
if residual is None:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
else:
hidden_states, residual = self.input_layernorm(hidden_states, residual)
hidden_states, coefficients = self.attention_conv.prepare(hidden_states)
hidden_states = self.self_attn(positions=positions, hidden_states=hidden_states)
hidden_states = self.attention_conv.finish(hidden_states, coefficients)
hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
hidden_states, coefficients = self.mlp_conv.prepare(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = self.mlp_conv.finish(hidden_states, coefficients)
return hidden_states, residual
def _score_edges(
predecessor_table: torch.Tensor,
successor_table: torch.Tensor,
candidate_ids: torch.Tensor,
unary_logits: torch.Tensor,
hidden: torch.Tensor,
anchor_token_ids: torch.Tensor,
top_k: int,
) -> torch.Tensor:
successors = successor_table[candidate_ids]
predecessor_ids = torch.cat(
(
anchor_token_ids[:, None, None].expand(-1, 1, top_k),
candidate_ids[:, :-1],
),
dim=1,
)
predecessors = predecessor_table[predecessor_ids]
return unary_logits[:, :, None] + torch.einsum(
"blpr,blcr->blpc", predecessors * hidden[:, :, None], successors
)
@support_torch_compile
class CandidateSelector(nn.Module):
def __init__(
self,
hidden_size: int,
vocab_size: int,
rank: int,
top_k: int,
params_dtype: torch.dtype,
prefix: str,
) -> None:
super().__init__()
self.top_k = top_k
self.predecessor_codebook = nn.Parameter(
torch.empty(vocab_size, rank, dtype=params_dtype), requires_grad=False
)
self.successor_codebook = nn.Parameter(
torch.empty(vocab_size, rank, dtype=params_dtype), requires_grad=False
)
self.hidden_projection = ReplicatedLinear(
hidden_size,
rank,
bias=False,
params_dtype=params_dtype,
quant_config=None,
prefix=maybe_prefix(prefix, "hidden_projection"),
return_bias=False,
)
def forward(
self,
candidate_ids: torch.Tensor,
unary_logits: torch.Tensor,
hidden_states: torch.Tensor,
anchor_token_ids: torch.Tensor,
) -> torch.Tensor:
hidden = self.hidden_projection(hidden_states)
return _score_edges(
self.predecessor_codebook,
self.successor_codebook,
candidate_ids,
unary_logits,
hidden,
anchor_token_ids,
self.top_k,
)
class DFlash2Qwen3Model(DFlashQwen3Model):
decoder_layer_cls = DFlash2Qwen3DecoderLayer
def __init__(
self,
*,
vllm_config: VllmConfig,
start_layer_id: int = 0,
prefix: str = "",
) -> None:
super().__init__(
vllm_config=vllm_config,
start_layer_id=start_layer_id,
prefix=prefix,
)
draft_config = self.config.dflash_config
self.input_embedding_scale = float(
draft_config.get("input_embedding_scale", 1.0)
)
# Without its own tag the selector shares the draft head's compile cache.
with set_model_tag("dflash2_candidate_selector"):
self.candidate_selector = CandidateSelector(
hidden_size=self.config.hidden_size,
vocab_size=self.config.vocab_size,
rank=int(draft_config["selector_rank"]),
top_k=int(draft_config["selector_top_k"]),
params_dtype=vllm_config.model_config.dtype,
prefix=maybe_prefix(prefix, "candidate_selector"),
)
def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
return super().embed_input_ids(input_ids) * self.input_embedding_scale
class DFlash2Qwen3ForCausalLM(DFlashQwen3ForCausalLM):
model_cls = DFlash2Qwen3Model
def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None:
super().__init__(vllm_config=vllm_config, prefix=prefix)
draft_config = self.config.dflash_config
softcap = float(draft_config.get("final_logit_softcapping") or 0.0)
self.candidate_logits_processor = LogitsProcessor(
vllm_config.model_config.get_vocab_size(),
scale=float(draft_config.get("output_multiplier", 1.0)),
soft_cap=softcap if softcap > 0 else None,
)
def compute_candidates(
self, hidden_states: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
return self.candidate_logits_processor.get_top_k_tokens(
self.lm_head, hidden_states, self.model.candidate_selector.top_k
)
EntryClass = DFlash2Qwen3ForCausalLM
Create files/overlay-dflash2/dflash2/__init__.py:
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
Create files/overlay-dflash2/dflash2/speculator.py:
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from typing import Any
import torch
from vllm.config import VllmConfig
from vllm.config.compilation import CUDAGraphMode
# SM121-PORT: tldevice is needed by the local gumbel_noised_argmax below.
from vllm.triton_utils import tl, tldevice, triton
# SM121-PORT: the image's gumbel.py (g487ecf187, ~Aug 15) predates PR #52816,
# which factored gumbel_noised_argmax out of gumbel_block_argmax. Import only
# the rand primitives and define the helper locally (verbatim from the merged
# gumbel.py) so the shared gumbel.py used by the production MTP path stays
# untouched. Same primitives => draft and verification draw identical noise.
from vllm.v1.worker.gpu.sample.gumbel import tl_rand32, tl_rand64
from vllm.v1.worker.gpu.spec_decode.dflash.speculator import DFlashSpeculator
@triton.jit
def gumbel_noised_argmax(
logits,
keys,
mask,
seed,
pos,
temp,
USE_FP64: tl.constexpr,
APPLY_TEMPERATURE: tl.constexpr = True,
):
"""Argmax of logits under Gumbel-max sampling, or plain argmax at temp 0.
`keys` indexes the noise, so the same token draws the same noise wherever it
appears; `pos` and `seed` place the draw in the request's stream, which is
what lets a draft and its verification agree.
SM121-PORT: copied verbatim from the merged vllm/v1/worker/gpu/sample/
gumbel.py (PR #52816, merge b389ac294); see the import note above.
"""
if temp != 0.0 and APPLY_TEMPERATURE:
# Match the behavior of _temperature_kernel: if that kernel uses
# tl.div_rn, this must too.
logits = logits / temp
# fp32 is the default reduction dtype; fp64 is ~1/32-1/64x the throughput
# on H100/Ada/Blackwell and empirically indistinguishable for Gumbel-max.
if USE_FP64:
logits = logits.to(tl.float64)
if temp != 0.0:
gumbel_seed = tl.randint(seed, pos)
if USE_FP64:
u = tl_rand64(gumbel_seed, keys, includes_zero=False)
gumbel_noise = -tl.log(-tl.log(u))
else:
u = tl_rand32(gumbel_seed, keys, includes_zero=False)
# log1p keeps the winning tail at u -> 0, where fp32 resolves it.
gumbel_noise = -tl.log(-tldevice.log1p(-u))
logits = tl.where(mask, logits + gumbel_noise, float("-inf"))
return tl.max(logits, axis=0, return_indices=True)
@triton.jit
def _selector_walk_kernel(
scores_ptr,
candidate_ptr,
sample_pos_ptr,
req_state_ptr,
temperature_ptr,
seeds_ptr,
tokens_ptr,
realized_scores_ptr,
num_steps: tl.constexpr,
top_k: tl.constexpr,
BLOCK_K: tl.constexpr,
SAMPLE_PROBABILISTIC: tl.constexpr,
USE_FP64: tl.constexpr,
):
row = tl.program_id(0)
offsets = tl.arange(0, BLOCK_K)
mask = offsets < top_k
req_state = tl.load(req_state_ptr + row * num_steps)
valid = req_state >= 0
temperature = tl.load(temperature_ptr + req_state, mask=valid, other=0.0)
seed = tl.load(seeds_ptr + req_state, mask=valid, other=0)
previous = 0
for step in range(num_steps):
flat = row * num_steps + step
score_base = (flat * top_k + previous) * top_k
scores = tl.load(
scores_ptr + score_base + offsets,
mask=mask & valid,
other=float("-inf"),
).to(tl.float64 if USE_FP64 else tl.float32)
candidate_base = flat * top_k
candidates = tl.load(
candidate_ptr + candidate_base + offsets,
mask=mask & valid,
other=0,
)
# Candidate ids key the noise, matching the target's own sampling.
position = tl.load(sample_pos_ptr + flat) - 1
_, index = gumbel_noised_argmax(
scores,
candidates,
mask & valid,
seed,
position,
temperature if SAMPLE_PROBABILISTIC else 0.0,
USE_FP64=USE_FP64,
)
tl.store(
realized_scores_ptr + candidate_base + offsets,
scores,
mask=mask & valid,
)
token = tl.load(candidate_ptr + candidate_base + index, mask=valid, other=0)
tl.store(tokens_ptr + flat, token, mask=valid)
previous = index
@triton.jit
def _cache_draft_logits_kernel(
draft_logits_ptr,
cached_candidate_ptr,
candidate_ptr,
scores_ptr,
req_state_ptr,
draft_logits_stride_0,
draft_logits_stride_1,
num_steps: tl.constexpr,
top_k: tl.constexpr,
BLOCK_K: tl.constexpr,
):
flat = tl.program_id(0)
req_state = tl.load(req_state_ptr + flat)
step = flat % num_steps
offsets = tl.arange(0, BLOCK_K)
mask = (req_state >= 0) & (offsets < top_k)
candidate_base = flat * top_k
cache_base = (req_state * num_steps + step) * top_k
old_token_ids = tl.load(cached_candidate_ptr + cache_base + offsets, mask=mask)
logits_base = (
draft_logits_ptr
+ req_state * draft_logits_stride_0
+ step * draft_logits_stride_1
)
tl.store(logits_base + old_token_ids, -float("inf"), mask=mask)
token_ids = tl.load(candidate_ptr + candidate_base + offsets, mask=mask)
scores = tl.load(scores_ptr + candidate_base + offsets, mask=mask)
tl.store(logits_base + token_ids, scores, mask=mask)
tl.store(cached_candidate_ptr + cache_base + offsets, token_ids, mask=mask)
class DFlash2Speculator(DFlashSpeculator):
_speculator_name = "DFlash2"
def __init__(self, vllm_config: VllmConfig, device: torch.device):
super().__init__(vllm_config, device)
draft_config = self.draft_model_config.hf_config.dflash_config
self.selector_top_k = int(draft_config["selector_top_k"])
self._anchor_indices = (
torch.arange(self.max_num_reqs, dtype=torch.int64, device=device)
* self.num_query_per_req
)
self._selector_scores = torch.empty(
self.max_num_reqs,
self.num_speculative_steps,
self.selector_top_k,
dtype=torch.float32,
device=device,
)
self._cached_candidate_ids = torch.zeros(
self._selector_scores.shape, dtype=torch.int64, device=device
)
def draft_logits_spec(self, vllm_config: VllmConfig) -> tuple[torch.dtype, float]:
# fp32 so the walk and the rejection that checks it read the same
# distribution; -inf because the cache kernel writes only the K
# candidates.
return torch.float32, -float("inf")
def _sample_path(
self,
candidate_ids: torch.Tensor,
scores: torch.Tensor,
num_reqs: int,
) -> None:
block_k = triton.next_power_of_2(self.selector_top_k)
_selector_walk_kernel[(num_reqs,)](
scores.contiguous(),
candidate_ids.contiguous(),
self.sample_pos,
self.sample_idx_mapping,
self.temperature,
self.seeds,
self.draft_tokens,
self._selector_scores,
num_steps=self.num_speculative_steps,
top_k=self.selector_top_k,
BLOCK_K=block_k,
SAMPLE_PROBABILISTIC=self.draft_logits is not None,
USE_FP64=self.use_fp64_gumbel,
num_warps=1,
)
def _cache_draft_logits(self, candidate_ids: torch.Tensor, num_sample: int) -> None:
draft_logits = self.draft_logits
assert draft_logits is not None
block_k = triton.next_power_of_2(self.selector_top_k)
_cache_draft_logits_kernel[(num_sample,)](
draft_logits,
self._cached_candidate_ids,
candidate_ids,
self._selector_scores,
self.sample_idx_mapping,
draft_logits.stride(0),
draft_logits.stride(1),
num_steps=self.num_speculative_steps,
top_k=self.selector_top_k,
BLOCK_K=block_k,
num_warps=1,
)
def _generate_draft(
self,
num_reqs: int,
num_tokens_padded: int,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
) -> None:
last_hidden_states = self._run_model(
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
num_sample = num_reqs * self.num_speculative_steps
hidden_states = last_hidden_states[self.sample_indices[:num_sample]].view(
num_reqs, self.num_speculative_steps, -1
)
candidate_ids, unary_logits = self.model.compute_candidates(
hidden_states.flatten(0, 1)
)
candidate_ids = candidate_ids.view(
num_reqs, self.num_speculative_steps, self.selector_top_k
)
unary_logits = unary_logits.view_as(candidate_ids)
anchor_token_ids = self.input_buffers.input_ids[self._anchor_indices[:num_reqs]]
scores = self.model.model.candidate_selector(
candidate_ids,
unary_logits,
hidden_states,
anchor_token_ids,
)
self._sample_path(candidate_ids, scores, num_reqs)
if self.draft_logits is not None:
self._cache_draft_logits(candidate_ids, num_sample)
Create files/overlay-dflash2/patch_registry_and_select.py:
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""SM121-PORT of vLLM PR #52816 (DFlash2: local convolution + candidate
selector, merge commit b389ac294) onto radixark/vllm-glm53-flash:sm121-v8
(vLLM 0.1.dev20051+g487ecf187).
Run INSIDE the image at build time, AFTER copying the new files in:
cp qwen3_dflash2.py $VLLM/model_executor/models/qwen3_dflash2.py
cp -r dflash2/ $VLLM/v1/worker/gpu/spec_decode/dflash2/
python3 patch_registry_and_select.py [optional-vllm-root]
Edits (all anchored string replacements, asserted, ast-checked, idempotent):
1. model_executor/models/registry.py DFlash2DraftModel entry
2. v1/worker/gpu/spec_decode/__init__.py route DFlash2 drafts to the
DFlash2 speculator
3. model_executor/models/qwen3_dflash.py decoder_layer_cls / model_cls
subclass hooks + is_causal
4. v1/worker/gpu/spec_decode/speculator.py draft_logits_spec hook
5. config/vllm.py force V2 model runner for a
DFlash2 draft
6. model_executor/layers/logits_processor.py get_top_k_tokens (verbatim
from the PR)
"""
import ast
import os
import sys
VLLM_ROOT = (
sys.argv[1]
if len(sys.argv) > 1
else "/usr/local/lib/python3.12/dist-packages/vllm"
)
def patch_file(rel_path: str, edits: list[tuple[str, str, str]]) -> None:
"""Apply (marker, old, new) edits to VLLM_ROOT/rel_path.
marker: substring whose presence means the edit is already applied (skip).
old: anchor text that must occur exactly once; replaced by new.
"""
path = os.path.join(VLLM_ROOT, rel_path)
with open(path, encoding="utf-8") as f:
src = f.read()
changed = False
for marker, old, new in edits:
if marker in src:
print(f"[skip] {rel_path}: already applied ({marker[:48]!r})")
continue
assert old in src, f"ANCHOR NOT FOUND in {rel_path}:\n{old[:200]!r}"
assert src.count(old) == 1, (
f"ANCHOR NOT UNIQUE ({src.count(old)}x) in {rel_path}:\n{old[:200]!r}"
)
src = src.replace(old, new)
changed = True
print(f"[edit] {rel_path}: applied ({marker[:48]!r})")
ast.parse(src, filename=path) # syntax gate before writing
if changed:
with open(path, "w", encoding="utf-8") as f:
f.write(src)
print(f"[ok] {rel_path}: written, ast.parse clean")
else:
print(f"[ok] {rel_path}: no changes needed, ast.parse clean")
# --------------------------------------------------------------------------
# 1. registry.py: DFlash2DraftModel entry
# --------------------------------------------------------------------------
patch_file(
"model_executor/models/registry.py",
[
(
'"DFlash2DraftModel"',
' "DFlashDraftModel": ("qwen3_dflash", "DFlashQwen3ForCausalLM"),\n',
' "DFlashDraftModel": ("qwen3_dflash", "DFlashQwen3ForCausalLM"),\n'
" # SM121-PORT PR#52816\n"
' "DFlash2DraftModel": ("qwen3_dflash2", "DFlash2Qwen3ForCausalLM"),\n',
),
],
)
# --------------------------------------------------------------------------
# 2. spec_decode/__init__.py: speculator selection
# --------------------------------------------------------------------------
patch_file(
"v1/worker/gpu/spec_decode/__init__.py",
[
(
"DFlash2Speculator",
' if speculative_config.method == "dflash":\n'
" from vllm.v1.worker.gpu.spec_decode.dflash.speculator import (\n",
' if speculative_config.method == "dflash":\n'
" # SM121-PORT PR#52816: route DFlash2 drafts (declared by\n"
" # architecture, or by a dflash_config carrying a candidate\n"
" # selector) to the DFlash2 speculator. On the plain DFlash path\n"
" # such a checkpoint would silently draft as DFlash1.\n"
" _draft_cfg = speculative_config.draft_model_config\n"
' _dflash_cfg = getattr(_draft_cfg.hf_config, "dflash_config", None) or {}\n'
' if "DFlash2DraftModel" in (_draft_cfg.architectures or []) or (\n'
' "selector_rank" in _dflash_cfg\n'
" ):\n"
" from vllm.v1.worker.gpu.spec_decode.dflash2.speculator import (\n"
" DFlash2Speculator,\n"
" )\n"
"\n"
" return DFlash2Speculator(vllm_config, device)\n"
" from vllm.v1.worker.gpu.spec_decode.dflash.speculator import (\n",
),
],
)
# --------------------------------------------------------------------------
# 3. qwen3_dflash.py: subclass hooks + explicit is_causal resolution
# --------------------------------------------------------------------------
patch_file(
"model_executor/models/qwen3_dflash.py",
[
(
"SM121-PORT causal",
' """``dflash_config.causal`` overrides all layers; else only SWA'
' layers causal."""\n'
' override = (getattr(config, "dflash_config", None) or {}).get("causal")\n',
' """Resolve explicit causality before falling back to legacy layer'
' defaults."""\n'
" # SM121-PORT causal (PR#52816, extended): honor a top-level `is_causal`\n"
" # and a dflash_config-level `is_causal` (the GLM53 DFlash2 drafter ships\n"
" # the latter) before the legacy `causal` key.\n"
' dflash_cfg = getattr(config, "dflash_config", None) or {}\n'
' is_causal = getattr(config, "is_causal", None)\n'
" if is_causal is None:\n"
' is_causal = dflash_cfg.get("is_causal")\n'
" if is_causal is not None:\n"
" return bool(is_causal)\n"
' override = dflash_cfg.get("causal")\n',
),
(
"decoder_layer_cls = DFlashQwen3DecoderLayer",
"class DFlashQwen3Model(nn.Module):\n"
" hf_to_vllm_mapper = WeightsMapper(\n",
"class DFlashQwen3Model(nn.Module):\n"
" # SM121-PORT PR#52816: subclass hook for DFlash2.\n"
" decoder_layer_cls = DFlashQwen3DecoderLayer\n"
"\n"
" hf_to_vllm_mapper = WeightsMapper(\n",
),
(
"self.decoder_layer_cls(",
" DFlashQwen3DecoderLayer(\n"
" current_vllm_config,\n",
" self.decoder_layer_cls( # SM121-PORT PR#52816\n"
" current_vllm_config,\n",
),
(
"model_cls = DFlashQwen3Model",
"class DFlashQwen3ForCausalLM(Qwen3ForCausalLM):\n"
' def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):\n',
"class DFlashQwen3ForCausalLM(Qwen3ForCausalLM):\n"
" # SM121-PORT PR#52816: subclass hook for DFlash2.\n"
" model_cls = DFlashQwen3Model\n"
"\n"
' def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):\n',
),
(
"self.model = self.model_cls(",
" self.model = DFlashQwen3Model(\n"
" vllm_config=vllm_config,\n",
" self.model = self.model_cls( # SM121-PORT PR#52816\n"
" vllm_config=vllm_config,\n",
),
],
)
# --------------------------------------------------------------------------
# 4. spec_decode/speculator.py: draft_logits_spec hook
# --------------------------------------------------------------------------
patch_file(
"v1/worker/gpu/spec_decode/speculator.py",
[
(
"self.draft_logits_spec(",
" self.draft_logits = torch.zeros(\n"
" self.max_num_reqs,\n"
" self.num_speculative_steps,\n"
" self.vocab_size,\n"
" dtype=vllm_config.model_config.head_dtype,\n"
" device=device,\n"
" )\n",
" # SM121-PORT PR#52816: dtype/fill via draft_logits_spec so\n"
" # DFlash2 can cache a sparse fp32/-inf distribution.\n"
" dtype, fill = self.draft_logits_spec(vllm_config)\n"
" self.draft_logits = torch.full(\n"
" (\n"
" self.max_num_reqs,\n"
" self.num_speculative_steps,\n"
" self.vocab_size,\n"
" ),\n"
" fill,\n"
" dtype=dtype,\n"
" device=device,\n"
" )\n",
),
(
"def draft_logits_spec(",
" def _validate_local_argmax_reduction(self) -> None:\n",
" def draft_logits_spec(\n"
" self, vllm_config: VllmConfig\n"
" ) -> tuple[torch.dtype, float]:\n"
' """Dtype and fill for the cached proposal distribution.\n'
"\n"
" Speculators that write only a subset of columns each step\n"
" override this. (SM121-PORT PR#52816)\n"
' """\n'
" return vllm_config.model_config.head_dtype, 0.0\n"
"\n"
" def _validate_local_argmax_reduction(self) -> None:\n",
),
],
)
# --------------------------------------------------------------------------
# 5. config/vllm.py: force the V2 model runner for a DFlash2 draft
# --------------------------------------------------------------------------
patch_file(
"config/vllm.py",
[
(
"_is_dflash2_draft():",
" if self._dflash_needs_multi_kv_group():\n"
" return True\n",
" if self._dflash_needs_multi_kv_group():\n"
" return True\n"
"\n"
" # SM121-PORT PR#52816: the DFlash2 candidate selector exists only\n"
" # in the V2 speculator; on V1 the same checkpoint would silently\n"
" # draft as DFlash1. Force V2 as for dspark.\n"
" if self._is_dflash2_draft():\n"
" return True\n",
),
(
"def _is_dflash2_draft(",
" def _dflash_needs_multi_kv_group(self) -> bool:\n",
" def _is_dflash2_draft(self) -> bool:\n"
' """SM121-PORT PR#52816: whether the DFlash draft is a DFlash2\n'
" one, by the same signals the speculator selection uses\n"
' (v1/worker/gpu/spec_decode/__init__.py)."""\n'
" spec = self.speculative_config\n"
' if spec is None or spec.method != "dflash":\n'
" return False\n"
' draft_config = getattr(spec, "draft_model_config", None)\n'
" if draft_config is None:\n"
" return False\n"
' if "DFlash2DraftModel" in (draft_config.architectures or []):\n'
" return True\n"
" dflash_cfg = (\n"
' getattr(draft_config.hf_config, "dflash_config", None) or {}\n'
" )\n"
' return "selector_rank" in dflash_cfg\n'
"\n"
" def _dflash_needs_multi_kv_group(self) -> bool:\n",
),
],
)
# --------------------------------------------------------------------------
# 6. logits_processor.py: get_top_k_tokens (verbatim from the merged PR)
# --------------------------------------------------------------------------
patch_file(
"model_executor/layers/logits_processor.py",
[
(
"_flashinfer_topk",
"from vllm.platforms import current_platform\n",
"from vllm.platforms import current_platform\n"
"\n"
"# SM121-PORT PR#52816: vocab-parallel top-k for the DFlash2 candidate\n"
"# selector -- verbatim from the merged logits_processor.py (b389ac294).\n"
"from collections.abc import Callable\n"
"from functools import cache\n"
"\n"
"from vllm.logger import init_logger\n"
"from vllm.utils.flashinfer import has_flashinfer\n"
"\n"
"logger = init_logger(__name__)\n"
"\n"
"\n"
"@cache\n"
"def _flashinfer_topk() -> (\n"
" Callable[..., tuple[torch.Tensor, torch.Tensor]] | None\n"
"):\n"
' """FlashInfer\'s radix top-k, or None for torch.topk.\n'
"\n"
" The top-k spans the vocabulary, where the radix kernel is about twice\n"
" torch.topk.\n"
' """\n'
" if not current_platform.is_cuda():\n"
" return None\n"
" if not has_flashinfer():\n"
" logger.info_once(\n"
' "flashinfer is unavailable; vocab-parallel top-k uses '
'torch.topk, "\n'
' "at roughly half the speed."\n'
" )\n"
" return None\n"
" from flashinfer import top_k\n"
"\n"
" return top_k\n"
"\n"
"\n"
"def _topk(scores: torch.Tensor, k: int) -> tuple[torch.Tensor, torch.Tensor]:\n"
" impl = _flashinfer_topk()\n"
" if impl is None or not scores.is_cuda:\n"
" return torch.topk(scores, k, dim=-1)\n"
" return impl(scores, k, sorted=True, deterministic=True)\n",
),
(
"def get_top_k_tokens(",
" def extra_repr(self) -> str:\n",
" # SM121-PORT PR#52816: verbatim from the merged logits_processor.py.\n"
" def get_top_k_tokens(\n"
" self,\n"
" lm_head: VocabParallelEmbedding,\n"
" hidden_states: torch.Tensor,\n"
" k: int,\n"
" embedding_bias: torch.Tensor | None = None,\n"
" ) -> tuple[torch.Tensor, torch.Tensor]:\n"
' """Vocab-parallel top-k without all-gathering full logits.\n'
"\n"
" The `get_top_tokens` reduction widened from one token to k,\n"
" returning the values as well as the global ids. Communication is\n"
" O(batch * 2k * tp_size) rather than O(batch * vocab_size).\n"
"\n"
" Scale and soft cap are applied to the k selected values rather\n"
" than the whole vocabulary; both are monotonic, so the selection\n"
" is the same and only k entries are touched.\n"
' """\n'
" if self.scale <= 0.0 and self.scale != 1.0:\n"
" raise ValueError(\n"
' "The local top-k reduction optimization is not supported '
'for "\n'
' "non-positive logit scaling factors."\n'
" )\n"
"\n"
" logits = self._apply_head(lm_head, hidden_states, embedding_bias)\n"
"\n"
" # Mask out padding entries beyond org_vocab_size on this shard.\n"
" num_pad = lm_head.shard_indices.num_org_vocab_padding\n"
" if num_pad > 0:\n"
' logits[..., -num_pad:] = -float("inf")\n'
"\n"
" values, ids = _topk(logits, k)\n"
" # Convert shard-local indices to global vocab indices.\n"
" ids = ids.to(torch.int64) + lm_head.shard_indices.org_vocab_start_index\n"
"\n"
" if lm_head.tp_size > 1:\n"
" values = tensor_model_parallel_all_gather(values, dim=-1)\n"
" ids = tensor_model_parallel_all_gather(ids, dim=-1)\n"
" values, selected = _topk(values, k)\n"
" ids = ids.gather(-1, selected)\n"
"\n"
" values = values.float()\n"
" if self.scale != 1.0:\n"
" values = values * self.scale\n"
" if self.soft_cap is not None:\n"
" values = torch.tanh(values / self.soft_cap) * self.soft_cap\n"
" return ids, values\n"
"\n"
" def extra_repr(self) -> str:\n",
),
],
)
print("\nAll patches applied and ast-checked. VLLM_ROOT =", VLLM_ROOT)
Create files/overlay-dflash2/patch_glm_aux_capture.py:
#!/usr/bin/env python3
"""
patch_glm_aux_capture.py -- DFlash2 / EAGLE-3 target-side glue for GLM-5.3-Flash.
Edits vllm/models/glm5next/nvidia/model.py IN PLACE (build-time, inside the
radixark/vllm-glm53-flash:sm121-v8 image) so the Glm5Next* target model can
feed aux hidden states to a DFlash2 draft model:
1. imports EagleModelMixin + SupportsEagle3 from
vllm.model_executor.models.interfaces
2. Glm5NextModel gains EagleModelMixin (provides the
`aux_hidden_state_layers` store + `_set_aux_hidden_state_layers`)
3. Glm5NextForCausalLM and Glm5NextForConditionalGeneration declare
SupportsEagle3 (the interface's concrete default methods implement
set_aux_hidden_state_layers / get_eagle3_default_aux_hidden_state_layers,
routing through `.language_model` / `.model` automatically)
4. the decoder-layer loop captures aux hidden states with mHC contraction
(deferred hc_post materialized, then hc_contract == mean over hc
streams), mirroring vllm/models/deepseek_v4/nvidia/model.py
DeepseekV4Model.forward (aux capture at ~lines 1160-1213:
`if idx + 1 in self.aux_hidden_state_layers: ... aux_recon.mean(dim=1)`
+ sp_all_gather)
5. Glm5NextModel.forward returns (hidden_states, aux_hidden_states) when
any layer was captured, exactly like DeepseekV4Model.forward
Layer-index semantics: gpu_model_runner._get_eagle3_aux_layers_from_config
converts DFlash `target_layer_ids` to id+1 before calling
set_aux_hidden_state_layers ("# Add 1 to convert DFlash's aux layer id
semantics"). The capture below therefore tests `idx + 1 in
self.aux_hidden_state_layers` -- identical to deepseek_v4 -- which captures
the OUTPUT of 0-based decoder layer `idx`. Net effect for
target_layer_ids=[5,14,24,33,42]: outputs of layers 5,14,24,33,42 are
captured, each contracted to [num_tokens, hidden_size].
Usage:
python3 patch_glm_aux_capture.py [--model-file PATH] [--dry-run]
Idempotent: re-running on an already-patched file is a no-op (exit 0).
Fails loudly (AssertionError, nonzero exit) if any anchor is missing.
"""
from __future__ import annotations
import argparse
import ast
import sys
DEFAULT_MODEL_FILE = (
"/usr/local/lib/python3.12/dist-packages/vllm/models/glm5next/nvidia/model.py"
)
MARKER = "DFLASH2-AUX-CAPTURE"
# ---------------------------------------------------------------------------
# Anchored edits. Every anchor must appear EXACTLY ONCE in the target file.
# ---------------------------------------------------------------------------
EDIT_IMPORTS_ANCHOR = """\
from vllm.model_executor.models.interfaces import (
HasInnerState,
IsHybrid,
MixtureOfExperts,
SupportsPP,
)
"""
EDIT_IMPORTS_NEW = """\
from vllm.model_executor.models.interfaces import (
EagleModelMixin,
HasInnerState,
IsHybrid,
MixtureOfExperts,
SupportsEagle3,
SupportsPP,
)
"""
EDIT_MODEL_CLASS_ANCHOR = """\
class Glm5NextModel(nn.Module):
"""
EDIT_MODEL_CLASS_NEW = """\
class Glm5NextModel(nn.Module, EagleModelMixin):
"""
EDIT_LOOP_ANCHOR = """\
for layer in self._active_layers:
hidden_states, residual, post, comb = layer(
positions, hidden_states, residual, post, comb
)
"""
EDIT_LOOP_NEW = """\
# DFLASH2-AUX-CAPTURE (EAGLE-3 aux hidden states; mirrors
# DeepseekV4Model.forward in vllm/models/deepseek_v4/nvidia/model.py).
aux_hidden_states: list[torch.Tensor] = []
for idx, layer in enumerate(self._active_layers, start=self.start_layer):
hidden_states, residual, post, comb = layer(
positions, hidden_states, residual, post, comb
)
if idx + 1 in self.aux_hidden_state_layers:
# `idx + 1` matches deepseek_v4: the runner already converted
# DFlash target_layer_ids to id+1 semantics
# (gpu_model_runner._get_eagle3_aux_layers_from_config), so
# this captures the OUTPUT of 0-based decoder layer `idx`.
if post is not None:
# Mid-stack mHC layer: its final hc_post is deferred to
# the next layer's fused pre. Materialize the multi-stream
# reconstruction here (pure op -- the deferred
# residual/post/comb state is not mutated), then contract
# hc streams exactly like the last layer does.
# hc_contract == mean over streams, the same contraction
# deepseek_v4 uses (aux_recon.mean(dim=1)).
aux_recon = layer.hc_post(hidden_states, residual, post, comb)
aux_hidden_state = hc_contract(aux_recon, layer.n)
else:
# Last mHC layer (already hc_post + hc_contract'ed inside
# the layer) or a non-mHC layer: the output is already
# plain [num_tokens, hidden_size].
aux_hidden_state = hidden_states
if self.is_sequence_parallel:
# Aux states are consumed at full-sequence granularity;
# gather the SP shard (deepseek_v4 pattern).
aux_hidden_state = sp_all_gather(aux_hidden_state)[
:full_num_tokens
]
aux_hidden_states.append(aux_hidden_state)
"""
EDIT_RETURN_ANCHOR = """\
hidden_states = self.norm(hidden_states)
return hidden_states
"""
EDIT_RETURN_NEW = """\
hidden_states = self.norm(hidden_states)
if len(aux_hidden_states) > 0:
# (final_hidden_states, list-of-aux) -- gpu_model_runner unpacks
# this tuple when use_aux_hidden_state_outputs is set; identical
# to DeepseekV4Model.forward's aux return.
return hidden_states, aux_hidden_states
return hidden_states
"""
EDIT_CAUSAL_LM_ANCHOR = """\
class Glm5NextForCausalLM(
nn.Module, HasInnerState, SupportsPP, MixtureOfExperts, IsHybrid
):
"""
EDIT_CAUSAL_LM_NEW = """\
class Glm5NextForCausalLM(
nn.Module, HasInnerState, SupportsPP, SupportsEagle3, MixtureOfExperts, IsHybrid
):
"""
EDIT_COND_GEN_ANCHOR = """\
class Glm5NextForConditionalGeneration(
Glm4vForConditionalGeneration, HasInnerState, IsHybrid
):
"""
EDIT_COND_GEN_NEW = """\
class Glm5NextForConditionalGeneration(
Glm4vForConditionalGeneration, HasInnerState, IsHybrid, SupportsEagle3
):
"""
EDITS: list[tuple[str, str, str]] = [
("import EagleModelMixin + SupportsEagle3", EDIT_IMPORTS_ANCHOR, EDIT_IMPORTS_NEW),
("Glm5NextModel gains EagleModelMixin", EDIT_MODEL_CLASS_ANCHOR, EDIT_MODEL_CLASS_NEW),
("decoder loop: aux hidden state capture + mHC contraction", EDIT_LOOP_ANCHOR, EDIT_LOOP_NEW),
("forward tail: return (hidden_states, aux_hidden_states)", EDIT_RETURN_ANCHOR, EDIT_RETURN_NEW),
("Glm5NextForCausalLM declares SupportsEagle3", EDIT_CAUSAL_LM_ANCHOR, EDIT_CAUSAL_LM_NEW),
("Glm5NextForConditionalGeneration declares SupportsEagle3", EDIT_COND_GEN_ANCHOR, EDIT_COND_GEN_NEW),
]
def patch_file(path: str, dry_run: bool = False) -> int:
with open(path, "r", encoding="utf-8") as f:
text = f.read()
if MARKER in text:
print(f"[patch_glm_aux_capture] {path}: already patched ({MARKER} marker found); no-op.")
return 0
# Sanity: the file we expect (guards against pointing at the wrong tree).
for required in ("class Glm5NextModel", "hc_contract", "sp_all_gather"):
assert required in text, (
f"ANCHOR PRECHECK FAILED: {required!r} not found in {path} -- "
"is this really glm5next/nvidia/model.py?"
)
applied = []
for name, anchor, replacement in EDITS:
n = text.count(anchor)
assert n == 1, (
f"ANCHOR FAILED for edit [{name}]: expected exactly 1 occurrence, "
f"found {n}. The upstream file has drifted -- re-derive the anchor "
f"before building.\n--- anchor ---\n{anchor}\n--------------"
)
text = text.replace(anchor, replacement, 1)
applied.append(name)
# The patched source must still be valid Python.
try:
ast.parse(text, filename=path)
except SyntaxError as e:
raise AssertionError(f"POST-EDIT ast.parse FAILED for {path}: {e}") from e
if dry_run:
print(f"[patch_glm_aux_capture] DRY RUN -- {path} not written.")
else:
with open(path, "w", encoding="utf-8") as f:
f.write(text)
print(f"[patch_glm_aux_capture] {path}: {len(applied)} edits applied:")
for name in applied:
print(f" - {name}")
print("[patch_glm_aux_capture] ast.parse OK.")
return 0
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__.splitlines()[1])
ap.add_argument("--model-file", default=DEFAULT_MODEL_FILE)
ap.add_argument("--dry-run", action="store_true", help="validate anchors + parse, write nothing")
args = ap.parse_args()
return patch_file(args.model_file, dry_run=args.dry_run)
if __name__ == "__main__":
sys.exit(main())
Create files/overlay-dflash2/patch_glm5_drafter_group.py:
#!/usr/bin/env python3
"""
patch_glm5_drafter_group.py -- teach the GLM-5-Next KV layout about the
DFlash2 drafter's SlidingWindowSpec layers.
Edits vllm/v1/core/kv_cache_utils.py IN PLACE (build-time, inside the
radixark/vllm-glm53-flash sm121 image).
PROBLEM
-------
`_get_kv_cache_groups_glm5_next` returns None the moment any non-mamba /
non-tail spec is not exactly MLAAttentionSpec. The DFlash2 drafter registers 5
plain SlidingWindowSpec layers, so the whole model drops to the generic
uniform-page path -- which provably cannot serve GLM-5.3-Flash: page
unification rescales the kpool tail's block away from its pool size and boot
dies at warmup's `assert tail_kv_cache.shape[2] == pool_size` (see
~/lane1_fail6.log / ~/lane1_fail7.log on Reddie).
DESIGN
------
Keep the GLM-5-Next fast path bit-for-bit identical for the base model and
extend it with ONE extra group for the drafter, appended LAST (existing group
ids stay stable). Two modes, decided from the geometry:
EXACT FIT (preferred; both deployed geometries land here): rescale the
drafter's block size so its REAL page equals the MLA page exactly
(block = mla_page // drafter_bytes_per_token), and let drafter layer i
co-own MLA tensor i (`shared_by`) at disjoint block ids from the one shared
BlockPool -- like mamba. The per-block byte cost of the pool is UNCHANGED,
so KV capacity stays at the base model's; the sliding window bounds the
drafter to a handful of block ids per request.
CRITICAL, learned from boot 8 (~/lane1_fail8.log): the drafter spec must
NOT use `page_size_padded`. A padded spec routes the runner into the
strided-view reshape (`_reshape_attention_kv_cache` in
vllm/v1/worker/gpu/attn_utils.py), and that view is INVALID whenever the
backend virtually splits the manager block into smaller kernel blocks
(FlashInfer registered int kernel sizes and picked 64 for a 2304-token
manager block; the strided path then applied the full per-page stride to
each KERNEL block: 5760 x 2,359,296 B demanded from a 160 x 2,359,296 B
tensor -> setStorage out of bounds). With an exact fit the ordinary
CONTIGUOUS view is correct under any kernel split: kernel block j of
manager block b lands at b * mla_page + j * kernel_page, inside block b's
own page, so slot-sharing with mamba stays sound. Exact fit is gated on:
- mla_page divisible by the drafter's bytes/token;
- fit block divisible by 64 (covers the 16/32/64 int kernel sizes the
SWA backends register, so select_common_block_size never fails);
- fit block and MLA block divide one another (keeps
resolve_kv_cache_block_sizes' scheduler LCM at their max);
- at most as many drafter layers as MLA tensors to ride in.
STANDALONE (fallback for geometries that cannot exactly fill the MLA
page): the drafter spec is kept as-is and its layers get compact per-layer
tensors of their own (size draft_page * num_blocks), added to the
per-block byte cost everywhere it is computed. Contiguous reshape again --
no padding, any kernel split valid.
`_glm5_next_tensor_layout` detects the drafter group (uniform SWA, never
padded) and returns it as a 9th tuple element; the three consumers are
updated in lock-step so detection, tensor emission, page accounting and the
available-memory check can never disagree:
- `get_kv_cache_config_from_groups`: exact fit -> drafter layer i joins MLA
tensor i's shared_by; standalone -> per-layer drafter tensors + per-block
cost;
- `_pool_bytes_per_block`: standalone drafter pages only (exact fit adds
no bytes);
- `_max_memory_usage_bytes_from_groups`: charges the drafter's window-
bounded block-id demand at the per-block byte sum (incl. standalone
drafter pages).
Runner-side audit (no edits needed there):
- init_attn_backend builds per-group AttentionGroups generically; the
drafter group's UniformTypeKVCacheSpecs unwraps to the per-layer SWA
spec; prepare_kernel_block_sizes may pick a smaller kernel block --
fine, both modes reshape through the contiguous path.
- _reshape_kv_cache: num_blocks = raw.numel() // page_size_bytes is the
pool's num_blocks in both modes (exact fit: the MLA tensor divided by
mla_page; standalone: the compact tensor divided by draft_page).
- _kv_first_layers_sharing_pool_with_mamba: blocks-first SWA backends
report block_dim 0, so no page-aligned restride is triggered; the
exact-fit contiguous view is already page-aligned per manager block.
- Scheduler: generate_scheduler_kv_cache_config unwraps the group to a
SlidingWindowSpec -> SlidingWindowManager; HybridKVCacheCoordinator's
verify_and_split handles an extra participating spec group generically.
- Speculator (dflash2): set_attn calls init_attn_backend with
active_layer_names=draft layers; the drafter group id indexes
BlockTables.input_block_tables generically.
Usage:
python3 patch_glm5_drafter_group.py [--kv-file PATH] [--dry-run]
Idempotent: re-running on an already-patched file is a no-op (exit 0).
Fails loudly (AssertionError, nonzero exit) if any anchor is missing.
"""
from __future__ import annotations
import argparse
import ast
import sys
DEFAULT_KV_FILE = (
"/usr/local/lib/python3.12/dist-packages/vllm/v1/core/kv_cache_utils.py"
)
MARKER = "DFLASH2-DRAFTER-GROUP"
# ---------------------------------------------------------------------------
# Anchored edits. Every anchor must appear EXACTLY ONCE in the target file.
# ---------------------------------------------------------------------------
# -- _get_kv_cache_groups_glm5_next: partition drafter layers out ------------
EDIT_PARTITION_ANCHOR = """\
attn_specs = {
k: v
for k, v in kv_cache_spec.items()
if not isinstance(v, (MambaSpec, KpoolTailSpec))
}
if not mamba_specs or not all(
type(s) is MLAAttentionSpec for s in attn_specs.values()
):
return None
"""
EDIT_PARTITION_NEW = """\
# DFLASH2-DRAFTER-GROUP: a spec-decode drafter (DFlash2) adds plain
# SlidingWindowSpec layers on top of the GLM-5-Next hybrid. Partition them
# out (exact type: KpoolTailSpec subclasses SlidingWindowSpec) so they do
# not disqualify the model from this fast path; they are appended as one
# extra group below.
draft_specs = {
k: v for k, v in kv_cache_spec.items() if type(v) is SlidingWindowSpec
}
attn_specs = {
k: v
for k, v in kv_cache_spec.items()
if not isinstance(v, (MambaSpec, KpoolTailSpec))
and type(v) is not SlidingWindowSpec
}
if not mamba_specs or not all(
type(s) is MLAAttentionSpec for s in attn_specs.values()
):
return None
"""
# -- _get_kv_cache_groups_glm5_next: build + append the drafter group --------
EDIT_GROUPS_RETURN_ANCHOR = """\
mamba_grouped_names: list[list[str]] = [[] for _ in range(num_groups)]
for k, name in enumerate(mamba_specs):
mamba_grouped_names[k % num_groups].append(name)
return (
[KVCacheGroupSpec(list(attn_specs), uniform_spec)]
+ ([tail_group] if tail_group is not None else [])
+ create_kv_cache_group_specs(padded_specs, mamba_grouped_names)
)
"""
EDIT_GROUPS_RETURN_NEW = """\
mamba_grouped_names: list[list[str]] = [[] for _ in range(num_groups)]
for k, name in enumerate(mamba_specs):
mamba_grouped_names[k % num_groups].append(name)
# Drafter group (DFLASH2-DRAFTER-GROUP): one extra group for the spec-
# decode drafter's SlidingWindowSpec layers, appended LAST so existing
# group ids stay stable. NEVER page_size_padded: a padded spec routes the
# runner into the strided-view reshape, which is invalid when the backend
# virtually splits the manager block into smaller kernel blocks (boot 8:
# FlashInfer picked kernel block 64 for a 2304-token manager block and
# the per-KERNEL-block page stride blew past the tensor). Both modes
# below use the ordinary contiguous reshape, valid under any split.
draft_group = None
if draft_specs:
any_draft = next(iter(draft_specs.values()))
assert all(spec == any_draft for spec in draft_specs.values()), (
"drafter SlidingWindowSpec layers must share one spec"
)
draft_bytes_per_token = any_draft.page_size_bytes // any_draft.block_size
mla_block = mla_specs[mla_names[0]].block_size
fit_block = (
mla_page // draft_bytes_per_token
if mla_page % draft_bytes_per_token == 0
else 0
)
if (
fit_block
# A 64-divisible manager block is divisible by every int kernel
# block size the SWA backends register (16/32/64), so
# select_common_block_size always finds a clean split.
and fit_block % 64 == 0
# Keep resolve_kv_cache_block_sizes' scheduler LCM at
# max(mla_block, fit_block) instead of exploding.
and (fit_block % mla_block == 0 or mla_block % fit_block == 0)
and len(draft_specs) <= len(mla_names)
):
# EXACT FIT: the drafter's real page equals the MLA page, so
# drafter layer i co-owns MLA tensor i at disjoint block ids
# (like mamba) with a contiguous view: kernel block j of manager
# block b lands at b * mla_page + j * kernel_page, inside block
# b's own page. Per-block pool cost unchanged.
new_draft_specs: dict[str, KVCacheSpec] = {
name: replace(s, block_size=fit_block)
for name, s in draft_specs.items()
}
else:
# STANDALONE: the drafter's geometry cannot exactly fill the MLA
# page; keep its spec as-is and give its layers compact tensors
# of their own (emitted in get_kv_cache_config_from_groups and
# charged in the per-block cost).
new_draft_specs = dict(draft_specs)
draft_uniform = UniformTypeKVCacheSpecs.from_specs(new_draft_specs)
assert draft_uniform is not None
draft_group = KVCacheGroupSpec(list(new_draft_specs), draft_uniform)
return (
[KVCacheGroupSpec(list(attn_specs), uniform_spec)]
+ ([tail_group] if tail_group is not None else [])
+ create_kv_cache_group_specs(padded_specs, mamba_grouped_names)
+ ([draft_group] if draft_group is not None else [])
)
"""
# -- _glm5_next_tensor_layout: return-type annotation ------------------------
EDIT_LAYOUT_ANNOT_ANCHOR = """\
list[str],
int,
]
| None
):
"""
EDIT_LAYOUT_ANNOT_NEW = """\
list[str],
int,
KVCacheGroupSpec | None,
]
| None
):
"""
# -- _glm5_next_tensor_layout: docstring Returns -----------------------------
EDIT_LAYOUT_DOC_ANCHOR = """\
- (attn_group, mamba_groups, mla_names, idx_names, mla_page, idx_page,
tail_names, tail_page)
"""
EDIT_LAYOUT_DOC_NEW = """\
- (attn_group, mamba_groups, mla_names, idx_names, mla_page, idx_page,
tail_names, tail_page, draft_group)
"""
# -- _glm5_next_tensor_layout: detect the drafter group ----------------------
EDIT_LAYOUT_DETECT_ANCHOR = """\
attn_group: KVCacheGroupSpec | None = None
tail_group: KVCacheGroupSpec | None = None
for g in uniform_groups:
group_inner = cast(UniformTypeKVCacheSpecs, g.kv_cache_spec).kv_cache_specs
if all(type(s) is MLAAttentionSpec for s in group_inner.values()):
attn_group = g
elif all(isinstance(s, KpoolTailSpec) for s in group_inner.values()):
tail_group = g
"""
EDIT_LAYOUT_DETECT_NEW = """\
attn_group: KVCacheGroupSpec | None = None
tail_group: KVCacheGroupSpec | None = None
draft_group: KVCacheGroupSpec | None = None
for g in uniform_groups:
group_inner = cast(UniformTypeKVCacheSpecs, g.kv_cache_spec).kv_cache_specs
if all(type(s) is MLAAttentionSpec for s in group_inner.values()):
attn_group = g
elif all(isinstance(s, KpoolTailSpec) for s in group_inner.values()):
tail_group = g
elif group_inner and all(
type(s) is SlidingWindowSpec for s in group_inner.values()
):
# DFLASH2-DRAFTER-GROUP: the spec-decode drafter's SWA group
# (validated below once mla_page is known).
draft_group = g
"""
# -- _glm5_next_tensor_layout: validate the drafter group --------------------
EDIT_LAYOUT_VALIDATE_ANCHOR = """\
if any(g.kv_cache_spec.page_size_bytes != mla_page for g in mamba_groups):
return None
tail_names: list[str] = []
"""
EDIT_LAYOUT_VALIDATE_NEW = """\
if any(g.kv_cache_spec.page_size_bytes != mla_page for g in mamba_groups):
return None
if draft_group is not None:
# DFLASH2-DRAFTER-GROUP: one uniform page across drafter layers and
# NEVER page_size_padded (a padded drafter view is invalid under
# kernel block splitting; see _get_kv_cache_groups_glm5_next).
# page == mla_page means exact-fit slot-sharing of the MLA tensors
# (needs one tensor per drafter layer); any other page means
# standalone drafter tensors.
draft_inner = cast(
UniformTypeKVCacheSpecs, draft_group.kv_cache_spec
).kv_cache_specs
draft_pages = {s.page_size_bytes for s in draft_inner.values()}
if len(draft_pages) != 1:
return None
if any(s.page_size_padded is not None for s in draft_inner.values()):
return None
if (
draft_pages.pop() == mla_page
and len(draft_group.layer_names) > len(mla_names)
):
return None
tail_names: list[str] = []
"""
# -- _glm5_next_tensor_layout: return the drafter group ----------------------
EDIT_LAYOUT_RETURN_ANCHOR = """\
return (
attn_group,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_pages.pop(),
tail_names,
tail_page,
)
"""
EDIT_LAYOUT_RETURN_NEW = """\
return (
attn_group,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_pages.pop(),
tail_names,
tail_page,
draft_group,
)
"""
# -- _pool_bytes_per_block: 9-tuple + standalone drafter bytes ---------------
EDIT_POOL_BYTES_ANCHOR = """\
_, _, mla_names, idx_names, mla_page, idx_page, _, _ = glm5
return len(mla_names) * mla_page + len(idx_names) * idx_page
"""
EDIT_POOL_BYTES_NEW = """\
# DFLASH2-DRAFTER-GROUP: an exact-fit drafter (page == mla_page)
# slot-shares the MLA tensors and adds no bytes; a standalone drafter
# adds one page per drafter layer.
_, _, mla_names, idx_names, mla_page, idx_page, _, _, draft_group = glm5
per_block = len(mla_names) * mla_page + len(idx_names) * idx_page
if draft_group is not None:
draft_page = next(
iter(
cast(
UniformTypeKVCacheSpecs, draft_group.kv_cache_spec
).kv_cache_specs.values()
)
).page_size_bytes
if draft_page != mla_page:
per_block += len(draft_group.layer_names) * draft_page
return per_block
"""
# -- get_kv_cache_config_from_groups: destructure + drafter mode -------------
EDIT_CONFIG_DESTRUCTURE_ANCHOR = """\
(
_,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_page,
tail_names,
_tail_page,
) = glm5n
"""
EDIT_CONFIG_DESTRUCTURE_NEW = """\
(
_,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_page,
tail_names,
_tail_page,
draft_group,
) = glm5n
draft_names: list[str] = []
draft_page = 0
draft_shared = False
if draft_group is not None:
draft_names = list(draft_group.layer_names)
draft_page = next(
iter(
cast(
UniformTypeKVCacheSpecs, draft_group.kv_cache_spec
).kv_cache_specs.values()
)
).page_size_bytes
# Exact fit: the drafter's real page equals the MLA page, so it
# rides the MLA tensors; otherwise it gets standalone tensors.
draft_shared = draft_page == mla_page
"""
# -- get_kv_cache_config_from_groups: per-block cost (standalone mode) -------
EDIT_CONFIG_PER_BLOCK_ANCHOR = """\
per_block = len(mla_names) * mla_page + len(idx_names) * idx_page
num_blocks = available_memory // per_block
"""
EDIT_CONFIG_PER_BLOCK_NEW = """\
per_block = len(mla_names) * mla_page + len(idx_names) * idx_page
if draft_names and not draft_shared:
# DFLASH2-DRAFTER-GROUP (standalone): drafter tensors are part of
# every block's byte cost.
per_block += len(draft_names) * draft_page
num_blocks = available_memory // per_block
"""
# -- get_kv_cache_config_from_groups: drafter co-owns MLA tensor i -----------
EDIT_CONFIG_SHARED_BY_ANCHOR = """\
shared_by=[mla_name]
+ [g.layer_names[i] for g in mamba_groups if i < len(g.layer_names)],
)
for i, mla_name in enumerate(mla_names)
"""
EDIT_CONFIG_SHARED_BY_NEW = """\
shared_by=[mla_name]
+ [g.layer_names[i] for g in mamba_groups if i < len(g.layer_names)]
# DFLASH2-DRAFTER-GROUP (exact fit): drafter layer i rides MLA
# tensor i (contiguous view, disjoint block ids), like mamba.
+ ([draft_names[i]] if draft_shared and i < len(draft_names) else []),
)
for i, mla_name in enumerate(mla_names)
"""
# -- get_kv_cache_config_from_groups: standalone drafter tensors -------------
EDIT_CONFIG_DRAFT_TENSORS_ANCHOR = """\
KVCacheTensor(
size=idx_page * num_blocks,
shared_by=(
[idx_names[i], tail_names[i]] if tail_names else [idx_names[i]]
),
)
for i in range(len(idx_names))
]
"""
EDIT_CONFIG_DRAFT_TENSORS_NEW = """\
KVCacheTensor(
size=idx_page * num_blocks,
shared_by=(
[idx_names[i], tail_names[i]] if tail_names else [idx_names[i]]
),
)
for i in range(len(idx_names))
] + [
# DFLASH2-DRAFTER-GROUP (standalone): compact per-layer drafter
# tensors; contiguous reshape, safe under kernel block splitting.
KVCacheTensor(size=draft_page * num_blocks, shared_by=[name])
for name in ([] if draft_shared else draft_names)
]
"""
# -- _max_memory_usage_bytes_from_groups: destructure ------------------------
EDIT_MAXMEM_DESTRUCTURE_ANCHOR = """\
(
attn_group,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_page,
tail_names,
_tail_page,
) = glm5n
"""
EDIT_MAXMEM_DESTRUCTURE_NEW = """\
(
attn_group,
mamba_groups,
mla_names,
idx_names,
mla_page,
idx_page,
tail_names,
_tail_page,
draft_group,
) = glm5n
"""
# -- _max_memory_usage_bytes_from_groups: drafter demand + per-block ---------
EDIT_MAXMEM_BLOCKS_ANCHOR = """\
if tail_names:
# Tail: 1 block/req (KpoolTailSpec.max_admission_blocks_per_request
# == 1), drawn from the shared pool.
blocks_needed += 1
return blocks_needed * (len(mla_names) * mla_page + len(idx_names) * idx_page)
"""
EDIT_MAXMEM_BLOCKS_NEW = """\
if tail_names:
# Tail: 1 block/req (KpoolTailSpec.max_admission_blocks_per_request
# == 1), drawn from the shared pool.
blocks_needed += 1
per_block = len(mla_names) * mla_page + len(idx_names) * idx_page
if draft_group is not None:
# DFLASH2-DRAFTER-GROUP: charge the drafter's window-bounded
# block-id demand; a standalone drafter also adds its pages to
# every block's byte cost (an exact-fit one rides the MLA
# tensors and adds none).
draft_uniform = draft_group.kv_cache_spec
assert isinstance(draft_uniform, UniformTypeKVCacheSpecs)
blocks_needed += draft_uniform.max_memory_usage_pages(vllm_config)
draft_page = next(
iter(draft_uniform.kv_cache_specs.values())
).page_size_bytes
if draft_page != mla_page:
per_block += len(draft_group.layer_names) * draft_page
return blocks_needed * per_block
"""
EDITS: list[tuple[str, str, str]] = [
(
"groups: partition drafter SlidingWindowSpec layers out",
EDIT_PARTITION_ANCHOR,
EDIT_PARTITION_NEW,
),
(
"groups: build + append drafter group (exact-fit / standalone)",
EDIT_GROUPS_RETURN_ANCHOR,
EDIT_GROUPS_RETURN_NEW,
),
(
"layout: return-type annotation gains draft_group",
EDIT_LAYOUT_ANNOT_ANCHOR,
EDIT_LAYOUT_ANNOT_NEW,
),
(
"layout: docstring Returns gains draft_group",
EDIT_LAYOUT_DOC_ANCHOR,
EDIT_LAYOUT_DOC_NEW,
),
(
"layout: detect drafter SWA uniform group",
EDIT_LAYOUT_DETECT_ANCHOR,
EDIT_LAYOUT_DETECT_NEW,
),
(
"layout: validate drafter (uniform page, never padded)",
EDIT_LAYOUT_VALIDATE_ANCHOR,
EDIT_LAYOUT_VALIDATE_NEW,
),
(
"layout: return draft_group (9th element)",
EDIT_LAYOUT_RETURN_ANCHOR,
EDIT_LAYOUT_RETURN_NEW,
),
(
"_pool_bytes_per_block: standalone drafter bytes",
EDIT_POOL_BYTES_ANCHOR,
EDIT_POOL_BYTES_NEW,
),
(
"config: destructure + drafter mode",
EDIT_CONFIG_DESTRUCTURE_ANCHOR,
EDIT_CONFIG_DESTRUCTURE_NEW,
),
(
"config: per-block cost includes standalone drafter",
EDIT_CONFIG_PER_BLOCK_ANCHOR,
EDIT_CONFIG_PER_BLOCK_NEW,
),
(
"config: exact-fit drafter layer i co-owns MLA tensor i",
EDIT_CONFIG_SHARED_BY_ANCHOR,
EDIT_CONFIG_SHARED_BY_NEW,
),
(
"config: standalone drafter tensors",
EDIT_CONFIG_DRAFT_TENSORS_ANCHOR,
EDIT_CONFIG_DRAFT_TENSORS_NEW,
),
(
"max-mem: destructure gains draft_group",
EDIT_MAXMEM_DESTRUCTURE_ANCHOR,
EDIT_MAXMEM_DESTRUCTURE_NEW,
),
(
"max-mem: charge drafter block-id demand + standalone bytes",
EDIT_MAXMEM_BLOCKS_ANCHOR,
EDIT_MAXMEM_BLOCKS_NEW,
),
]
def patch_file(path: str, dry_run: bool = False) -> int:
with open(path, "r", encoding="utf-8") as f:
text = f.read()
if MARKER in text:
print(
f"[patch_glm5_drafter_group] {path}: already patched "
f"({MARKER} marker found); no-op."
)
return 0
# Sanity: the file we expect (guards against pointing at the wrong tree).
for required in (
"def _get_kv_cache_groups_glm5_next",
"def _glm5_next_tensor_layout",
"def _pool_bytes_per_block",
"SlidingWindowSpec",
"UniformTypeKVCacheSpecs",
):
assert required in text, (
f"ANCHOR PRECHECK FAILED: {required!r} not found in {path} -- "
"is this really vllm/v1/core/kv_cache_utils.py?"
)
applied = []
for name, anchor, replacement in EDITS:
n = text.count(anchor)
assert n == 1, (
f"ANCHOR FAILED for edit [{name}]: expected exactly 1 occurrence, "
f"found {n}. The upstream file has drifted -- re-derive the anchor "
f"before building.\n--- anchor ---\n{anchor}\n--------------"
)
text = text.replace(anchor, replacement, 1)
applied.append(name)
# The patched source must still be valid Python.
try:
ast.parse(text, filename=path)
except SyntaxError as e:
raise AssertionError(f"POST-EDIT ast.parse FAILED for {path}: {e}") from e
if dry_run:
print(f"[patch_glm5_drafter_group] DRY RUN -- {path} not written.")
else:
with open(path, "w", encoding="utf-8") as f:
f.write(text)
print(f"[patch_glm5_drafter_group] {path}: {len(applied)} edits applied:")
for name in applied:
print(f" - {name}")
print("[patch_glm5_drafter_group] ast.parse OK.")
return 0
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__.splitlines()[1])
ap.add_argument("--kv-file", default=DEFAULT_KV_FILE)
ap.add_argument(
"--dry-run",
action="store_true",
help="validate anchors + parse, write nothing",
)
args = ap.parse_args()
return patch_file(args.kv_file, dry_run=args.dry_run)
if __name__ == "__main__":
sys.exit(main())
Now build — from the working-directory root, because the build context is this directory: the COPY lines expect files/overlay-dflash2/ at its top, with the Dockerfile at tony-lane/Dockerfile (accepts arm64 on the Spark; do not add --platform):
docker build --network=host -t glm53-dflash2-tony:v1 -f tony-lane/Dockerfile .
Both build gates must print: registry: DFlash2 OK, dflash2 modules import OK.
Optional — build the InstantTensor fast-loader variant too (this is the :v1-it image the day-to-day launch uses):
docker build --network=host --build-arg INSTANTTENSOR=1 -t glm53-dflash2-tony:v1-it -f tony-lane/Dockerfile .
Step 6 — ship the image to the worker
docker save glm53-dflash2-tony:v1 | ssh user@worker-node docker load
Step 7 — create the launcher (both nodes)
One script, rank as argument. Copy it to both machines as tony-lane/launch-tp2.sh (inside the working directory) and make it executable (chmod +x tony-lane/launch-tp2.sh). Set WORKER_SSH to your real worker target before running. The CX7 point-to-point addresses follow the upstream recipe's convention (head 10.0.0.1, worker 10.0.0.2); override MASTER_ADDR and the NIC pins (CX_NIC, CX_HCA) if your cabling differs:
#!/usr/bin/env bash
# ============================================================================
# launch-tp2.sh — tonyd2wild-style mp-executor launcher for GLM-5.3-Flash
# DFlash2 on our 2x DGX Spark kit (head 10.0.0.1 / worker 10.0.0.2).
#
# Mirrors tonyd2wild's launch-glm53-vllm-tp2.sh structure but drops Ray
# entirely (`--distributed-executor-backend mp`, worker-first). Serves the
# MTP-stripped target copy with the DFlash2 drafter.
#
# ORDER MATTERS: run rank 1 (worker) FIRST, wait ~20 s, then rank 0 (head).
# ssh user@worker-node 'cd ~/glm5.3-flash-spark && ./tony-lane/launch-tp2.sh 1'
# ./tony-lane/launch-tp2.sh 0
#
# Usage: ./tony-lane/launch-tp2.sh [stop | <0|1> [stop]]
# stop stop BOTH ranks (head here + worker over ssh)
# <0|1> launch that rank (0=head, 1=worker)
# <0|1> stop remove that rank's container only
# ============================================================================
set -euo pipefail
RANK="${1:?usage: launch-tp2.sh [stop | <0|1> [stop]]}"
CMD="${2:-start}"
# "stop" (no rank) = tear down BOTH ranks: head here, worker via ssh.
# Head first — once the driver/API dies the worker rank is doomed anyway.
if [ "$RANK" = "stop" ]; then
HEAD_NAME="${HEAD_NAME:-glm53-tony-rank0}"
WORKER_NAME="${WORKER_NAME:-glm53-tony-rank1}"
WORKER_SSH="${WORKER_SSH:-user@worker-node}"
docker rm -f "$HEAD_NAME" 2>/dev/null || true
echo "stopped $HEAD_NAME (head)"
if ssh "$WORKER_SSH" "docker rm -f '$WORKER_NAME' 2>/dev/null || true"; then
echo "stopped $WORKER_NAME on $WORKER_SSH (worker)"
else
echo "WARN: ssh $WORKER_SSH failed — stop the worker manually:" >&2
echo " ssh $WORKER_SSH 'docker rm -f $WORKER_NAME'" >&2
fi
exit 0
fi
[[ "$RANK" == "0" || "$RANK" == "1" ]] || { echo "rank must be 0 or 1" >&2; exit 2; }
IMAGE="${IMAGE:-glm53-dflash2-tony:v1}"
NAME="${NAME:-glm53-tony-rank$RANK}"
MASTER_ADDR="${MASTER_ADDR:-10.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29521}"
API_PORT="${API_PORT:-8888}"
SERVED_NAME="${SERVED_NAME:-LibertAIDAI/GLM-5.3-Flash-NVFP4}"
HF_HUB="${HF_HOME:-$HOME/.cache/huggingface}/hub"
# Target model HF cache dir name. Swap checkpoints by pointing this at
# another cache dir — any glm5_next modelopt-NVFP4 checkpoint with the
# standard layout works.
MODEL_CACHE_NAME="${MODEL_CACHE_NAME:-models--LibertAIDAI--GLM-5.3-Flash-NVFP4-no-MTP}"
DRAFT_CACHE_NAME="${DRAFT_CACHE_NAME:-models--incoai--GLM-5.3-Flash-DFlash2}"
if [ "$RANK" = "0" ]; then
HOST_IP=10.0.0.1
HEADLESS=""
else
HOST_IP=10.0.0.2
HEADLESS="--headless"
fi
# CX7 fabric pins. Both nodes use the SAME CX7 port here (f1 <> f1 w/ cable),
# so NIC/HCA names are identical on both ranks. If your cabling differs (e.g.
# the Mia recipe's crossed f1-next-to-f0), change per rank below and/or set
# the WORKER_* envs — never guess: the UP check below fails fast.
CX_NIC="${CX_NIC:-enp1s0f1np1}"
CX_HCA="${CX_HCA:-rocep1s0f1}"
if [ "$CMD" = "stop" ]; then
docker rm -f "$NAME" 2>/dev/null || true
echo "stopped $NAME"
exit 0
fi
# fail fast: the pinned NIC must exist AND have carrier on THIS node.
# (after the stop path so a down/renamed NIC can never block teardown)
if ! ip -br link show "$CX_NIC" 2>/dev/null | grep -q UP; then
echo "FATAL: $CX_NIC is not UP on $(hostname) — wrong port/cable or wrong pin." >&2
echo " active fabric candidates:" >&2
ip -br link show | grep -E "enp1s0f[01]" || true
echo " (both ranks pinned to f1/f1; override with CX_NIC/CX_HCA if cabling changes)" >&2
exit 1
fi
# resolve snapshot RELATIVE paths (same under host HF_HUB and container mount)
MODEL_HASH="$(cat "$HF_HUB/$MODEL_CACHE_NAME/refs/main")"
DRAFT_HASH="$(cat "$HF_HUB/$DRAFT_CACHE_NAME/refs/main")"
MODEL_REL="$MODEL_CACHE_NAME/snapshots/$MODEL_HASH"
DRAFT_REL="$DRAFT_CACHE_NAME/snapshots/$DRAFT_HASH"
# host paths — existence check only
MODEL_SNAP_HOST="$HF_HUB/$MODEL_REL"
DRAFT_SNAP_HOST="$HF_HUB/$DRAFT_REL"
for d in "$MODEL_SNAP_HOST" "$DRAFT_SNAP_HOST"; do
[ -d "$d" ] || { echo "missing $d" >&2; exit 1; }
done
# container paths — what vLLM / huggingface_hub must see inside (the hub is
# mounted at /root/.cache/huggingface/hub; host paths do not exist there)
MODEL_SNAP="/root/.cache/huggingface/hub/$MODEL_REL"
DRAFT_SNAP="/root/.cache/huggingface/hub/$DRAFT_REL"
docker image inspect "$IMAGE" >/dev/null 2>&1 \
|| { echo "image $IMAGE missing — build via:" >&2
echo " docker build --network=host -t $IMAGE -f tony-lane/Dockerfile ." >&2; exit 1; }
docker rm -f "$NAME" 2>/dev/null || true
# ---- NCCL / fabric environment (per-node CX7 pins) ------------------------
NCCL_ENV=(
-e NCCL_NET=IB -e NCCL_IB_DISABLE=0
-e "NCCL_IB_HCA=$CX_HCA" -e NCCL_IB_GID_INDEX=3 -e NCCL_IB_ROCE_VERSION_NUM=2
-e "NCCL_SOCKET_IFNAME=$CX_NIC" -e "GLOO_SOCKET_IFNAME=$CX_NIC"
-e NCCL_IB_ADDR_FAMILY=AF_INET -e NCCL_IB_ADDR_RANGE=10.0.0.0/24
-e NCCL_NVLS_ENABLE=0 -e NCCL_CROSS_NIC=0 -e NCCL_IB_MERGE_NICS=0
-e NCCL_CUMEM_ENABLE=0 -e NCCL_IGNORE_CPU_AFFINITY=1 -e NCCL_DEBUG=WARN
-e TORCH_NCCL_ASYNC_ERROR_HANDLING=1 -e NCCL_IB_TIMEOUT=22
)
LOAD_FORMAT="${LOAD_FORMAT:-}" # empty = default safetensors; "instanttensor" = opt-in
if [ -n "$LOAD_FORMAT" ]; then
echo "NOTE: LOAD_FORMAT=$LOAD_FORMAT — confirmed working on this kit (2026-08-29, boot + serve)." >&2
echo " tonyd2wild's md reports 4/4 silent TP2 deaths on their stack; our same-layer nccl re-pin to 2.30.7 may be why we're fine." >&2
fi
LOAD_ARGS=()
[ -n "$LOAD_FORMAT" ] && LOAD_ARGS=(--load-format "$LOAD_FORMAT")
# ---- KV / scheduler knobs (env-overridable) --------------------------------
# MAX_MODEL_LEN: 500k-class context (2 full sessions in the pool). The DFlash2
# drafter is 1M-native, so no drafter change is needed. KV_CACHE_MEMORY
# set-but-EMPTY drops the pin: vLLM then profiles KV from GPU_MEM_UTIL,
# which is conservative on this box (~5 GB residual at 0.89). Use the SAME
# values on BOTH ranks. Previous proven defaults: max-len 262144, gmu 0.90,
# pin 10247694336 (989,727 tokens) — set via env to walk back.
MAX_NUM_SEQS="${MAX_NUM_SEQS:-6}"
MAX_MODEL_LEN="${MAX_MODEL_LEN:-524288}"
GPU_MEM_UTIL="${GPU_MEM_UTIL:-0.90}"
# ~10,353 B/token (fp8_e4m3, TP=2): 10247694336 -> 989,727 tokens.
# Holds one full 524,288-token request with room to spare, or two
# 500k-class sessions (~494k each). Proven to boot inside gmu 0.90
# headroom on this kit (2026-08-29).
KV_CACHE_MEMORY="${KV_CACHE_MEMORY-10247694336}"
KV_ARGS=()
[ -n "$KV_CACHE_MEMORY" ] && KV_ARGS=(--kv-cache-memory "$KV_CACHE_MEMORY")
# CPU threads: 'nobind' stops vLLM's OMP thread-pinning from limiting engine
# aux work (tokenizer, scheduler, MoE weight repack) to a tiny core subset —
# the known Spark fix for "vLLM only uses 2 cores". Same value on BOTH ranks.
VLLM_CPU_OMP_THREADS_BIND="${VLLM_CPU_OMP_THREADS_BIND:-nobind}"
# vLLM clamps torch CPU threads to 1 for serving ("Reducing Torch threads
# from 20 to 1 … Set OMP_NUM_THREADS to override") — give the tokenizer and
# scheduler a real budget without oversubscribing the 20-core GB10.
OMP_NUM_THREADS="${OMP_NUM_THREADS:-8}"
# Batching budget: spec-decode clamps max_num_scheduled_tokens to 2048 by
# default (vLLM itself warns this is suboptimal with 7-token draft blocks).
# Raise it so batched prefill/decode steps scale with concurrency.
MAX_NUM_BATCHED_TOKENS="${MAX_NUM_BATCHED_TOKENS:-8192}"
# ---- performance knobs (env-overridable) -----------------------------------
# Spec-decode block length. Measured per-position acceptance on real traffic
# (0.70/0.49/0.37/0.29/...) makes positions 5-7 of the block nearly
# free-riding: n=4 keeps the accepted tokens while cutting the verify batch
# from 8 to 5 positions. Measured +21-23% tok/s vs n=7 (2026-08-31 sweep).
# Set NUM_SPEC_TOKENS=7 to match the upstream benchmark conditions.
NUM_SPEC_TOKENS="${NUM_SPEC_TOKENS:-4}"
# MoE kernel family. MARLIN is the only backend that works on this build —
# FLASHINFER_TRTLLM / FLASHINFER_CUTLASS / VLLM_CUTLASS / FLASHINFER_CUTEDSL
# all fail to boot (2026-08-31). Keep marlin unless a new image lands.
MOE_BACKEND="${MOE_BACKEND:-marlin}"
# 0 = CUDA graphs via Breakable mode (captured cleanly incl. the DFlash2
# drafter, 1.66-1.76 GiB, +2-6% steps/s). 1 = --enforce-eager fallback.
EAGER="${EAGER:-0}"
EAGER_ARGS=()
[ "$EAGER" = "1" ] && EAGER_ARGS=(--enforce-eager)
docker run -d --name "$NAME" --restart no \
--gpus all --network host --ipc host --shm-size 16g \
--ulimit memlock=-1:-1 --cap-add IPC_LOCK \
--device /dev/infiniband:/dev/infiniband \
-v "$HF_HUB:/root/.cache/huggingface/hub:ro" \
-e "VLLM_HOST_IP=$HOST_IP" \
-e "VLLM_CPU_OMP_THREADS_BIND=$VLLM_CPU_OMP_THREADS_BIND" \
-e "OMP_NUM_THREADS=$OMP_NUM_THREADS" \
-e HF_HUB_OFFLINE=1 -e TRANSFORMERS_OFFLINE=1 \
-e VLLM_ENGINE_READY_TIMEOUT_S=3600 \
-e TORCH_CUDA_ARCH_LIST=12.1a -e FLASHINFER_CUDA_ARCH_LIST=12.1a \
-e FLASHINFER_DISABLE_VERSION_CHECK=1 \
-e PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True \
"${NCCL_ENV[@]}" \
"$IMAGE" \
"$MODEL_SNAP" \
--served-model-name "$SERVED_NAME" \
--host 0.0.0.0 --port "$API_PORT" \
--trust-remote-code --tensor-parallel-size 2 \
"${LOAD_ARGS[@]}" \
--distributed-executor-backend mp \
--nnodes 2 --node-rank "$RANK" \
--master-addr "$MASTER_ADDR" --master-port "$MASTER_PORT" \
$HEADLESS \
--max-model-len "$MAX_MODEL_LEN" --max-num-seqs "$MAX_NUM_SEQS" --block-size 2304 \
--max-num-batched-tokens "$MAX_NUM_BATCHED_TOKENS" \
--moe-backend "$MOE_BACKEND" "${EAGER_ARGS[@]}" \
--kv-cache-dtype fp8_e4m3 "${KV_ARGS[@]}" \
--gpu-memory-utilization "$GPU_MEM_UTIL" \
--speculative-config "{\"method\":\"dflash\",\"model\":\"$DRAFT_SNAP\",\"num_speculative_tokens\":$NUM_SPEC_TOKENS}" \
--tool-call-parser glm47 --enable-auto-tool-choice \
--reasoning-parser glm45
echo "launched $NAME rank=$RANK host=$HOST_IP"
echo " model (container): $MODEL_SNAP"
echo " draft (container): $DRAFT_SNAP"
echo " kv: pin=${KV_CACHE_MEMORY:-<none — gmu-profiled>} gmu=$GPU_MEM_UTIL seqs=$MAX_NUM_SEQS mlen=$MAX_MODEL_LEN"
echo " perf: spec=$NUM_SPEC_TOKENS moe=$MOE_BACKEND eager=$EAGER omp=$OMP_NUM_THREADS"
sleep 3
docker ps --format '{{.Names}} {{.Status}}' | grep -q "$NAME" \
|| { echo "container exited; logs:"; docker logs "$NAME" 2>&1 | tail -30; exit 1; }
if [ "$RANK" = "0" ]; then
echo "waiting /health (VLLM_ENGINE_READY_TIMEOUT_S=3600) ..."
for i in $(seq 1 360); do
if curl -sf "http://127.0.0.1:$API_PORT/health" >/dev/null 2>&1; then
echo "HEALTHY after ~$((i * 10)) s"; exit 0
fi
sleep 10
done
echo "timed out waiting /health; logs follow:" >&2
docker logs "$NAME" 2>&1 | tail -80
exit 1
fi
echo "rank $RANK up (tail logs: docker logs -f $NAME)"
Step 8 — launch (worker first, then head)
Order matters: worker (rank 1) first, wait ~20 s, then head (rank 0).
IMAGE=glm53-dflash2-tony:v1-it LOAD_FORMAT=instanttensor ./tony-lane/launch-tp2.sh 1 # on the worker
IMAGE=glm53-dflash2-tony:v1-it LOAD_FORMAT=instanttensor ./tony-lane/launch-tp2.sh 0 # on the head, ~20 s later
That is the day-to-day launch (InstantTensor fast-loader image). Drop the leading IMAGE=…/LOAD_FORMAT=… assignments to run the plain safetensors :v1 image instead — the tuning is identical.
One more thing we do on both nodes: cap the GPU clocks.
sudo nvidia-smi -lgc 0,2000
Desk-side boxes — we trade a little peak boost for cool, quiet 24/7 operation. Every number in this post is measured at that 2 GHz cap.
Every knob the launcher takes — tuned values are the defaults, so an empty environment reproduces the numbers below:
| knob | default | what it does |
|---|---|---|
IMAGE |
glm53-dflash2-tony:v1 |
serving image; :v1-it = InstantTensor variant |
SERVED_NAME |
LibertAIDAI/GLM-5.3-Flash-NVFP4 |
model name the API reports |
MODEL_CACHE_NAME |
models--LibertAIDAI--GLM-5.3-Flash-NVFP4-no-MTP |
which weights to serve |
DRAFT_CACHE_NAME |
models--incoai--GLM-5.3-Flash-DFlash2 |
drafter weights |
LOAD_FORMAT |
(empty = safetensors) | instanttensor = fast-loader opt-in |
EAGER |
0 |
1 = --enforce-eager (no CUDA graphs) |
NUM_SPEC_TOKENS |
4 |
draft block length (upstream default 7) |
GPU_MEM_UTIL |
0.90 |
vLLM --gpu-memory-utilization |
KV_CACHE_MEMORY |
10247694336 |
the 989,727-token KV pin; set-but-empty drops it |
OMP_NUM_THREADS |
8 |
tokenizer/detokenizer threads |
API_PORT |
8888 |
OpenAI-compatible API port |
MAX_MODEL_LEN |
524288 |
max context per request |
MAX_NUM_SEQS |
6 |
max concurrent sequences |
MOE_BACKEND |
marlin |
only NVFP4 MoE backend that boots |
WORKER_SSH |
user@worker-node |
where the head pushes rank 1 |
MASTER_ADDR |
10.0.0.1 |
head's CX7 address |
CX_NIC / CX_HCA |
enp1s0f1np1 / rocep1s0f1 |
CX7 NIC + RDMA device |
The launcher also echoes the resolved config at boot (kv: pin=… gmu=… seqs=… mlen=…, perf: spec=… moe=… eager=…), so you always see what you actually launched.
320B MoE init takes a while. The launcher echoes the resolved snapshots and the perf config; in the logs, confirm these lines:
Detected ModelOpt NVFP4 checkpoint (quant_algo=NVFP4)
Using 'MARLIN' NvFp4 MoE backend
Using Eagle3 auxiliary layers from config: (6, 15, 25, 34, 43)
GPU KV cache size: 989,727 tokens
Breakable CUDA graph enabled
Graph capturing finished in ~30 secs, took ~1.3 GiB
Step 9 — verify: first request
curl -s http://localhost:8888/v1/chat/completions \
-H 'Content-Type: application/json' \
-d '{"model": "LibertAIDAI/GLM-5.3-Flash-NVFP4", "messages": [{"role": "user", "content": "hello!"}]}'
A JSON completion comes back and you are serving.
The numbers (all measured on this config)
All measured at the 2 GHz clock cap from Step 8 — nothing here is a boost-clock number.
KV geometry, exact from the pinned boot:
| GPU memory utilization | KV pin | KV pool | KV tokens | vs one full 524k request |
|---|---|---|---|---|
| 0.85 | 5.5 GiB (upstream recipe) | 5.5 GiB | ~570k | 1.1× |
| 0.89 | none — self-profiled | ~4.6 GiB | 481,488 | 0.92× |
| 0.90 | 10,247,694,336 B | 9.54 GiB | 989,727 | 1.89× |
Without the pin, vLLM's profiler is conservative (~5 GB residual at 0.89) — the explicit pin is what unlocked the pool. We serve --max-model-len 524288: half a million tokens per request, 1.89× covered by the pool.
Throughput tuning, all live-measured on this kit:
| config | steps/s | output tok/s (live avg) |
|---|---|---|
| n=7, eager (upstream default) | 7.5 | ~19.0 |
| n=4, eager | 9.1 | ~23.4 |
| n=4 + CUDA graphs | ~9.3 | ~25 avg, 26–29 peak |
Why n=4 beats the upstream default of 7: vLLM logs per-position draft acceptance, and ours reads 0.70 / 0.46 / 0.29 / 0.22 for the first four positions with the tail collapsing. Since output tok/s = steps/s × acceptance length and each verify step pays for the whole block, positions 5–7 were costing half the verify batch for ~0.1 tokens per step. Truncating to 4 kept the tokens and cut the batch — +21–23%. The Breakable CUDA graphs (which capture the DFlash2 drafter too, 1.3–1.8 GiB) add 2–6%; decode here is memory-bandwidth-bound, so launch-overhead elimination is the small lever.
And the ceiling: asked to count from 1 to 1000 — next token near-deterministic — the drafter hit per-position acceptance 1.000 / 0.99 / 0.97 / 0.94, acceptance length 4.89 of 5.0, and 46.5–47.8 tok/s. That is tonyd2wild's 46.9 reference, reproduced at ceiling: the pipeline is mechanically perfect, and tok/s differences between stacks and posts are workload predictability, not quality. Quote steps/s (what the kit delivers) and acceptance length (what your content allows) separately, and vendor benchmarks stop being able to fool you.
We serve at temperature 1.0 / top-p 0.95 — the model's own generation defaults. Greedy decoding raises draft acceptance; benchmarking at temp 0 and serving at temp 1 flatters the published number.
The quantization lesson: NVFP4 is a format, not a speed
We swap-tested a third-party abliterated NVFP4 checkpoint. On disk smaller; in RAM 90.77 GiB/rank vs 90.67 — no headroom gained. It serves through vLLM's compressed-tensors NVFP4 path, which this build forces into W4A4 (activations quantized to FP4 too). GB10 has native FP4 tensor cores, but the upstream kernels haven't landed, so W4A4 runs as software emulation: measured 5.9 steps/s vs 7.5 — a ~25% per-step penalty at identical acceptance, and a config-only conversion to weight-only is silently ignored for this format family. The checkpoint's quant format family determines which kernels you get. Ours is modelopt-format and that is the speed decision. When SM121 FP4 kernels land upstream, W4A4 checkpoints stop paying the emulation tax and this flips.
Caveats
The DFlash2 drafter weights are CC BY-NC-ND 4.0 — non-commercial. Fine for a desk, not a product.
- Long-run soak and a formal fixed-prompt throughput suite are still open. Every number above is from live-traffic windows or targeted tests on our boxes; direction is reproducible, decimals carry noise.
- MARLIN is the only NVFP4 MoE backend that boots on this build (TRT-LLM and CUTLASS flavors fail). First thing to re-test on a newer image.
- Spec block n=2–3 untested; measured acceptance says the optimum is near 4.
Try it
curl -s http://localhost:8888/v1/chat/completions \
-H 'Content-Type: application/json' \
-d '{"model": "LibertAIDAI/GLM-5.3-Flash-NVFP4", "messages": [{"role": "user", "content": "hello!"}]}'
Tool calling and reasoning-parser flags are on; images and video work through standard multimodal content parts. Two ~$1k-class boxes, one CX7 cable, no cloud, no per-token bill — and a 320B model that can't be deprecated out from under you. Credit one last time: tonyd2wild's recipe and repos made this a weekend build instead of a research project.