Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion docker/sglang_disagg_inference.ubuntu.amd.Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,11 @@ RUN pip install --upgrade sglang-router

WORKDIR /sgl-workspace/mori

ARG MORI_COMMIT="158c7e8335a0b19b3f1f422ff134d7869252135e"
# MoRI >= #363 (guard dispatch kernels vs out-of-range expert id) and #505 (AsyncLL slot
# double-alloc when top-k does not divide warpSize) are REQUIRED for DeepSeek-V4-Flash decode
# CUDA-graph capture (topk6 -> 6 does not divide warpSize 64). The older 158c7e83 pin (2026-06-08)
# predates both and crashes at capture (mori low_latency_async.cpp:360 pe-out-of-range).
ARG MORI_COMMIT="7c51d18fda59457cc9238ed262bd93c8cad906c9"
# Set INSTALL_MORI=1 to build/install MoRI at MORI_COMMIT; any other value skips it.
ARG INSTALL_MORI=1

Expand Down
37 changes: 37 additions & 0 deletions scripts/sglang_disagg/README.MD
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ MoE Models
- DeepSeek-V3 (https://huggingface.co/deepseek-ai/DeepSeek-V3)
- DeepSeek-R1 (https://huggingface.co/deepseek-ai/DeepSeek-R1)
- Mixtral-8x7B-v0.1 (https://huggingface.co/mistralai/Mixtral-8x7B-v0.1)
- DeepSeek-V4-Flash-FP8 (https://huggingface.co/sgl-project/DeepSeek-V4-Flash-FP8) — MoRI-EP wide-EP; use the FP8-E4M3 checkpoint, DP-attention mandatory (num_key_value_heads=1). EP8 1P1D and EP16 2P2D validated on MI308X/gfx942 + Broadcom Thor2 (set USE_CX7_NICS=0).

This repository contains scripts and documentation to launch PD Disaggregation for above models. You will find setup instructions, node assignment details and benchmarking commands.

Expand Down Expand Up @@ -129,3 +130,39 @@ curl -X POST http://127.0.0.1:2322/generate \
For larger models, such as DeepSeekV3 and Llama-3.1-405B-Instruct-FP8-KV and higher concurrency(512+), errors with below signature is observed:<br>
_'<TransferEncodingError: 400, message:\n Not enough data to satisfy transfer length header.\n\nThe above exception was the direct cause of the following exception:\n\nTraceback (most recent call last):\n '_<br>
This leads to dropping requests and lower throughput.This issue is being discussed on the SGLang forums.

## DeepSeek-V4-Flash-FP8 (MoRI-EP wide-EP, EP8 + EP16)

DeepSeek-V4-Flash is served through the MoRI-EP wide-EP path (`DP_MODE=1`). Notes specific to this model:

- **Checkpoint:** use the **FP8-E4M3 block-quant** weights (e.g. `sgl-project/DeepSeek-V4-Flash-FP8`). The model uses the dedicated `dsv4` attention backend (not `aiter`), and `num_key_value_heads=1` so attention cannot be tensor-sharded — **DP-attention is mandatory** for any EP run.
- **MoRI-EP mode:** prefill uses `--deepep-mode normal` (high-throughput), decode uses `--deepep-mode low_latency` (matches the MoRI HT-prefill / LL-decode split). Set in `models.yaml`.
- **Prefill is eager** (`--disable-cuda-graph`); decode uses cudagraph capture.

### Topologies (via `models.json`)

| Entry | xP | yD | Nodes |
|-------|----|----|-------|
| `pyt_sglang_disagg_mori_io_dsv4-flash_ep8_1p1d` | 1 | 1 | 2 (EP8, 1 prefill + 1 decode) |
| `pyt_sglang_disagg_mori_io_dsv4-flash_ep16_2p2d` | 2 | 2 | 4 (EP16, 2 prefill + 2 decode) |

### Docker image

A prebuilt image (SGLang + MoRI on ROCm 7.2, gfx942) is available:

```bash
docker pull rocmshared/sglang-disagg-dsv4:mori-mi308-pr
```

Or build via `docker/sglang_disagg_inference.ubuntu.amd.Dockerfile` (base `lmsysorg/sglang:v0.5.15-rocm720-mi30x`, MoRI from source).

### Broadcom Thor2 (bnxt) / non-CX7 clusters

Set `USE_CX7_NICS=0`. On bnxt the image's baked `libbnxt_re` may be ABI-stale versus the host kernel; the launcher exposes the host's ABI-correct `libbnxt_re` and runs `ldconfig` in-container automatically for this path. RDMA NIC selection (`IB_DEVICES` / `MORI_RDMA_DEVICES`) should be set to the node's RoCE devices.

### Validation (MI308X/gfx942 + Broadcom Thor2)

Both topologies validated end-to-end (serve via `sglang_router`, greedy `/v1/completions`):

- **Correctness:** needle-in-a-haystack **18/18** (context 1K→200K, needle at 10/50/90% depth) on **both** EP8 1P1D and EP16 2P2D. Factual length-sweep coherent to 200K.
- **Perf** (functional, eager decode; all requests succeeded): 8K/1K and 16K/1K input/output at concurrency 16 and 32.
184 changes: 184 additions & 0 deletions scripts/sglang_disagg/dsv4_flash/PERF_REPORT.html

Large diffs are not rendered by default.

20 changes: 20 additions & 0 deletions scripts/sglang_disagg/dsv4_flash/niah.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
import json, urllib.request
import os
URL=os.environ.get("ENDPOINT","http://127.0.0.1:2322")+"/v1/completions"; M="/models/DeepSeek-V4-Flash-FP8-E4M3"
F="The quick brown fox jumps over the lazy dog. " # ~9 tokens
def ask(p, mx=5):
data=json.dumps({"model":M,"prompt":p,"max_tokens":mx,"temperature":0}).encode()
r=urllib.request.Request(URL,data=data,headers={"Content-Type":"application/json"})
return json.load(urllib.request.urlopen(r,timeout=300))["choices"][0]["text"]
# token targets -> filler repeats (~9 tok each)
lengths={"1K":110,"4K":440,"16K":1780,"32K":3560,"100K":11100,"200K":22200}
pass_=tot=0
for name,reps in lengths.items():
for d in [0.1,0.5,0.9]:
pre=int(reps*d)
p=F*pre+"Marie was born in the city of Paris. "+F*(reps-pre)+"Marie was born in the city of"
try: t=ask(p); ok="paris" in t.lower()
except Exception as e: t=f"ERR {str(e)[:40]}"; ok=False
tot+=1; pass_+=ok
print(f"{name} d={int(d*100)}% -> {t[:28]!r} [{'PASS' if ok else 'FAIL'}]",flush=True)
print(f"NIAH TOTAL: {pass_}/{tot}",flush=True)
29 changes: 29 additions & 0 deletions scripts/sglang_disagg/dsv4_flash/perf.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
import json,urllib.request,time,threading
import os
URL=os.environ.get("ENDPOINT","http://127.0.0.1:2322")+"/v1/completions"; M="/models/DeepSeek-V4-Flash-FP8-E4M3"
F="The quick brown fox jumps over the lazy dog. "
def make(isl):
reps=isl//9; return (F*reps)[:isl*5] # approx isl tokens
def one(prompt,osl,res,i):
t0=time.time()
data=json.dumps({"model":M,"prompt":prompt,"max_tokens":osl,"temperature":0,"stream":False}).encode()
r=urllib.request.Request(URL,data=data,headers={"Content-Type":"application/json"})
try:
d=json.load(urllib.request.urlopen(r,timeout=300)); dt=time.time()-t0
ct=d["usage"]["completion_tokens"]; res[i]=(dt,ct)
except Exception as e: res[i]=(None,str(e)[:40])
def run(isl,osl,con):
prompt=make(isl)
res=[None]*con; ths=[]
t0=time.time()
for i in range(con):
th=threading.Thread(target=one,args=(prompt,osl,res,i)); th.start(); ths.append(th)
for th in ths: th.join()
wall=time.time()-t0
ok=[r for r in res if r and r[0]]
if not ok: print(f"ISL={isl} OSL={osl} CON={con}: ALL FAILED {res[0]}"); return
lat=sum(r[0] for r in ok)/len(ok); toks=sum(r[1] for r in ok)
print(f"ISL={isl} OSL={osl} CON={con}: {len(ok)}/{con} ok | mean_latency={lat:.1f}s | out_tok={toks} | wall={wall:.1f}s | out_tps={toks/wall:.0f} tok/s",flush=True)
for isl in [8192,16384]:
for con in [16,32]:
run(isl,1024,con)
69 changes: 69 additions & 0 deletions scripts/sglang_disagg/dsv4_flash/run_dsv4_nonslurm.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
#!/bin/bash
# Non-SLURM driver for sglang_disagg DSV4-Flash EP8/EP16 on a bnxt/Thor2 (Broadcom) cluster.
# Mirrors the docker-run + env block of run_xPyD_models.slurm, but launches per-node
# over SSH (no SLURM). Runs the framework's own sglang_disagg_mori_io_ep.sh inside.
#
# Usage:
# EP8 1P1D: xP=1 yD=1 PREFILL_NODES=nodeA DECODE_NODES=nodeB bash run_dsv4_nonslurm.sh
# EP16 2P2D: xP=2 yD=2 PREFILL_NODES="nodeA nodeB" DECODE_NODES="nodeC nodeD" bash run_dsv4_nonslurm.sh
set -u

IMG="${IMG:-localhost/mad_dsv4_disagg:pr}"
MODEL_NAME="${MODEL_NAME:-DeepSeek-V4-Flash-FP8}"
MODEL_PATH="${MODEL_PATH:-/models/DeepSeek-V4-Flash-FP8-E4M3}"
COOKBOOK_HOST="${COOKBOOK_HOST:-/root/sglang_disagg_cookbook}"
COOKBOOK_IN="/sgl-cookbook"
xP="${xP:?set xP}"; yD="${yD:?set yD}"
DP_MODE="${DP_MODE:-1}"
PREFILL_NODES="${PREFILL_NODES:?space-separated prefill mgmt hostnames}"
DECODE_NODES="${DECODE_NODES:?space-separated decode mgmt hostnames}"
JIT_CACHE="${JIT_CACHE:-/root/.mad_jit_cache/$(echo "$IMG"|tr '/:' '__')}"

ALL_NODES=($PREFILL_NODES $DECODE_NODES)
# Fabric .200-subnet IP per node (MoRI/disagg bootstrap rides the RDMA fabric).
fabip(){ ssh -n "$1" "ip -br -4 addr show | awk '\$3 ~ /^192\.168\.200\./{print \$3}' | cut -d/ -f1 | head -1"; }
IPADDRS=""
for n in "${ALL_NODES[@]}"; do IPADDRS+="${IPADDRS:+,}$(fabip "$n")"; done
MASTER_ADDR=$(echo "$IPADDRS" | cut -d, -f1)
echo "IPADDRS(fabric)=$IPADDRS MASTER=$MASTER_ADDR xP=$xP yD=$yD DP_MODE=$DP_MODE"

launch_node(){
local node="$1" rank="$2"
local mgmtif; mgmtif=$(ssh -n "$node" "ip route | awk '/^default/{for(i=1;i<=NF;i++)if(\$i==\"dev\")print \$(i+1)}' | head -1")
local hostlib; hostlib=$(ssh -n "$node" "ls /usr/local/lib/libbnxt_re-rdmav34.so 2>/dev/null | head -1")
local bnxtmnt=""; [[ -n "$hostlib" ]] && bnxtmnt="-v ${hostlib}:${hostlib}:ro"
local DATA_MNTS=""; for d in /mnt/md0 /mnt/nvme1 /mnt/nvme2 /mnt/nvme3; do ssh -n "$node" "[ -d $d ]" 2>/dev/null && DATA_MNTS+=" -v $d:$d"; done
# bnxt/Thor2 RDMA devices, sorted by fabric subnet octet (rank i -> same rail both ends)
local RAILDEVS; RAILDEVS=$(ssh -n "$node" 'for dv in /sys/class/infiniband/bnxt_re_bond*; do d=$(basename $dv); nd=$(ls $dv/device/net 2>/dev/null|head -1); m=$(basename $(readlink /sys/class/net/$nd/master 2>/dev/null) 2>/dev/null); ip=$(ip -br -4 addr show $m 2>/dev/null|awk "{print \$3}"|cut -d/ -f1); echo "$(echo $ip|cut -d. -f3) $d"; done | sort -n | awk "{print \$2}" | paste -sd,')
# persistent JIT cache + orphan-lock sweep (no live compiler)
ssh -n "$node" "mkdir -p ${JIT_CACHE}/mori ${JIT_CACHE}/aiter; if ! pgrep -x cc1plus>/dev/null && ! pgrep -x hipcc>/dev/null; then find ${JIT_CACHE} \\( -name 'lock_module_*' -o -name '*.lock' -o -name 'lock' \\) -delete 2>/dev/null; fi"
ssh -n "$node" "docker rm -f dsv4_${MODEL_NAME}_r${rank} >/dev/null 2>&1
docker run -d --name dsv4_${MODEL_NAME}_r${rank} \
--network host --ipc host --privileged \
--device /dev/kfd --device /dev/dri --device /dev/infiniband \
--group-add video --cap-add SYS_PTRACE --cap-add IPC_LOCK --security-opt seccomp=unconfined \
--ulimit memlock=-1 --ulimit nofile=1048576 --ulimit nproc=-1 \
-v /models:/models ${DATA_MNTS} \
-v /sys/kernel/config:/sys/kernel/config -v /sys/kernel/debug:/sys/kernel/debug \
-v /etc/libibverbs.d:/etc/libibverbs.d:ro ${bnxtmnt} \
-v ${JIT_CACHE}:/jit_cache \
-v ${COOKBOOK_HOST}:${COOKBOOK_IN} \
-e MORI_JIT_CACHE_DIR=/jit_cache/mori -e AITER_JIT_DIR=/jit_cache/aiter \
-e OPENBLAS_NUM_THREADS=4 -e OMP_NUM_THREADS=4 -e MKL_NUM_THREADS=4 -e NUMEXPR_NUM_THREADS=4 -e GOTO_NUM_THREADS=4 -e PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True \
-e IB_DEVICES=${RAILDEVS} -e MORI_RDMA_DEVICES=${RAILDEVS} -e NCCL_IB_HCA=${RAILDEVS} -e MORI_IB_GID_INDEX=3 -e NCCL_IB_GID_INDEX=3 -e MORI_RDMA_TC=41 -e MORI_RDMA_SL=0 \
-e MODEL_NAME=${MODEL_NAME} -e MODEL_PATH=${MODEL_PATH} \
-e xP=${xP} -e yD=${yD} -e DP_MODE=${DP_MODE} -e RUN_MORI=1 -e USE_CX7_NICS=0 -e SKIP_BENCHMARK=${SKIP_BENCHMARK:-1} -e SKIP_CURL_TEST=${SKIP_CURL_TEST:-1} -e KEEP_ALIVE=${KEEP_ALIVE:-1} \
-e MASTER_ADDR=${MASTER_ADDR} -e NODE_RANK=${rank} -e IPADDRS=${IPADDRS} \
-e NCCL_SOCKET_IFNAME=${mgmtif} -e GLOO_SOCKET_IFNAME=${mgmtif} \
-e MOONCAKE_COOKBOOK_PATH=${COOKBOOK_IN} \
--entrypoint /bin/bash ${IMG} -lc '
if [ -e /usr/local/lib/libbnxt_re-rdmav34.so ]; then echo /usr/local/lib>/etc/ld.so.conf.d/bnxt.conf; ldconfig; fi
cd ${COOKBOOK_IN}
bash ${COOKBOOK_IN}/sglang_disagg_mori_io_ep.sh
' > /dev/null && echo \" ${node} rank ${rank} launched (mgmtif=${mgmtif})\""
}

rank=0
for n in $PREFILL_NODES; do launch_node "$n" "$rank"; rank=$((rank+1)); done
for n in $DECODE_NODES; do launch_node "$n" "$rank"; rank=$((rank+1)); done
echo "All nodes launched. Router/API on first prefill node fabric IP ${MASTER_ADDR} (port per framework, default 3000/8000)."
83 changes: 83 additions & 0 deletions scripts/sglang_disagg/dsv4_flash/slo_harness.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
#!/usr/bin/env python3
# SLO harness: streaming TTFT/TPOT percentiles for DSV4-Flash via sglang router.
# Measures per-request TTFT (first token) + TPOT (mean inter-token), reports
# p50/p95/p99 TTFT, p95 TPOT, aggregate output tok/s, E2E percentiles.
import json, urllib.request, time, threading, os, sys, statistics

URL = os.environ.get("ENDPOINT", "http://192.168.200.55:2322") + "/v1/completions"
M = os.environ.get("MODEL", "/models/DeepSeek-V4-Flash-FP8-E4M3")
ISL = int(os.environ.get("ISL", "100000"))
OSL = int(os.environ.get("OSL", "1100"))
# unique-ish filler to avoid radix cache sharing across requests
FILL = "The quick brown fox jumps over the lazy dog. "

def make_prompt(isl, salt):
reps = isl // 9 + 1
base = (f"[req{salt}] " + FILL * reps)
return base[: isl * 5]

def one(prompt, osl, res, i):
t0 = time.time()
ttft = None; nevent = 0; usage_tok = 0
body = json.dumps({"model": M, "prompt": prompt, "max_tokens": osl,
"temperature": 0, "stream": True,
"stream_options": {"include_usage": True}}).encode()
req = urllib.request.Request(URL, data=body, headers={"Content-Type": "application/json"})
try:
r = urllib.request.urlopen(req, timeout=1200)
for raw in r:
line = raw.decode("utf-8", "ignore").strip()
if not line.startswith("data:"): continue
data = line[5:].strip()
if data == "[DONE]": break
try: obj = json.loads(data)
except Exception: continue
u = obj.get("usage")
if u and u.get("completion_tokens"): usage_tok = u["completion_tokens"]
txt = obj.get("choices", [{}])[0].get("text", "") if obj.get("choices") else ""
if txt:
now = time.time()
if ttft is None: ttft = now - t0
nevent += 1
e2e = time.time() - t0
ntok = usage_tok or nevent # real output tokens (usage) — correct under MTP bursts
# eff_tpot = decode wall time / REAL output tokens — burst-insensitive, correct for
# speculative decode (MTP). This is the TRUE per-output-token latency over decode.
eff_tpot = ((e2e - ttft) / max(ntok - 1, 1)) * 1000.0 if ttft is not None else 0.0
res[i] = {"ttft": ttft, "eff_tpot": eff_tpot, "ntok": ntok, "e2e": e2e}
except Exception as e:
res[i] = {"err": str(e)[:80]}

def pct(xs, p):
if not xs: return 0.0
xs = sorted(xs); k = (len(xs)-1) * p/100.0
f = int(k); c = min(f+1, len(xs)-1)
return xs[f] + (xs[c]-xs[f]) * (k-f)

def run(con):
prompts = [make_prompt(ISL, i) for i in range(con)]
res = [None]*con; ths=[]; t0=time.time()
for i in range(con):
th = threading.Thread(target=one, args=(prompts[i], OSL, res, i)); th.start(); ths.append(th)
for th in ths: th.join()
wall = time.time()-t0
ok = [r for r in res if r and "err" in r is False or (r and "ttft" in r and r["ttft"] is not None)]
ok = [r for r in res if r and r.get("ttft") is not None]
errs = [r for r in res if r and "err" in r]
if not ok:
print(f"CON={con}: ALL FAILED. sample_err={errs[0]['err'] if errs else '?'}", flush=True); return
ttfts=[r["ttft"] for r in ok]
eff=[r["eff_tpot"] for r in ok if r["eff_tpot"]>0]
ntoks=sum(r["ntok"] for r in ok); e2es=[r["e2e"] for r in ok]
print(f"CON={con} ISL={ISL} OSL={OSL}: {len(ok)}/{con} ok"
f" | TTFT p50={pct(ttfts,50):.2f}s p95={pct(ttfts,95):.2f}s p99={pct(ttfts,99):.2f}s"
f" | TPOT p50={pct(eff,50):.0f}ms p95={pct(eff,95):.0f}ms"
f" | E2E p50={pct(e2es,50):.1f}s p95={pct(e2es,95):.1f}s"
f" | agg_out={ntoks/wall:.0f} tok/s | wall={wall:.1f}s"
+ (f" | ERRS={len(errs)}" if errs else ""), flush=True)

if __name__ == "__main__":
cons = [int(x) for x in os.environ.get("CONS", "12,24,36").split(",")]
print(f"# SLO sweep ISL={ISL} OSL={OSL} cons={cons} endpoint={URL}", flush=True)
for c in cons:
run(c)
Loading