diff --git a/deploy/CURRENT.md b/deploy/CURRENT.md index 34e0a1e..3f8ec0a 100644 --- a/deploy/CURRENT.md +++ b/deploy/CURRENT.md @@ -7,7 +7,7 @@ | 机器 | 在役 | 口径 / 归属 | 对应 profile | |---|---|---|---| -| 60.1 (6000D-1) | `glm53-pp4`(Up,8 卡满载,:30000) | **方案 D 生产**。09-08 PD 压测窗口停机 ~2.5h 后已恢复并核验 | `profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4.env` | +| 60.1 (6000D-1) | `glm53-pp4`(Up,8 卡满载,:30000) | **方案 D 生产**。09-09 DSV4 优化点迁移实验窗口停机 ~7h 后已恢复并核验(09-08 PD 窗口 ~2.5h 前史)。**升级路径已交付未执行**:latest 镜像(d6e72886=sglang0.5.19+fi0.6.18)+开 autotune+持久 SGLANG_CACHE_DIR/fi_jit_cache 挂载,16k/512 cc8-64 五点 +1.2~+3.2%(A/B/A 回切确认因果),代价 ~3.9GB/卡 tactic 缓冲、KV 池不缩;profile `..._tp2pp4_latest_autotune.env`,实验全量 `experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/`。同轮判决:PCIe-IPC AR 包五点全降判负勿叠用;page-mark 向量化对 GLM 双证据判死(GLM 入口不经过该 kernel)。`dsv4_scan`(DSV4 团队扫描容器)经授权保持停止,还原说明 `/root/dsv4_scan_restore_note.txt` | `profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4.env`(现役);升级候选 `..._tp2pp4_latest_autotune.env` | | 60.2 (6000D-2) | `glm53-nvfp4` 实验容器(09-08 晚 TP1PP8 phase,run_phase.sh 实验链进行中) | **他人实验进行中,勿动**(方案 F PD 链已拆除;GPU7 曾有外部裸金属任务)。动卡前仍先核实归属 | `profiles/pro6000/glm53_nvfp4_pro6000_pd_{prefill→decode 侧}.env` | | 60.3 | 无容器,但 8 卡被外部裸金属训练占用(`/data/mas/larm`,09-08 晚实测) | 外部任务,勿动(此前台账漏记) | — | | 60.4 | `glm53-nvfp4`(Up,:30000,restart=unless-stopped) | **TP8+EAGLE+custom-AR 1stage(E7b 配方)在役**(09-09 下午部署:EAGLE 4/1/5/mem0.90/MRR16/chunk8192/ctxlen270336/fp8KV+hicache3/decode 图桶 1-8/双 parser;CAR 补丁三处全注入、8 rank `SSKJ_CAR_PATCH_ACTIVE` 确认,health 200、生成冒烟、质量门 7/7 见 `/root/qg_604_car.log`)。启动 `/root/deploy_glm53_604_exp.sh`(=60.7 实验版逐字拷贝,`RESTART=yes CAR_PATCH=1`),补丁 `/root/patches/`(md5 与仓库 car_patch 归档一致)。当日早间曾短暂部署 TP2PP4 D 配方复刻(deploy_s2_test_604.sh 留盘可切回)后被本方案替换;同日经授权清退外部 vllm 评测流水线(tmux `mas` 的 run_multiseed.sh 链,--resume 可续跑) | `experiments/pro6000/glm53_nvfp4_pro6000d_sglang_dual_scenario_bench/scripts/`(deploy_glm53_604_exp.sh + car_patch/ 补丁快照) | diff --git a/deploy/profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune.env b/deploy/profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune.env new file mode 100644 index 0000000..9c73eac --- /dev/null +++ b/deploy/profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune.env @@ -0,0 +1,49 @@ +# GLM-5.3-NVFP4 SGLang TP=2 PP=4 + latest镜像 + autotune profile(方案 D 升级版)。 +# 2026-09-09 DSV4 优化点迁移实验优胜配置:五点(16k/512 cc8-64)全胜 +1.2~+3.2%, +# A/B/A 回切确认因果成立。见 experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/。 +# 已入档未执行;60.1 生产仍跑 nightly 版(glm53-pp4,见同目录 tp2pp4.env)。 +# +# 关键点(实测,勿随意改): +# - 相对 tp2pp4.env 仅两处变化:镜像 nightly-dev-20260828 → latest(d6e72886); +# 删 --disable-flashinfer-autotune。启动参数其余逐字相同 +# - 镜像中性实测:latest 与 nightly 同配置五点差 ≤0.7%(换镜像无风险无收益) +# - autotune 增益钉在 tactic 缓存上:SGLANG_CACHE_DIR 与 flashinfer JIT 缓存 +# 必须挂宿主持久盘,否则重部署重抽签(增益消失/不可复现,DSV4 §12.6 同款教训) +# - 代价:available_gpu_mem 15.64→11.79 GB(tactic 缓冲 ≈3.9GB/卡), +# KV 池不缩(1,040,384) +# - PCIe-IPC AllReduce 包在同一实验中五点全降(-0.4~-3.5%)判负勿叠用: +# TP2 单对端 NCCL AR 同 switch P2P 已近最优 +# - 上生产须补 --tool-call-parser glm47 与 --reasoning-parser glm45(D 系既有缺口) +# - 部署器:experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/deploy_glm53_exp.sh +# (EXP_NAME=<名> EXP_IMAGE=lmsysorg/sglang:latest AUTOTUNE=1) + +PLATFORM=pro6000 +EXPERIMENT=glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune +MODEL_NAME=GLM-5.3-NVFP4 +ENGINE=sglang +RUNTIME=docker +DOCKER_IMAGE=lmsysorg/sglang:latest +DOCKER_IMAGE_DIGEST=sha256:d6e7288627be8b02be88e4bba38e73f6d50e2826869f753c13a4c4385ab3eda9 +CONTAINER_NAME=glm53-pp4-autotune +MODEL_PATH=/data/hf_models/GLM-5.3-NVFP4 +SERVED_MODEL_NAME=/data/hf_models/GLM-5.3-NVFP4 +PORT=30000 +HEALTH_PATH=/health +HEALTH_WAIT_S=600 +CONTAINER_PYTHON=python3 + +TP=2 +PP=4 +MEM_FRACTION_STATIC=0.85 +MAX_RUNNING_REQUESTS=48 +CHUNKED_PREFILL_SIZE=16384 + +DEVICE_VARS="CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7" +ENGINE_ENV="PYTHONUNBUFFERED=1 HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1 SGLANG_CACHE_DIR=/root/.cache/sglang" + +DOCKER_FLAGS="--gpus all --shm-size 64g --ipc=host -p ${PORT}:${PORT}" +VOLUMES="/data/hf_models:/data/hf_models /data/glm53_exp/fi_jit_cache:/root/.cache/flashinfer /data/glm53_exp/sglang_cache:/root/.cache/sglang /data/glm53_exp/triton_cache:/root/.triton" + +BOOTSTRAP="python3 -m sglang.launch_server ${LAUNCH_ARGS}" + +LAUNCH_ARGS="--model-path ${MODEL_PATH} --tp-size ${TP} --pp-size ${PP} --mem-fraction-static ${MEM_FRACTION_STATIC} --max-running-requests ${MAX_RUNNING_REQUESTS} --disable-radix-cache --disable-shared-experts-fusion --moe-runner-backend flashinfer_cutlass --disable-custom-all-reduce --chunked-prefill-size ${CHUNKED_PREFILL_SIZE} --host 0.0.0.0 --port ${PORT} --json-model-override-args {\"index_topk_freq\": 4}" diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/REPORT.md b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/REPORT.md new file mode 100644 index 0000000..eb2cb9e --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/REPORT.md @@ -0,0 +1,81 @@ +# DSV4-Flash 优化点迁移实验报告 — GLM-5.3-NVFP4 @ 60.1(6000D-1) + +日期:2026-09-09 | 机器:174.1.60.1(8×RTX 6000D,SM120,无 NVLink 纯 PCIe) +基线:方案 D(TP2 PP4 mono,60.1 生产 glm53-pp4 原样配方) +口径:i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐高者优(用户判定规则) +源报告:飞书《DSV4-Flash 优化日报》(A3Fmw6YePikH9lkRx8scN1UYngw) + +《DSV4-Flash 单机八卡推理优化完整报告》(Bby6wiQ9Bi8yGjkB6EIcG1CqnJe) + +## 一、迁移性判定总表(先判后测) + +| DSV4 优化点 | GLM-5.3 迁移判定 | 依据 | +|---|---|---| +| P0 native-heads(deepseek_v4.py) | **N/A,不实验** | GLM 已原生直通 flashinfer(8/16/32-head 模板齐全),无中间层可剥 | +| P1 page-mark 向量化(BLOCK=1024) | **判死(双证据),压测前撤臂** | GLM 入口 `flashinfer_sparse_mla_forward:607` 直调 trtllm_batch_decode_with_kv_cache_mla,不经过 `_split_kv_pages_to_64`(page-mark 唯一入口,DSV4 SWA 专用);运行时 GLM serving 0 个 Triton kernel 启动。证据见 `patches/page_mark_verdict_evidence.md` | +| pdi/SWA 相关整族 | **N/A** | GLM-5.3 无滑动窗口注意力 | +| chunk2048+radix 命中 | **N/A** | 本场景 16k 独立输入无共享前缀,radix OFF 为既定口径 | +| flashinfer autotune + 持久 tactic 缓存 | **可迁移 → 实验(arm2)** | DSV4 E2 实测 autotune-on + 持久缓存胜 autotune-off ~7%;GLM 同为 flashinfer 栈 | +| PCIe-IPC AllReduce 包(PR#34528 + fi#4393) | **可迁移 → 实验(arm3)** | 同代 SM120 无 NVLink 平台同痛点;DSV4 TP4 实测 +1.8% | +| 最新镜像(sglang 0.5.19 + flashinfer 0.6.18) | **可迁移 → 实验(0b)** | 用户加测;DSV4 团队验证过、IPC manifest 指定镜像 | +| DeepGEMM MoE runner / dsa-topk flashinfer / AR-fusion | **勿碰** | DSV4 报告负结果 + 本集群已知 SM120 门禁(见双场景报告勿碰清单) | + +## 二、实验臂与判决 + +| 臂 | 配置 | 判决 | +|---|---|---| +| 0a | D 配置 @ nightly-dev-20260828(生产镜像) | 基线带 102.2/148.2/191.1/198.4/208.2,与 mono 战役发表 D 数字交叉一致 | +| 0b | D 配置 @ latest(d6e72886, 0.5.19+fi0.6.18) | 与 0a 差 ≤0.7% 全点持平 → **按用户规则转 latest 跑后续臂**(镜像中性,无回归) | +| 1 | page-mark 向量化 | **压测前撤销**(双证据判死,省一整轮),见判定总表 | +| 2 | latest + autotune-on + 持久 SGLANG_CACHE_DIR/fi_jit_cache | **五点全胜 +1.2~+3.2%(唯一胜者)**:104.8/151.1/195.0/200.6/210.0;cc16/32/64 分布与基线不重叠;TPOT 全点改善(−2~−4.6%);代价 available_gpu_mem 15.64→11.79 GB(≈3.9GB/卡 tactic 缓冲),KV 池不缩(1,040,384);门 6/7 不变 | +| 3 | latest + PCIe-IPC 12 文件包 + 双 env 开关 | **五点全降判负**:97.9/143.2/187.1/197.3/204.5(−0.4~−3.5%),cc8/16/32/64 分布不重叠(方向为劣),TPOT 全点变差;8 rank 均确认 `FlashInfer PCIe-IPC all-reduce enabled (world=2)`——激活正确、真输。机制:TP2 单对端 NCCL AR 走同 switch P2P 已近最优,IPC 工作区拷贝+跨块同步纯倒贴;DSV4 的 +1.8% 本就是 TP4 口径,TP 度越低收益越小,TP2 翻负。与既有结论「TP2 AR≈free」互证 | +| 4 | 冠军组合 | 不需要(唯一胜者 autotune 已完整测量=arm2) | +| 5 | 回切确认:stock-latest 全新部署 1 轮 | **PASS**:102.6/147.0/190.9/200.6/207.8 全落 0b 基线带(缓存已热,无 0b r1 式冷 JIT 压低),且低于 arm2 → A/B/A 闭环,**autotune 增益判定为因果非环境漂移** | + +原始数据:`results/results_raw.md`。 + +## 三、优胜配置交付(已入档未执行) + +**推荐部署 = 方案 D 配方升级两处:镜像换 latest + 开 autotune(持久缓存)** + +``` +镜像: lmsysorg/sglang:latest (sha256:d6e7288627be8b02be88e4bba38e73f6d50e2826869f753c13a4c4385ab3eda9 + = sglang 0.5.19 + flashinfer 0.6.18;生产现役 nightly-dev-20260828 同 digest 已本地 tag 为 dev) +启动: python3 -m sglang.launch_server --model-path /data/hf_models/GLM-5.3-NVFP4 \ + --tp-size 2 --pp-size 4 --mem-fraction-static 0.85 --max-running-requests 48 \ + --disable-radix-cache --disable-shared-experts-fusion \ + --moe-runner-backend flashinfer_cutlass \ + --disable-custom-all-reduce --chunked-prefill-size 16384 \ + --host 0.0.0.0 --port 30000 --json-model-override-args '{"index_topk_freq": 4}' + # 相对现役:镜像 nightly-dev→latest;删 --disable-flashinfer-autotune;其余逐字不变 +挂载(新增,持久化必需): -v /fi_jit_cache:/root/.cache/flashinfer \ + -v /sglang_cache:/root/.cache/sglang \ + -e SGLANG_CACHE_DIR=/root/.cache/sglang +``` + +**关键运维纪律(DSV4 §12.6 同款教训)**:autotune 增益钉在 tactic 缓存上。 +SGLANG_CACHE_DIR / flashinfer JIT 缓存**必须挂载到宿主持久盘**——否则每次重部署 +重抽 tactic 签,轻则增益消失、重则抽到差 tactic 比关闭还慢,且引入不可复现性。 +缓存目录随部署资产同迁移(本实验用 /data/glm53_exp/{fi_jit_cache,sglang_cache})。 + +上生产另须补 `--tool-call-parser glm47 --reasoning-parser glm45`(D 配置既有缺口)。 + +profile:`deploy/profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune.env` +部署器:`scripts/deploy_glm53_exp.sh`(EXP_NAME=x EXP_IMAGE=lmsysorg/sglang:latest AUTOTUNE=1) + +## 四、机器状态与还原 + +- 原生产容器 `glm53-pp4` 配置全程未动,实验后 `docker start` 还原(health 核验) +- `dsv4_scan`(DSV4 团队扫描容器,实验前确认已跑完 45min 无日志无 GPU 占用,经授权 stop) + 保持停止,还原说明在 60.1:/root/dsv4_scan_restore_note.txt(`docker start dsv4_scan` 即还原) +- 实验容器 glm53-exp0a/0b/2/3/5 全部拆除,/data/glm53_exp/(缓存+补丁+日志)留盘 + +## 五、资产清单 + +- `scripts/`:deploy_glm53_exp.sh(参数化部署器)/ run_exp_s2.sh(五点单轮)/ + run_arm_chain.sh(门+3轮链) +- `patches/`:page_mark_verdict_evidence.md(判死证据)+ base_flash_mla_sm120.py + (两镜像共有基础模块 md5 9df4d50d)+ page_mark_vectorized/ + native_heads/(N/A 留档) + + pcie_ipc_v1/(12 文件判负快照,sha256 校验 ok=12) +- `results/results_raw.md`:逐轮原始数据 + 中位对比 + TPOT + 显存口径 +- 远端:60.1 /root/{deploy_glm53_exp.sh,run_exp_s2.sh,run_arm_chain.sh}、 + /root/bench_logs/exp_*.log、/data/glm53_exp/{logs,patches,inject} diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/base_flash_mla_sm120.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/base_flash_mla_sm120.py new file mode 100644 index 0000000..6f39fa0 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/base_flash_mla_sm120.py @@ -0,0 +1,644 @@ +"""SM120 FlashMLA sparse decode implementation. + +On SM120 (Blackwell Desktop / RTX PRO 6000) the flash_mla CUDA kernel +is not available, so this module provides alternative implementations: + +- A fused Triton kernel (default, ``SGLANG_SM120_TRITON_FLASHMLA=1``) +- A pure-PyTorch fallback (``SGLANG_SM120_TRITON_FLASHMLA=0``) + +The FP8 KV cache uses a page-internal layout where NOPE+ROPE data has +stride (nope_dim + rope_dim*2) per token, and scales are stored in a +separate region at the end of each page. +""" + +import logging +import math +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.srt.environ import envs +from sglang.srt.utils import is_hip + +logger = logging.getLogger(__name__) +_is_hip = is_hip() + +_GLM_DSA_MODEL_ARCHS = ( + "GlmMoeDsaForCausalLM", + "GlmMoeDsaForCausalLMNextN", +) + +# Page layout constants for DSv4-Flash (MODEL1): +# nope_dim = 448, rope_dim = 64, quantize_block_size = 64 +# nope_rope_stride = 448 + 64*2 = 576 bytes per token +# scale_stride = ceil(448/64) + 1 = 8 bytes per token (7 scales + 1 pad) +# bytes_per_token = 448 + 128 + 8 = 584 +# page_bytes = ceil_div(page_size * 584, 576) * 576 + +_NOPE_DIM = 448 +_ROPE_DIM = 64 +_NOPE_ROPE_STRIDE = _NOPE_DIM + _ROPE_DIM * 2 # 576 +_TILE_SIZE = 64 +_NUM_TILES = _NOPE_DIM // _TILE_SIZE # 7 +_SCALE_STRIDE = _NUM_TILES + 1 # 8 (7 scales + 1 pad) +_D = _NOPE_DIM + _ROPE_DIM # 512 + + +def _gather_and_dequant(k_cache, indices, page_size): + """Gather KV entries from the paged buffer using correct page-internal addressing. + + Args: + k_cache: (num_pages, page_size, 1, bytes_per_token) float8_e4m3fn + Non-contiguous view of the raw page buffer. + indices: (...) int32/int64, token-level indices. -1 = invalid. + page_size: tokens per page (256) + + Returns: + kv: (..., _D) bfloat16, dequantized KV vectors + """ + idx_shape = indices.shape + flat_idx = indices.reshape(-1) # (N,) + N = flat_idx.shape[0] + device = k_cache.device + + # Page-level addressing + page_bytes = k_cache.stride(0) # actual byte stride between pages + pages = flat_idx // page_size + offsets = flat_idx % page_size + + # Clamp invalid indices + safe_pages = pages.clamp(min=0) + safe_offsets = offsets.clamp(min=0) + + # Access raw buffer as uint8 — use as_strided to get full page view + num_pages = k_cache.shape[0] + raw_pages = k_cache.as_strided( + (num_pages, page_bytes), + (page_bytes, 1), + ).view( + torch.uint8 + ) # (num_pages, page_bytes) uint8 + # Note: float8_e4m3fn and uint8 are both 1 byte, view is safe + + # Compute byte offsets within each page + # NOPE: page[safe_page, safe_offset * 576 + 0:448] + # ROPE: page[safe_page, safe_offset * 576 + 448:576] + # SCALES: page[safe_page, page_size * 576 + safe_offset * 8 + 0:7] + + nope_base = safe_offsets * _NOPE_ROPE_STRIDE # (N,) + nope_offsets = nope_base.unsqueeze(-1) + torch.arange( + _NOPE_DIM, device=device, dtype=torch.long + ) # (N, 448) + + rope_base = nope_base + _NOPE_DIM # (N,) + rope_offsets = rope_base.unsqueeze(-1) + torch.arange( + _ROPE_DIM * 2, device=device, dtype=torch.long + ) # (N, 128) + + scale_section_offset = page_size * _NOPE_ROPE_STRIDE # 147456 + scale_base = scale_section_offset + safe_offsets * _SCALE_STRIDE # (N,) + scale_offsets = scale_base.unsqueeze(-1) + torch.arange( + _NUM_TILES, device=device, dtype=torch.long + ) # (N, 7) + + # Gather bytes per page — use advanced indexing + # raw_pages[safe_pages, nope_offsets] → (N, 448) + page_idx_nope = safe_pages.unsqueeze(-1).expand_as(nope_offsets) + nope_bytes = raw_pages[page_idx_nope, nope_offsets] # (N, 448) uint8 + + page_idx_rope = safe_pages.unsqueeze(-1).expand_as(rope_offsets) + rope_bytes = raw_pages[page_idx_rope, rope_offsets] # (N, 128) uint8 + + page_idx_scale = safe_pages.unsqueeze(-1).expand_as(scale_offsets) + scale_bytes = raw_pages[page_idx_scale, scale_offsets] # (N, 7) uint8 + + # Reinterpret dtypes + nope_fp8 = nope_bytes.view(torch.float8_e4m3fn) # (N, 448) + rope_bf16 = rope_bytes.contiguous().view(torch.bfloat16) # (N, 64) + scale_e8m0 = scale_bytes.view(torch.float8_e8m0fnu) # (N, 7) + + # Dequantize: nope_tile * scale_tile → bf16 (vectorized) + result = torch.empty(N, _D, dtype=torch.bfloat16, device=device) + result[:, :_NOPE_DIM] = ( + ( + nope_fp8.view(N, _NUM_TILES, _TILE_SIZE).float() + * scale_e8m0.view(N, _NUM_TILES, 1).float() + ) + .view(N, _NOPE_DIM) + .to(torch.bfloat16) + ) + result[:, _NOPE_DIM:] = rope_bf16 + + return result.reshape(*idx_shape, _D) + + +def _sm120_sparse_decode_fwd( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache=None, + extra_indices=None, + extra_topk_length=None, +): + B, s_q, H_q, D_qk = q.shape + num_pages, page_size, H_k, bpt = k_cache.shape + topk = indices.shape[-1] + + invalid_mask = indices < 0 + safe_indices = indices.clamp(min=0) + + if topk_length is not None: + topk_range = torch.arange(topk, device=topk_length.device).view(1, 1, topk) + invalid_mask = invalid_mask | (topk_range >= topk_length.view(B, 1, 1)) + + # Gather and dequantize using page-aware addressing + gathered_kv = _gather_and_dequant(k_cache, safe_indices, page_size) + + if extra_k_cache is not None and extra_indices is not None: + extra_topk = extra_indices.shape[-1] + extra_page_size = extra_k_cache.shape[1] + extra_invalid = extra_indices < 0 + extra_safe = extra_indices.clamp(min=0) + if extra_topk_length is not None: + extra_range = torch.arange( + extra_topk, device=extra_topk_length.device + ).view(1, 1, extra_topk) + extra_invalid = extra_invalid | ( + extra_range >= extra_topk_length.view(B, 1, 1) + ) + extra_kv = _gather_and_dequant(extra_k_cache, extra_safe, extra_page_size) + gathered_kv = torch.cat([gathered_kv, extra_kv], dim=2) + invalid_mask = torch.cat([invalid_mask, extra_invalid], dim=2) + + gathered_kv[invalid_mask] = 0.0 + + q_f = q.float() + kv_f = gathered_kv.float() + kv_d = kv_f.shape[-1] + if D_qk != kv_d: + q_f = q_f[..., :kv_d] + + scores = torch.einsum("bshd,bstd->bsht", q_f, kv_f) * softmax_scale + scores.masked_fill_(invalid_mask.unsqueeze(2).expand_as(scores), float("-inf")) + + lse = torch.logsumexp(scores, dim=-1) + + if attn_sink is not None: + lse_for_out = torch.logsumexp( + torch.stack([lse, attn_sink.view(1, 1, H_q).expand_as(lse)], dim=0), dim=0 + ) + else: + lse_for_out = lse.clone() + + lonely = lse == float("-inf") + lse_for_out[lonely] = float("inf") + weights = torch.exp(scores - lse_for_out.unsqueeze(-1)) + out = torch.einsum("bsht,bstv->bshv", weights, kv_f[..., :head_dim_v]) + out[lonely.unsqueeze(-1).expand_as(out)] = 0.0 + + return out.to(torch.bfloat16), lse.permute(0, 2, 1) + + +# SM120 FlashMLA: default FlashInfer (CUTLASS SM120 sparse MLA decode). +# Override with SGLANG_SM120_FLASHMLA_BACKEND=triton|torch to force fallback. +_sm120_default_backend = envs.SGLANG_SM120_FLASHMLA_BACKEND.get() + + +def flash_mla_with_kvcache_sm120(**kwargs): + """SM120 FlashMLA sparse decode entry point. + + Dispatches to FlashInfer (default if available), Triton, or PyTorch fallback. + """ + q = kwargs["q"] + k_cache = kwargs["k_cache"] + indices = kwargs["indices"] + topk_length = kwargs.get("topk_length") + attn_sink = kwargs.get("attn_sink") + head_dim_v = kwargs["head_dim_v"] + softmax_scale = kwargs.get("softmax_scale") + if softmax_scale is None: + softmax_scale = q.shape[-1] ** (-0.5) + extra_k_cache = kwargs.get("extra_k_cache") + extra_indices = kwargs.get("extra_indices_in_kvcache") + extra_topk_length = kwargs.get("extra_topk_length") + + if _sm120_default_backend == "flashinfer": + return _flash_mla_flashinfer( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + + if _sm120_default_backend == "triton": + from sglang.kernels.ops.attention.flash_mla_sm120_triton import ( + flash_mla_sparse_decode_triton, + ) + + out, lse = flash_mla_sparse_decode_triton( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + return (out, lse) + + out, lse = _sm120_sparse_decode_fwd( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + return (out, lse) + + +# --- Page-split utilities: pbs=256 → pbs=64 --- +# SGLang SWA KV cache footer layout per 256-token page: +# [data: 256 * 576 bytes] [scale: 256 * 8 bytes] [padding] +# FlashInfer decode_dsv4 expects per 64-token page: +# [data: 64 * 576 bytes] [scale: 64 * 8 bytes] [padding to 37440] +_PBS_SRC = 256 # SGLang physical page size +_PBS_DST = 64 # FlashInfer page_block_size +_NOPE_ROPE_STRIDE = 576 # bytes per token for nope+rope +_SCALE_STRIDE = 8 # bytes per token for scale (7 + 1 pad) +_BYTES_PER_DST_PAGE = ( + _PBS_DST * _NOPE_ROPE_STRIDE + _PBS_DST * _SCALE_STRIDE +) # 64*576 + 64*8 = 37376 + 512 = 37888 +# Padded to 576 alignment + +_BYTES_PER_DST_PAGE_PADDED = math.ceil(_BYTES_PER_DST_PAGE / 576) * 576 # 37440 + + +@triton.jit +def _page_split_kernel( + src_ptr, + dst_ptr, + N_pages, + src_stride0: tl.constexpr, + dst_stride0: tl.constexpr, + DATA_PER_SUB: tl.constexpr, # 64 * 576 = 36864 + SCALE_PER_SUB: tl.constexpr, # 64 * 8 = 512 + SRC_SCALE_OFF: tl.constexpr, # 256 * 576 = 147456 + DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864 + RATIO: tl.constexpr, # 4 + BLOCK_SIZE: tl.constexpr, + mask_ptr, + HAS_MASK: tl.constexpr, +): + """Fused page-split: copy data+scale for all sub-pages in one kernel. + + When HAS_MASK is set, only pages flagged in ``mask_ptr`` (int8, 1=touched) + are copied; untouched pages are skipped so the kernel no longer rewrites the + entire KV pool every decode step. + """ + pid = tl.program_id(0) + page_idx = pid // RATIO + sub = pid % RATIO + + if page_idx >= N_pages: + return + + if HAS_MASK: + if tl.load(mask_ptr + page_idx) == 0: + return + + src_base = src_ptr + page_idx * src_stride0 + dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0 + + # Copy data region: DATA_PER_SUB bytes from src offset sub*DATA_PER_SUB + data_src_off = sub * DATA_PER_SUB + for start in tl.range(0, DATA_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < DATA_PER_SUB + vals = tl.load(src_base + data_src_off + offs, mask=mask) + tl.store(dst_base + offs, vals, mask=mask) + + # Copy scale region: SCALE_PER_SUB bytes + scale_src_off = SRC_SCALE_OFF + sub * SCALE_PER_SUB + for start in tl.range(0, SCALE_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < SCALE_PER_SUB + vals = tl.load(src_base + scale_src_off + offs, mask=mask) + tl.store(dst_base + DST_SCALE_OFF + offs, vals, mask=mask) + + +@triton.jit +def _page_mark_kernel( + indices_ptr, + mask_ptr, + N_idx, + SRC_PBS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Mark touched source pages (1 byte each) from token-level indices. + + ``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid. + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ + pid = tl.program_id(0) + if pid >= N_idx: + return + idx = tl.load(indices_ptr + pid) + if idx < 0: + return + page = idx // SRC_PBS + tl.store(mask_ptr + page, 1) + + +def _split_kv_pages_to_64( + kv_u8: torch.Tensor, + src_pbs: int, + touched_indices: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Split pbs=N footer-format pages into pbs=64 footer-format pages. + + When ``touched_indices`` (token-level int32 indices into the pbs=src_pbs + SWA pool, -1 = invalid) is provided, only the source pages that actually + contain a referenced token are copied. This avoids rewriting the entire KV + pool on every decode step (only ~2*batch pages are touched vs the full + pool). The output buffer is persistent and reused across steps; untouched + dst pages simply retain their (unreferenced) stale data. + """ + assert src_pbs % _PBS_DST == 0 and src_pbs >= _PBS_DST + if src_pbs == _PBS_DST: + return kv_u8 + + N = kv_u8.shape[0] + ratio = src_pbs // _PBS_DST + num_dst_pages = N * ratio + + from sglang.srt.runtime_context import get_resources + + # Pre-allocated grow-only buffer for page-split output per device. + dev = kv_u8.device + buffers = get_resources().buffers + key = f"flash_mla_sm120_split:{dev}" + buf = buffers.get(key) + if buf is None or buf.shape[0] < num_dst_pages: + # The first allocation can happen under inference mode (autotune), but + # the buffer is written again during CUDA graph capture outside + # inference mode, where an inference tensor cannot be mutated. + with torch.inference_mode(False): + buf = torch.empty( + num_dst_pages, + _BYTES_PER_DST_PAGE_PADDED, + dtype=torch.uint8, + device=dev, + ) + buffers[key] = buf + out = buf[:num_dst_pages] + + # Get raw 2D view of source + src_2d = kv_u8 + if src_2d.ndim == 4: + src_stride0 = src_2d.stride(0) + src_2d = torch.as_strided(src_2d, (N, src_stride0), (src_stride0, 1)) + else: + src_stride0 = src_2d.stride(0) + + use_mask = touched_indices is not None and touched_indices.numel() > 0 + mask_ptr = src_2d # dummy, never dereferenced when HAS_MASK is False + if use_mask: + # Persistent per-device int8 mask, zeroed each call (cheap memset, + # captured cleanly by CUDA graph). 1 = page is referenced this step. + mkey = f"flash_mla_sm120_mask:{dev}" + mbuf = buffers.get(mkey) + if mbuf is None or mbuf.shape[0] < N: + # The first allocation can happen under inference mode (autotune), + # but the buffer is zeroed again later during CUDA graph capture + # outside inference mode -- an inference tensor cannot be mutated + # there, so force a normal tensor. + with torch.inference_mode(False): + mbuf = torch.empty(N, dtype=torch.int8, device=dev) + buffers[mkey] = mbuf + mask = mbuf[:N] + mask.zero_() + idx_flat = touched_indices.reshape(-1).contiguous() + if idx_flat.dtype != torch.int32: + idx_flat = idx_flat.to(torch.int32) + _page_mark_kernel[(idx_flat.numel(),)]( + idx_flat, + mask, + idx_flat.numel(), + src_pbs, # SRC_PBS + 1024, # BLOCK (unused, kept for JIT signature) + ) + mask_ptr = mask + + grid = (N * ratio,) + _page_split_kernel[grid]( + src_2d, + out, + N, + src_stride0, + _BYTES_PER_DST_PAGE_PADDED, + _PBS_DST * _NOPE_ROPE_STRIDE, # DATA_PER_SUB = 36864 + _PBS_DST * _SCALE_STRIDE, # SCALE_PER_SUB = 512 + src_pbs * _NOPE_ROPE_STRIDE, # SRC_SCALE_OFF = 147456 + _PBS_DST * _NOPE_ROPE_STRIDE, # DST_SCALE_OFF = 36864 + ratio, # RATIO = 4 + 1024, # BLOCK_SIZE + mask_ptr, + use_mask, # HAS_MASK + ) + + bpt = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 584 + return out.as_strided( + (num_dst_pages, _PBS_DST, 1, bpt), + (_BYTES_PER_DST_PAGE_PADDED, bpt, bpt, 1), + ) + + +def _flash_mla_flashinfer( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, +): + """FlashInfer SM120 sparse MLA via the paged-attention dispatcher. + + SGLang SWA pool uses page_size=256 (footer format: 256*576 bytes data + 256*8 bytes scale). + FlashInfer decode_dsv4 fast path requires page_block_size=64 (footer: 64*576 + 64*8). + We split 256-token pages into 4 virtual 64-token pages. + Token indices are invariant under page-split (identity mapping). + """ + from flashinfer.mla._sparse_mla_sm120 import ( + _DECODE_MAX_TOKENS as _FI_DECODE_MAX_TOKENS, + ) + from flashinfer.mla._sparse_mla_sm120 import ( + _sparse_mla_sm120_paged_attention, + ) + + B, _, H, D = q.shape # (batch, 1, num_heads, head_dim) + dev = q.device + + # Indices: no remapping needed (page-split preserves token addressing). + idx = indices.squeeze(1) if indices.dim() == 3 else indices + + # --- Page-split: convert pbs=N kv_cache to pbs=64 view --- + # Only the SWA pages actually referenced by `idx` are copied (the rest of + # the persistent dst buffer is left untouched and never read). + kv_u8 = k_cache.view(torch.uint8) if k_cache.dtype != torch.uint8 else k_cache + src_pbs = k_cache.shape[1] if k_cache.ndim >= 3 else _PBS_SRC + kv_64 = ( + _split_kv_pages_to_64(kv_u8, src_pbs, touched_indices=idx) + if src_pbs != _PBS_DST + else kv_u8 + ) + + extra_kv_u8 = ( + extra_k_cache.view(torch.uint8) + if extra_k_cache is not None and extra_k_cache.dtype != torch.uint8 + else extra_k_cache + ) + extra_kv_64 = extra_kv_u8 + + extra_idx = ( + extra_indices.squeeze(1) + if extra_indices is not None and extra_indices.dim() == 3 + else extra_indices + ) + + output = torch.empty(B, H, head_dim_v, dtype=torch.bfloat16, device=dev) + out_lse = torch.empty(B, H, dtype=torch.float32, device=dev) + + # Use split-K for decode-sized batches and paged attention otherwise. + if B <= _FI_DECODE_MAX_TOKENS: + topk = idx.shape[-1] + extra_topk = extra_idx.shape[-1] if extra_idx is not None else 0 + _BI = 64 + num_splits = (topk + _BI - 1) // _BI + ( + (extra_topk + _BI - 1) // _BI if extra_topk > 0 else 0 + ) + mid_out = torch.empty( + B, H, num_splits, head_dim_v, dtype=torch.bfloat16, device=dev + ) + mid_lse = torch.empty(B, H, num_splits, dtype=torch.float32, device=dev) + else: + mid_out = None + mid_lse = None + + _sparse_mla_sm120_paged_attention( + q.squeeze(1) if q.ndim == 4 else q, + kv_64, + idx, + output, + out_lse, + softmax_scale, + d_v=head_dim_v, + topk_length=topk_length, + attn_sink=attn_sink, + extra_kv_cache=extra_kv_64, + extra_indices=extra_idx, + extra_topk_length=extra_topk_length, + mid_out=mid_out, + mid_lse=mid_lse, + ) + + return (output.unsqueeze(1), None) + + +def _validate_flashinfer_sparse_mla_backend( + *, + model_arch: str, + device_sm_major: int, + kv_cache_dtype: torch.dtype, + prefill_impl: str, + decode_impl: str, +) -> bool: + selected = {prefill_impl, decode_impl} + uses_flashinfer_sparse_mla = "flashinfer_sparse_mla" in selected + is_glm_sm12_fp8 = ( + model_arch in _GLM_DSA_MODEL_ARCHS + and device_sm_major == 12 + and kv_cache_dtype == torch.float8_e4m3fn + and not _is_hip + ) + if uses_flashinfer_sparse_mla and not is_glm_sm12_fp8: + raise ValueError( + "flashinfer_sparse_mla supports only GLM DSA with FP8 KV cache " + "on NVIDIA SM120/SM121; " + f"got model_arch={model_arch!r}, sm_major={device_sm_major}, " + f"kv_cache_dtype={kv_cache_dtype}, prefill_impl={prefill_impl!r}, " + f"decode_impl={decode_impl!r}." + ) + if is_glm_sm12_fp8: + unsupported = selected - {"flashinfer_sparse_mla"} + if unsupported: + raise ValueError( + "GLM DSA with FP8 KV cache on NVIDIA SM120/SM121 supports " + "only flashinfer_sparse_mla, " + f"but got {sorted(unsupported)}." + ) + return uses_flashinfer_sparse_mla + + +def flashinfer_sparse_mla_forward( + q: torch.Tensor, + kv_cache: torch.Tensor, + indices: torch.Tensor, + seq_lens: torch.Tensor, + workspace_buffer: torch.Tensor, + *, + page_size: int, + kv_cache_dim: int, + qk_nope_head_dim: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + sm_scale: float, + skip_softmax_threshold_scale_factor: float | None, +) -> torch.Tensor: + """Run FlashInfer's SM120 sparse MLA kernel on SGLang's packed DSA cache.""" + from flashinfer.mla import trtllm_batch_decode_with_kv_cache_mla + + topk = indices.shape[1] + result = trtllm_batch_decode_with_kv_cache_mla( + query=q.unsqueeze(1), + kv_cache=kv_cache.view(torch.uint8) + .view(-1, page_size, kv_cache_dim) + .unsqueeze(1), + workspace_buffer=workspace_buffer, + qk_nope_head_dim=qk_nope_head_dim, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + block_tables=indices.unsqueeze(1), + seq_lens=seq_lens, + max_seq_len=topk, + sparse_mla_top_k=topk, + bmm1_scale=float(sm_scale), + bmm2_scale=1.0, + kv_scale_format="arbitrary_fp32", + skip_softmax_threshold_scale_factor=skip_softmax_threshold_scale_factor, + ) + return result.squeeze(1) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/deepseek_v4.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/deepseek_v4.py new file mode 100644 index 0000000..ca7dbe2 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/deepseek_v4.py @@ -0,0 +1,4061 @@ +from __future__ import annotations + +import concurrent.futures +import functools +import logging +import time +from contextlib import contextmanager, nullcontext +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Iterable, + List, + NamedTuple, + Optional, + Set, + Tuple, + Union, +) + +import torch +import torch.nn as nn +import torch.nn.functional as F + +import sglang.srt.models.deepseek_v2 as deepseek_v2 +from sglang.kernels.ops.attention.dsv4 import ( + fused_norm_rope_inplace, + fused_q_norm_rope, + fused_rope_inplace, + sglang_per_token_group_quant_fp8_dsv4_wo_a, +) +from sglang.kernels.ops.quantization.fp8_kernel import ( + sglang_per_token_group_quant_fp8, +) +from sglang.srt.compilation.compilation_config import register_split_op +from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config +from sglang.srt.distributed import ( + get_pp_group, + get_tp_group, +) +from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + use_symmetric_memory, +) +from sglang.srt.environ import envs +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder +from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation +from sglang.srt.hardware_backend.npu.dsv4.dsv4_rope import Dsv4NpuRoPE +from sglang.srt.layers.attention.dsa.utils import ( + can_dsa_cp_split, + dsa_use_prefill_cp, + is_dsa_enable_prefill_cp, + is_dsa_prefill_cp_round_robin_split, +) +from sglang.srt.layers.attention.dsv4.compressor import Compressor +from sglang.srt.layers.attention.dsv4.indexer import C4Indexer +from sglang.srt.layers.communicator import get_attn_tp_context +from sglang.srt.layers.communicator_dsa_cp import ( + dsa_cp_gather_hidden_states, + dsa_cp_reduce_scatter_hidden_states, +) +from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx +from sglang.srt.layers.cp.utils import ( + cp_materialize_global_token_order, + cp_round_robin_input_ids_v2, + is_cp_v2_active, +) +from sglang.srt.layers.dp_attention import ( + _tbo_event, + attn_cp_overlap_all_gather_into_tensor, + attn_cp_overlap_reduce_scatter_tensor, + attn_tp_all_gather, + attn_tp_all_reduce, + dp_gather_partial, + dp_gather_replicate, + dp_reduce_scatter_tensor, + dp_reduce_scatterv_async, + dp_scatter, + get_dp_global_num_tokens, + get_dp_tbo_comm_stream, + get_global_dp_buffer, + get_global_dp_buffer_len, + get_local_dp_buffer, + get_local_dp_buffer_len, + get_tbo_persistent_buffer, + is_allocation_symmetric, + is_dp_attention_enabled, + is_dp_gatherv_active, +) +from sglang.srt.layers.layernorm import RMSNorm +from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear +from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv +from sglang.srt.layers.moe.fused_moe_triton import FusedMoE +from sglang.srt.layers.moe.utils import ( + is_shared_experts_fusion_disabled, + uses_per_rank_fused_shared_slots, +) +from sglang.srt.layers.quantization.fp8_utils import ( + view_aiter_fused_rms_transposed_fp8_scale, +) +from sglang.srt.layers.rotary_embedding import get_rope_wrapper +from sglang.srt.layers.utils import PPMissingLayer, get_layer_id +from sglang.srt.layers.utils.cp_utils import ( + cp_all_gather_rerange_finish, + cp_all_gather_rerange_launch, + cp_all_gather_rerange_output, + cp_round_robin_input_ids, + cp_split_and_rebuild_data, + cp_split_and_rebuild_position, + prepare_context_parallel_metadata, +) +from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding +from sglang.srt.mem_cache.memory_pool import RadixAttention +from sglang.srt.model_executor.cuda_graph_config import ( + Backend, + Phase, + check_cuda_graph_backend, +) +from sglang.srt.model_executor.forward_batch_info import PPProxyTensors +from sglang.srt.model_executor.forward_context import ( + get_attn_backend, + get_token_to_kv_pool, +) +from sglang.srt.model_executor.runner import ( + compile_in_capture_mode, + get_is_capture_mode, +) +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.breakable_cuda_graph import ( + eager_on_graph, +) +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( + is_in_breakable_cuda_graph, +) +from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( + get_tc_piecewise_forward_context, +) +from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load +from sglang.srt.model_loader.weight_utils import ( + RUNAI_STREAMER_TENSOR_ATTR, + default_weight_loader, +) +from sglang.srt.models.dbrx import ReplicatedLinear +from sglang.srt.models.deepseek_common.amd.deepseek_v4_fused_mhc import ( + apply_mhc_post_pre_boundary, + is_cross_layer_mhc_fusion_enabled, +) +from sglang.srt.models.deepseek_common.utils import ( + _use_aiter_bpreshuffle_gfx95, + is_wint4afp8_or_wint4a16_config, + quant_blocks_shared_experts_fusion, +) +from sglang.srt.models.deepseek_v2 import ( + ParallelLMHead, + _is_cuda, + _is_hip, + _is_npu, + _is_xpu, +) +from sglang.srt.runtime_context import ( + get_device, + get_exec, + get_forward, + get_parallel, + get_platform, +) + +if not _is_hip: + from sglang.srt.layers.utils.cp_utils import ( + prepare_context_parallel_metadata, + ) + +from sglang.srt.utils import ( + LazyValue, + add_prefix, + get_bool_env_var, + is_gfx95_supported, + is_gfx942_supported, + is_gfx1250_supported, + log_info_on_rank0, + make_layers, +) +from sglang.srt.utils.custom_op import register_custom_op +from sglang.srt.utils.hf_transformers_utils import get_rope_config + +# NPU-only: bind torch_npu here so _compute_q_b / _forward_prepare can call +# torch_npu.npu_rms_norm directly (imports elsewhere aren't visible in this module). +if _is_npu: + import torch_npu + + +class MhcOps(NamedTuple): + hc_split_sinkhorn: Callable[..., Any] + mhc_fused_post_pre: Optional[Callable[..., Any]] + npu_hc_pre: Optional[Callable[..., Any]] + + +@functools.cache +def _get_mhc_ops() -> MhcOps: + """Load MHC kernels only when a DeepSeek-V4 layer needs them. + + Model modules are imported eagerly by the registry. Importing + ``sglang.kernels.ops.layernorm.mhc`` owns TileLang-backed MHC kernels. + Import it only when a DeepSeek-V4 layer executes so registry discovery + cannot initialize an optional CUDA runtime before unrelated models set up + their communication workspaces. DeepSeek-V4 is the sole consumer here. + """ + if _is_xpu: + from sgl_kernel import hc_split_sinkhorn + + return MhcOps(hc_split_sinkhorn, None, None) + + from sglang.kernels.ops.layernorm.mhc import ( + hc_split_sinkhorn, + mhc_fused_post_pre, + npu_hc_pre, + ) + + return MhcOps(hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre) + + +logger = logging.getLogger(__name__) + +_FP8_WO_A_GEMM = envs.SGLANG_OPT_FP8_WO_A_GEMM.get() +_MHC_POST_MULT_VALUE = 2.0 + +DEEPSEEK_V4_STACKED_PARAMS_MAPPING: List[Tuple[str, str, int]] = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), +] + + +# FlashInfer's mhc_pre_big_fuse only accepts these split-K counts. +_FLASHINFER_MHC_PRE_SPLITS = (1, 2, 4, 8, 16) + + +@functools.cache +def _cuda_sm_count() -> int: + return torch.cuda.get_device_properties(0).multi_processor_count + + +def _flashinfer_mhc_pre_num_splits(num_tokens: int, hc_hidden_size: int) -> int: + block_m = block_k = 64 + grid_m = (num_tokens + block_m - 1) // block_m + num_block_k = (hc_hidden_size + block_k - 1) // block_k + raw = max(1, min(_cuda_sm_count() // max(grid_m, 1), num_block_k // 4)) + best = 1 + for split in _FLASHINFER_MHC_PRE_SPLITS: + if split <= raw: + best = split + return best + + +def _flashinfer_hc_pre( + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + *, + rms_eps: float, + hc_eps: float, + sinkhorn_iters: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + from flashinfer.mhc import mhc_pre_big_fuse + + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import tf32_hc_prenorm_gemm + + num_tokens, hc_mult, hidden_size = x.shape + hc_hidden_size = hc_mult * hidden_size + mix_dim = hc_fn.shape[0] # hc_mult * (2 + hc_mult) == 24 + n_splits = _flashinfer_mhc_pre_num_splits(num_tokens, hc_hidden_size) + + dot_mix = torch.empty( + (n_splits, num_tokens, mix_dim), dtype=torch.float32, device=x.device + ) + sqrsum = torch.empty((n_splits, num_tokens), dtype=torch.float32, device=x.device) + tf32_hc_prenorm_gemm( + x.reshape(num_tokens, hc_hidden_size), hc_fn, dot_mix, sqrsum, n_splits + ) + if n_splits == 1: + dot_mix = dot_mix.squeeze(0) + sqrsum = sqrsum.squeeze(0) + + post, comb, layer_input = mhc_pre_big_fuse( + dot_mix, + sqrsum, + x, + hc_scale, + hc_base, + hc_hidden_size, + rms_eps=rms_eps, + mhc_pre_eps=hc_eps, + mhc_sinkhorn_eps=hc_eps, + mhc_post_mult_value=_MHC_POST_MULT_VALUE, + sinkhorn_repeat=sinkhorn_iters, + num_splits=n_splits, + ) + return layer_input, post.squeeze(-1), comb + + +_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip +# PoC: compute the (replicated TP1) shared expert on LOCAL hidden before the dp +# gather instead of on the gathered global buffer. Requires +# SGLANG_SHARED_EXPERT_TP1=1 (replicated shared expert). Default OFF. +_SHARED_EXPERT_LOCAL = get_bool_env_var("SGLANG_DP_SHARED_EXPERT_LOCAL") +_is_gfx95_supported = is_gfx95_supported() +_is_gfx942_supported = is_gfx942_supported() +_is_gfx1250_supported = is_gfx1250_supported() + +if _use_aiter: + if _is_gfx95_supported or _is_gfx1250_supported: + from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant + + +def _wo_a_aiter_gemm_eligible( + flag: bool, use_aiter: bool, is_hip: bool, is_gfx95: bool +) -> bool: + """Static eligibility for the aiter ``wo_a`` reroute. + + Folds the opt-in flag, the global ``SGLANG_USE_AITER`` switch, and the + HIP/gfx95 platform gates into one predicate. Evaluated once at import (see + ``_wo_a_aiter_batched_gemm_enabled``) so none of it runs on the per-token + decode critical path. + """ + return bool(flag and use_aiter and is_hip and is_gfx95) + + +# Read the opt-in flag and import the aiter kernel ONCE at module import: the +# decode ``wo_a`` matmul runs per layer/token on the critical path, so it must +# not pay an ``EnvBool.get()`` plus a function-local import on every call. If the +# path is eligible but the kernel import fails, disable it here and fall back to +# the einsum for the process (logged once) instead of retrying every step. +_wo_a_aiter_batched_gemm_enabled = _wo_a_aiter_gemm_eligible( + envs.SGLANG_OPT_USE_AITER_BATCHED_GEMM.get(), + _use_aiter, + _is_hip, + _is_gfx95_supported, +) +_wo_a_batched_gemm_bf16 = None +if _wo_a_aiter_batched_gemm_enabled: + try: + from aiter.ops.triton.gemm.batched.batched_gemm_bf16 import ( + batched_gemm_bf16 as _wo_a_batched_gemm_bf16, + ) + except Exception as err: # pragma: no cover - env-dependent + _wo_a_aiter_batched_gemm_enabled = False + logger.warning( + "aiter wo_a batched_gemm_bf16 import failed; using einsum for wo_a " + "for the rest of this process: %s", + err, + ) + +# Flipped once if the (already-imported) aiter kernel raises at runtime, so a +# per-call kernel failure falls back to the einsum for the rest of the process +# instead of re-raising (and re-logging) on every layer/token. +_wo_a_aiter_batched_gemm_disabled = False + + +def _apply_wo_a_bf16_matmul( + o: torch.Tensor, wo_a: torch.Tensor, is_decode: bool +) -> torch.Tensor: + """wo_a (attn output -> o_proj low-rank) bf16 batched matmul. + + ``o`` is ``[T, G, D]`` (tokens, groups, head_dim) and ``wo_a`` is + ``[G, R, D]`` (groups, o_lora_rank, head_dim); the result is ``[T, G, R]``. + + Dispatch contract: on the decode path, when the reroute is enabled + (``_wo_a_aiter_batched_gemm_enabled``, computed once at import) and has not + been disabled by a prior runtime failure, call the pre-imported aiter + ``batched_gemm_bf16`` (``Y[i] = X[i] @ W[i]^T``). Otherwise -- prefill, any + gate off, or after a failure -- use the numerically-equivalent + ``torch.einsum("tgd,grd->tgr", ...)``. The first runtime kernel failure + disables the reroute for the process (logged once). + """ + global _wo_a_aiter_batched_gemm_disabled + if ( + is_decode + and _wo_a_aiter_batched_gemm_enabled + and not _wo_a_aiter_batched_gemm_disabled + ): + try: + # aiter batched_gemm_bf16: XQ[B,M,K] @ WQ[B,N,K]^T -> [B,M,N]. + # Here batch = group G: XQ = o.transpose(0,1) [G,T,D], WQ = wo_a + # [G,R,D] -> [G,T,R] -> transpose back to [T,G,R]. + xq = o.transpose(0, 1).contiguous() + y = _wo_a_batched_gemm_bf16(xq, wo_a, dtype=torch.bfloat16) + return y.transpose(0, 1).contiguous() + except Exception as err: + _wo_a_aiter_batched_gemm_disabled = True + logger.warning( + "aiter wo_a batched_gemm_bf16 failed; disabling the reroute and " + "falling back to einsum for the rest of this process: %s", + err, + ) + return torch.einsum("tgd,grd->tgr", o, wo_a) + + +def _fused_rmsnorm_fp8_quant(hidden_states, weight, eps): + x_quant, x_bf16, _, _ = fused_rms_fp8_group_quant( + hidden_states, + weight, + eps, + inp2=None, + inp2_weight=None, + inp2_epsilon=None, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=True, + transpose_scale=_use_aiter_bpreshuffle_gfx95, + ) + if _use_aiter_bpreshuffle_gfx95: + x_quant = ( + x_quant[0], + view_aiter_fused_rms_transposed_fp8_scale(x_quant[1]), + ) + return x_quant, x_bf16 + + +def make_hc_mixing_params( + hc_mult: int, hidden_size: int +) -> Tuple[ + nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter, nn.Parameter +]: + mix_hc = (2 + hc_mult) * hc_mult + hc_dim = hc_mult * hidden_size + return ( + nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)), + nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)), + nn.Parameter(torch.empty(3, dtype=torch.float32)), + nn.Parameter(torch.empty(3, dtype=torch.float32)), + ) + + +def make_hc_head_params( + hc_mult: int, hidden_size: int +) -> Tuple[nn.Parameter, nn.Parameter, nn.Parameter]: + hc_dim = hc_mult * hidden_size + return ( + nn.Parameter(torch.empty(hc_mult, hc_dim, dtype=torch.float32)), + nn.Parameter(torch.empty(hc_mult, dtype=torch.float32)), + nn.Parameter(torch.empty(1, dtype=torch.float32)), + ) + + +def hc_head_torch( + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + *, + norm_eps: float, + hc_eps: float, +) -> torch.Tensor: + shape, dtype = x.size(), x.dtype + x = x.flatten(-2).float() + rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + norm_eps) + mixes = F.linear(x, hc_fn) * rsqrt + pre = torch.sigmoid(mixes * hc_scale + hc_base) + hc_eps + y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=-2) + return y.to(dtype) + + +_FREQS_CIS_TO_COS_SIN: dict[ + Tuple[int, torch.dtype, torch.device], Tuple[torch.Tensor, torch.Tensor] +] = {} + + +def _freqs_cis_to_cos_sin( + freqs_cis: torch.Tensor, dtype: torch.dtype, device: torch.device +) -> Tuple[torch.Tensor, torch.Tensor]: + """Derive (cos, sin) bf16 contiguous tables from a complex64 `freqs_cis`, + cached by `(id(freqs_cis), dtype, device)` so that all layers sharing the + same `freqs_cis` (via `precompute_freqs_cis`'s lru_cache) reuse one pair.""" + key = (id(freqs_cis), dtype, device) + cached = _FREQS_CIS_TO_COS_SIN.get(key) + if cached is not None: + return cached + fr = torch.view_as_real(freqs_cis) + cos = fr[..., 0].to(device=device, dtype=dtype).contiguous() + sin = fr[..., 1].to(device=device, dtype=dtype).contiguous() + _FREQS_CIS_TO_COS_SIN[key] = (cos, sin) + return cos, sin + + +def _apply_gguf_grouped_wo_a( + o: torch.Tensor, + qweight: torch.Tensor, + qweight_type: int, + o_lora_rank: int, + matmul_fn: Optional[Callable] = None, +) -> torch.Tensor: + if matmul_fn is None: + from sglang.srt.layers.quantization.gguf import fused_mul_mat_gguf + + matmul_fn = fused_mul_mat_gguf + + group_outputs = [] + for group_id in range(o.shape[1]): + start = group_id * o_lora_rank + group_outputs.append( + matmul_fn( + o[:, group_id, :].contiguous(), + qweight[start : start + o_lora_rank], + qweight_type, + ) + ) + return torch.stack(group_outputs, dim=1) + + +if TYPE_CHECKING: + from sglang.srt.layers.attention.deepseek_v4_backend import ( + DeepseekV4AttnBackend, + ) + from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + ) + from sglang.srt.layers.quantization import QuantizationConfig + from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +@register_custom_op(mutates_args=["output"]) +@register_split_op() +def deepseek_v4_attention_with_output( + query: torch.Tensor, + key_value: torch.Tensor, + output: torch.Tensor, + layer_id: int, + compress_ratio: int, + attn_sink: torch.Tensor, + save_kv_cache: bool, +) -> None: + context = get_tc_piecewise_forward_context() + forward_batch = context.forward_batch + attention_layers = context.attention_layers + attention_layer = attention_layers[layer_id] + real_num_tokens = forward_batch.num_token_non_padded_cpu + + if real_num_tokens == 0: + output.zero_() + return + + query = query[:real_num_tokens] + key_value = key_value[:real_num_tokens] + + original_out_cache_loc = forward_batch.out_cache_loc + forward_batch.out_cache_loc = original_out_cache_loc[:real_num_tokens] + + attn_backend = get_attn_backend() + try: + ret = attn_backend.forward( + q=query, + k=key_value, + v=key_value, + layer=attention_layer, + forward_batch=forward_batch, + compress_ratio=compress_ratio, + attn_sink=attn_sink, + save_kv_cache=save_kv_cache, + ) + finally: + forward_batch.out_cache_loc = original_out_cache_loc + + assert ( + output[:real_num_tokens].numel() == ret.numel() + ), f"Output tensor element mismatch: {output[:real_num_tokens].numel()} != {ret.numel()}" + + output[:real_num_tokens].view(ret.shape).copy_(ret) + return + + +bcg_deepseek_v4_attention_with_output = eager_on_graph(True)( + deepseek_v4_attention_with_output +) + + +class MqaAttentionBase(nn.Module): + + def __init__( + self, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig], + prefix: str, + *, + attn_tp_rank: Optional[int] = None, + attn_tp_size: Optional[int] = None, + compress_ratio: Optional[int] = None, + fuse_wqa_wkv: Optional[bool] = None, + wo_a_fp8: Optional[bool] = None, + wo_a_keeps_quant_config: Optional[bool] = None, + wo_b_reduce_results: Optional[bool] = None, + rope_original_seq_len: Optional[int] = None, + ) -> None: + super().__init__() + self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + if attn_tp_rank is None or attn_tp_size is None: + attn_tp_rank = get_parallel().attn_tp_rank + attn_tp_size = get_parallel().attn_tp_size + if self.dsa_enable_prefill_cp: + self.cp_size = get_parallel().attn_cp_size + attn_tp_rank, attn_tp_size = 0, 1 + self.attn_tp_rank: int = attn_tp_rank + self.attn_tp_size: int = attn_tp_size + + self.layer_id = layer_id + self.dim = config.hidden_size + self.hidden_size = config.hidden_size + self.qk_rope_head_dim = config.qk_rope_head_dim + self.qk_nope_head_dim = config.head_dim - config.qk_rope_head_dim + self.head_dim = self.qk_rope_head_dim + self.qk_nope_head_dim + self.rope_head_dim = config.qk_rope_head_dim + self.n_heads = config.num_attention_heads + self.n_local_heads = self.n_heads // self.attn_tp_size + self.n_groups = config.o_groups + self.n_local_groups = self.n_groups // self.attn_tp_size + self.q_lora_rank = config.q_lora_rank + self.o_lora_rank = config.o_lora_rank + self.eps = config.rms_norm_eps + self.softmax_scale = self.head_dim**-0.5 + + self.compress_ratio: int = ( + compress_ratio + if compress_ratio is not None + else config.compress_ratios[layer_id] + ) + assert self.compress_ratio in ( + 0, + 4, + 128, + ), f"V4 compress_ratio: expected one of (0, 4, 128), got {self.compress_ratio}" + + assert self.head_dim == config.head_dim + assert config.num_key_value_heads == 1 + + fuse: bool = ( + envs.SGLANG_OPT_FUSE_WQA_WKV.get() if fuse_wqa_wkv is None else fuse_wqa_wkv + ) + fp8: bool = _FP8_WO_A_GEMM if wo_a_fp8 is None else wo_a_fp8 + reduce_results: bool = ( + (self.attn_tp_size == get_parallel().tp_size and self.attn_tp_size > 1) + if wo_b_reduce_results is None + else wo_b_reduce_results + ) + if wo_a_keeps_quant_config is None: + keep_source_quant = ( + quant_config is not None and quant_config.get_name() == "expert_pack" + ) + wo_a_quant_config: Optional[QuantizationConfig] = ( + quant_config if fp8 or keep_source_quant else None + ) + elif wo_a_keeps_quant_config: + wo_a_quant_config = quant_config + else: + wo_a_quant_config = None + + self.fuse_wqa_wkv = fuse + + self.attn_sink = nn.Parameter(torch.empty(self.n_heads, dtype=torch.float32)) + self._attn_sink_local: Optional[torch.Tensor] = None + if fuse: + self.wqkv_a = ReplicatedLinear( + self.hidden_size, + self.q_lora_rank + self.head_dim, + bias=False, + quant_config=quant_config, + prefix=add_prefix("wqkv_a", prefix), + ) + else: + self.wq_a = ReplicatedLinear( + self.hidden_size, + self.q_lora_rank, + bias=False, + quant_config=quant_config, + prefix=add_prefix("wq_a", prefix), + ) + self.wkv = ReplicatedLinear( + self.hidden_size, + self.head_dim, + bias=False, + quant_config=quant_config, + prefix=add_prefix("wkv", prefix), + ) + self.q_norm = RMSNorm(self.q_lora_rank, eps=self.eps) + self.wq_b = ColumnParallelLinear( + self.q_lora_rank, + self.n_heads * self.head_dim, + bias=False, + quant_config=quant_config, + prefix=add_prefix("wq_b", prefix), + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + ) + self.kv_norm = RMSNorm(self.head_dim, eps=self.eps) + self.wo_a = ColumnParallelLinear( + self.n_heads * self.head_dim // self.n_groups, + self.n_groups * self.o_lora_rank, + bias=False, + quant_config=wo_a_quant_config, + prefix=add_prefix("wo_a", prefix), + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + **({} if fp8 else {"params_dtype": torch.bfloat16}), + ) + if fp8: + from sglang.srt.layers import deep_gemm_wrapper + + assert hasattr( + self.wo_a, "weight_scale_inv" + ), "FP8 quant_config must create weight_scale_inv" + self.wo_a.weight_scale_inv.format_ue8m0 = ( + deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 + ) + self.wo_b = RowParallelLinear( + self.n_groups * self.o_lora_rank, + self.hidden_size, + bias=False, + quant_config=quant_config, + reduce_results=reduce_results, + prefix=add_prefix("wo_b", prefix), + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + ) + + from sglang.kernels.ops.attention.deepseek_v4_rope import precompute_freqs_cis + + rope_theta, rope_scaling = get_rope_config(config) + self.rope_scaling = dict(rope_scaling) if rope_scaling else None + scaling = self.rope_scaling or {} + + # RoPE is selected at layer granularity in the reference model. Pure + # SWA layers use the main unscaled RoPE, while C4/C128 layers use the + # compressed YaRN RoPE for Q, their SWA branch, and compressed KV. + self.rope_base = ( + config.compress_rope_theta if self.compress_ratio else rope_theta + ) + original_seq_len: int = ( + rope_original_seq_len + if rope_original_seq_len is not None + else ( + scaling["original_max_position_embeddings"] + if self.compress_ratio + else 0 + ) + ) + freqs_cis = precompute_freqs_cis( + dim=self.qk_rope_head_dim, + seqlen=config.max_position_embeddings, + original_seq_len=original_seq_len, + base=self.rope_base, + factor=scaling.get("factor", 1.0), + beta_fast=scaling.get("beta_fast", 32), + beta_slow=scaling.get("beta_slow", 1), + ) + self.register_buffer("freqs_cis", freqs_cis, persistent=False) + self.freqs_cis: torch.Tensor + + def _padded_attn_heads(self) -> int: + # FlashInfer SM120 specializes native TP-local heads; the generic + # FlashMLA path still requires padding to 64 or 128. + if ( + get_platform().is_sm120 + and envs.SGLANG_SM120_FLASHMLA_BACKEND.get() == "flashinfer" + and self.n_local_heads in (8, 16, 32, 64, 128) + ): + return self.n_local_heads + return 64 if self.n_local_heads <= 64 else self.n_heads + + def _local_attn_sink(self) -> torch.Tensor: + if self.attn_tp_size == 1: + return self.attn_sink + if self._attn_sink_local is None: + rank = self.attn_tp_rank + num_heads = self.n_local_heads + padded_num_heads = self._padded_attn_heads() + sink = self.attn_sink.new_zeros(padded_num_heads) + sink[:num_heads] = self.attn_sink[rank * num_heads : (rank + 1) * num_heads] + self._attn_sink_local = sink + return self._attn_sink_local + + @contextmanager + def maybe_use_decode_attn_tp(self, forward_batch: ForwardBatch): + ctx = get_cp_decode_attn_tp_ctx() + attn = self.attn_mqa if isinstance(self, MQALayer) else self.attn + with ctx.maybe_use_decode_attn_tp( + forward_batch, + [self.wq_b, self.wo_a, self.wo_b], + radix_attn=attn, + ): + if ctx.use_decode_attn_tp: + orig = ( + self.n_local_heads, + self.n_local_groups, + self.attn_tp_rank, + self.attn_tp_size, + ) + decode_tp_size = ctx.decode_tp_size + self.n_local_heads = self.n_heads // decode_tp_size + self.n_local_groups = self.n_groups // decode_tp_size + self.attn_tp_rank = ctx.decode_tp_rank + self.attn_tp_size = decode_tp_size + try: + yield + finally: + ( + self.n_local_heads, + self.n_local_groups, + self.attn_tp_rank, + self.attn_tp_size, + ) = orig + else: + yield + + +class MQALayer(MqaAttentionBase): + def __init__( + self, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_streams: Optional[List[torch.cuda.Stream]] = None, + compress_ratio_override: Optional[int] = None, + ) -> None: + super().__init__( + config, + layer_id, + quant_config, + prefix, + compress_ratio=compress_ratio_override, + ) + + active_rope_scaling = None + if self.compress_ratio in (4, 128): + active_rope_scaling = dict(self.rope_scaling or {}) + active_rope_scaling["rope_type"] = "deepseek_yarn" + self.rotary_emb = get_rope_wrapper( + head_size=self.rope_head_dim, + rotary_dim=self.rope_head_dim, + max_position=config.max_position_embeddings, + base=self.rope_base, + rope_scaling=active_rope_scaling, + is_neox_style=False, + device=get_device().device, + ) + + if _is_npu: + Dsv4NpuRoPE.for_freqs( + self.freqs_cis, getattr(self, "rotary_emb", None) + ).ensure_tables(torch.float32) + + if _is_hip: + cos_cache = ( + self.freqs_cis.real.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) + ) + sin_cache = ( + self.freqs_cis.imag.to(torch.bfloat16).unsqueeze(-2).unsqueeze(-2) + ) + self.register_buffer("cos_cache", cos_cache, persistent=False) + self.register_buffer("sin_cache", sin_cache, persistent=False) + + if alt_streams is not None and ( + (_is_cuda and envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get()) + or (_is_npu and envs.SGLANG_NPU_USE_MULTI_STREAM.get()) + ): + self.alt_streams = alt_streams[:3] + self.alt_streams_indexer = alt_streams[-2:] + else: + self.alt_streams = None + self.alt_streams_indexer = None + + self._multi_stream_bs_limit = 128 if get_platform().is_blackwell else 64 + + self.compressor = None + self.indexer = None + if self.compress_ratio in (4, 128): + expert_pack_quant_config = ( + quant_config + if quant_config is not None and quant_config.get_name() == "expert_pack" + else None + ) + self.compressor = Compressor( + config, + layer_id=self.layer_id, + is_in_indexer=False, + freqs_cis=self.freqs_cis, + compress_ratio=self.compress_ratio, + head_dim=self.head_dim, + rotate=False, + prefix=add_prefix("compressor", prefix), + quant_config=expert_pack_quant_config, + rotary_emb=self.rotary_emb, + ) + if self.compress_ratio == 4: + self.indexer = C4Indexer( + config, + freqs_cis=self.freqs_cis, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("indexer", prefix), + alt_streams=self.alt_streams_indexer, + rotary_emb=self.rotary_emb, + ) + + self.attn_mqa = RadixAttention( + self.n_local_heads, + self.head_dim, + self.softmax_scale, + num_kv_heads=1, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("attn_mqa", prefix), + ) + + self.use_fused_qk_norm_rope = ( + _is_hip and envs.SGLANG_OPT_USE_FUSED_QK_NORM_ROPE.get() + ) + + # KV cache write is always fused into the K kernel + # (`_compute_kv_to_cache`), so the legacy "overlap store cache" flag + # has no effect here -- the fused path is on by default. + + def _get_npu_rope_position_cache( + self, positions: torch.Tensor, dtype: torch.dtype, inverse: bool = False + ) -> Tuple[torch.Tensor, torch.Tensor]: + # ``rotary_emb`` is shared by layers with the same RoPE configuration and + # can also be shared by the target and NextN models. Only cache the + # immutable full table on it. A position-gathered tensor is specific to + # this forward and reusing it based on shape alone gives MTP decode the + # previous step's RoPE values when positions change but batch size does not. + return Dsv4NpuRoPE.for_freqs( + self.freqs_cis, getattr(self, "rotary_emb", None) + ).get_cos_sin( + positions, + dtype, + view_4d=True, + inverse=inverse, + allow_build=False, + cache_dtype=torch.float32, + ) + + def _compute_q_a( + self, + x: torch.Tensor, + qkv_a: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if qkv_a is not None: + q = qkv_a[..., : self.q_lora_rank] + else: + q, _ = self.wq_a(x) + return self.q_norm(q) + + def _compute_q_b( + self, + q: torch.Tensor, + positions: torch.Tensor, + q_out: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + q, _ = self.wq_b(q) + q = q.view(-1, self.n_local_heads, self.head_dim) + if q_out is None: + q_out = torch.empty_like(q) + # Fused warp-per-(token, head) rmsnorm-self + RoPE + write to q_out. + fused_q_norm_rope(q, q_out, self.eps, self.freqs_cis, positions) + return q_out + + def _compute_kv_to_cache( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + attn_backend, + qkv_a: Optional[torch.Tensor] = None, + ) -> None: + """Fused: rmsnorm + RoPE + write directly to FlashMLA paged cache. + + Replaces the bf16-kv-intermediate path. Used everywhere except the DSA + prefill-CP case (which needs bf16 kv for the cross-rank all-gather). + """ + if envs.SGLANG_DSV4_USE_BF16_KV_QUANT_SOURCE.get(): + # Quantize the nope payload from bf16-rounded values (the fused + # kernel quantizes from fp32 registers; the bf16 rounding moves + # values across fp8 bins relative to bf16-sourced consumers). + kv = self._compute_kv_bf16(x, positions, qkv_a=qkv_a) + attn_backend.store_cache( + layer_id=self.layer_id, swa_k=kv, forward_batch=forward_batch + ) + return + if qkv_a is not None: + kv = qkv_a[..., self.q_lora_rank :] + else: + kv, _ = self.wkv(x) + token_to_kv_pool = get_token_to_kv_pool() + if TYPE_CHECKING: + assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool) + token_to_kv_pool.set_swa_key_buffer_radix_fused_norm_rope( + layer_id=self.layer_id, + swa_loc=attn_backend.get_swa_out_cache_loc(forward_batch), + kv=kv, + kv_weight=self.kv_norm.weight.data, + eps=self.eps, + freqs_cis=self.freqs_cis, + positions=positions, + ) + + def _compute_kv_bf16( + self, + x: torch.Tensor, + positions: torch.Tensor, + qkv_a: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Bf16-kv path used by the DSA prefill-CP case (needs all-gather).""" + if qkv_a is not None: + kv = qkv_a[..., self.q_lora_rank :] + else: + kv, _ = self.wkv(x) + kv = kv.contiguous() + fused_norm_rope_inplace( + kv, + self.kv_norm.weight.data, + self.eps, + self.freqs_cis, + positions, + ) + return kv + + def _forward_prepare_multi_stream( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + attn_backend, + q_out: Optional[torch.Tensor] = None, + x_quant=None, + ) -> torch.Tensor: + assert self.alt_streams is not None + assert len(self.alt_streams) >= 3 + + current_stream = torch.cuda.current_stream() + stream_kv = self.alt_streams[0] + stream_compressor = self.alt_streams[1] + stream_indexer = self.alt_streams[2] + + stream_kv.wait_stream(current_stream) + stream_compressor.wait_stream(current_stream) + stream_indexer.wait_stream(current_stream) + + x_linear = x_quant if x_quant is not None else x + qkv_a: Optional[torch.Tensor] = None + qkv_a_ready: Optional[torch.cuda.Event] = None + if self.fuse_wqa_wkv: + qkv_a, _ = self.wqkv_a(x_linear) + qkv_a_ready = current_stream.record_event() + + q_lora = self._compute_q_a(x_linear, qkv_a=qkv_a) + q_lora_ready = current_stream.record_event() + + if self.indexer is not None: + with torch.cuda.stream(stream_indexer): + self.indexer( + x=x, + q_lora=q_lora, + forward_batch=forward_batch, + attn_backend=attn_backend, + enable_multi_stream=True, + q_lora_ready=q_lora_ready, + ) + + with torch.cuda.stream(stream_kv): + if qkv_a_ready is not None: + stream_kv.wait_event(qkv_a_ready) + # Fused norm + rope + cache write -- no bf16 KV intermediate. + self._compute_kv_to_cache( + x_linear, positions, forward_batch, attn_backend, qkv_a=qkv_a + ) + + if self.compressor is not None: + with torch.cuda.stream(stream_compressor): + attn_backend.forward_core_compressor( + x, forward_batch, self.layer_id, self.compressor + ) + + q = self._compute_q_b(q_lora, positions, q_out) + current_stream.wait_stream(stream_kv) + current_stream.wait_stream(stream_compressor) + current_stream.wait_stream(stream_indexer) + del qkv_a + + return q + + def _forward_prepare_multi_stream_npu( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + attn_backend, + q_out: Optional[torch.Tensor] = None, + x_quant=None, + ) -> torch.Tensor: + # NPU multi-stream: KV on stream_kv, Q on stream_q, overlapped with + # indexer/compressor on current. rope is split; the kv-only call passes + # kv.unsqueeze(1) as q_rope so the op sees [T,1,1,head_dim] like the + # fused path. + assert self.alt_streams is not None + current_stream = torch.npu.current_stream() + stream_kv = self.alt_streams[0] + stream_q = self.alt_streams[1] + stream_kv.wait_stream(current_stream) + stream_q.wait_stream(current_stream) + + x_linear = x_quant if x_quant is not None else x + qkv_a: Optional[torch.Tensor] = None + qkv_a_ready = None + if self.fuse_wqa_wkv: + qkv_a, _ = self.wqkv_a(x_linear) + qkv_a_ready = current_stream.record_event() + if qkv_a is not None: + q_lora = qkv_a[..., : self.q_lora_rank] + else: + q_lora, _ = self.wq_a(x_linear) + q_lora = self.q_norm(q_lora) + q_lora_ready = current_stream.record_event() + + # KV block on stream_kv. + with torch.npu.stream(stream_kv): + if qkv_a_ready is not None: + stream_kv.wait_event(qkv_a_ready) + if qkv_a is not None: + kv = qkv_a[..., self.q_lora_rank :] + else: + kv, _ = self.wkv(x) + kv = self.kv_norm(kv) + cos4_k, sin4_k = self._get_npu_rope_position_cache( + positions, kv.dtype, inverse=False + ) + Dsv4NpuRoPE.apply_rotary_mul_inplace( + kv.unsqueeze(1), + None, + cos4_k, + sin4_k, + qk_nope_dim=self.qk_nope_head_dim, + ) + attn_backend.store_cache( + layer_id=self.layer_id, + swa_k=kv, + forward_batch=forward_batch, + ) + + # Q block on stream_q (needs only q_lora). + with torch.npu.stream(stream_q): + stream_q.wait_event(q_lora_ready) + q, _ = self.wq_b(q_lora) + q = q.view(-1, self.n_local_heads, self.head_dim) + _dummy = q.new_ones(q.shape[-1]) + q = torch_npu.npu_rms_norm(q, _dummy, self.eps)[0] + cos4_q, sin4_q = self._get_npu_rope_position_cache( + positions, q.dtype, inverse=False + ) + Dsv4NpuRoPE.apply_rotary_mul_inplace( + q, + None, + cos4_q, + sin4_q, + qk_nope_dim=self.qk_nope_head_dim, + ) + if q_out is not None: + q_out.copy_(q) + q.record_stream(stream_q) + + # Indexer + compressor: serial on current. + if self.indexer is not None: + self.indexer( + x=x, + q_lora=q_lora, + forward_batch=forward_batch, + attn_backend=attn_backend, + ) + if self.compressor is not None: + attn_backend.forward_core_compressor( + x, + forward_batch, + self.layer_id, + self.compressor, + ) + + # Join stream_kv + stream_q before downstream attention. + current_stream.wait_stream(stream_kv) + current_stream.wait_stream(stream_q) + del qkv_a + return q + + def _forward_prepare_multi_stream_hip( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + attn_backend, + q_out: Optional[torch.Tensor] = None, + x_quant=None, + ) -> torch.Tensor: + """ATOM-style ROCm path: overlap compressors, keep Q/KV on main stream.""" + assert self.alt_streams is not None + assert len(self.alt_streams) >= 1 + + current_stream = torch.cuda.current_stream() + stream_compressor = self.alt_streams[0] + stream_indexer_compressor = ( + self.alt_streams[1] if len(self.alt_streams) > 1 else None + ) + + if self.compressor is not None: + stream_compressor.wait_stream(current_stream) + with torch.cuda.stream(stream_compressor): + attn_backend.forward_core_compressor( + x, forward_batch, self.layer_id, self.compressor + ) + + if self.indexer is not None and stream_indexer_compressor is not None: + stream_indexer_compressor.wait_stream(current_stream) + with torch.cuda.stream(stream_indexer_compressor): + attn_backend.forward_indexer_compressor( + x=x, + forward_batch=forward_batch, + layer_id=self.indexer.layer_id, + compressor=self.indexer.compressor, + ) + + x_linear = x_quant if x_quant is not None else x + if self.fuse_wqa_wkv: + qkv_a, _ = self.wqkv_a(x_linear) + q_lora = qkv_a[..., : self.q_lora_rank] + else: + q_lora, _ = self.wq_a(x_linear) + qkv_a = None + + if self.use_fused_qk_norm_rope: + if _is_gfx95_supported or _is_gfx1250_supported: + q_for_wqb, q_lora = _fused_rmsnorm_fp8_quant( + q_lora, + self.q_norm.weight, + self.q_norm.variance_epsilon, + ) + q, _ = self.wq_b(q_for_wqb) + else: + q_lora = self.q_norm(q_lora) + q, _ = self.wq_b(q_lora) + + kv = ( + qkv_a[..., self.q_lora_rank :] + if qkv_a is not None + else self.wkv(x_linear)[0] + ) + + from sglang.kernels.ops.attention.fused_qk_norm_rope_store import ( + fused_qk_norm_rope_swa_store, + ) + + token_to_kv_pool = get_token_to_kv_pool() + swa_loc = attn_backend.get_swa_out_cache_loc(forward_batch) + swa_cache = token_to_kv_pool.get_swa_raw_buffer(self.layer_id) + swa_page_size = token_to_kv_pool.swa_kv_pool.page_size + + q = fused_qk_norm_rope_swa_store( + q=q, + kv=kv, + q_norm_weight=None, + kv_norm_weight=self.kv_norm.weight, + q_rms_eps=self.eps, + kv_rms_eps=self.eps, + rope_head_dim=self.qk_rope_head_dim, + cos_cache=self.cos_cache, + sin_cache=self.sin_cache, + positions=positions, + swa_cache=swa_cache, + swa_loc=swa_loc, + swa_page_size=swa_page_size, + q_out=q_out, + dtype=x.dtype, + ) + else: + q_lora = self.q_norm(q_lora) + q = self._compute_q_b(q_lora, positions, q_out) + self._compute_kv_to_cache( + x_linear, positions, forward_batch, attn_backend, qkv_a=qkv_a + ) + + del qkv_a + + if self.indexer is not None: + current_stream.wait_stream(stream_compressor) + if stream_indexer_compressor is not None: + current_stream.wait_stream(stream_indexer_compressor) + self.indexer( + x=x, + q_lora=q_lora, + forward_batch=forward_batch, + attn_backend=attn_backend, + skip_compressor=True, + ) + elif self.compressor is not None: + current_stream.wait_stream(stream_compressor) + + return q + + def _forward_prepare( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + attn_backend, + q_out: Optional[torch.Tensor] = None, + x_quant=None, + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + x_linear = x_quant if x_quant is not None else x + # kv_score depends only on x, so its CP all-gather can start before the + # projections and be collected inside forward_core_compressor below -- + # the projections are what hides it. No-op unless the CP+TBO path armed + # _cp_prefetch_comm_stream. + if _is_hip and self.compressor is not None: + self.compressor.prelaunch_kv_score(x, forward_batch) + if self.indexer is not None: + self.indexer.compressor.prelaunch_kv_score(x, forward_batch) + + if self.fuse_wqa_wkv: + qkv_a, _ = self.wqkv_a(x_linear) + q_lora = qkv_a[..., : self.q_lora_rank] + else: + q_lora, _ = self.wq_a(x_linear) + qkv_a = None + + use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) + kv: Optional[torch.Tensor] + kv_handle = None + + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( + is_unified_kv_triton, + ) + + unified = is_unified_kv_triton() + is_decode = forward_batch.forward_mode.is_decode_or_idle() + # The kernel is token-indexed (q, kv and positions are all length M), so + # a verify batch carrying several draft tokens per request is a shape it + # already handles. Only the cache store differs between decode and + # verify, and that half is left off below. + fuse_verify = ( + envs.SGLANG_OPT_FUSED_QK_NORM_ROPE_VERIFY.get() + and forward_batch.forward_mode.is_target_verify() + ) + do_fused_qk_norm_rope = (unified and (is_decode or fuse_verify)) or ( + not unified and self.use_fused_qk_norm_rope + ) + + if do_fused_qk_norm_rope: + if _is_gfx95_supported or _is_gfx1250_supported: + q_for_wqb, q_lora = _fused_rmsnorm_fp8_quant( + q_lora, + self.q_norm.weight, + self.q_norm.variance_epsilon, + ) + q, _ = self.wq_b(q_for_wqb) + else: + q_lora = self.q_norm(q_lora) + q, _ = self.wq_b(q_lora) + + kv = ( + qkv_a[..., self.q_lora_rank :] + if qkv_a is not None + else self.wkv(x_linear)[0] + ) + + token_to_kv_pool = get_token_to_kv_pool() + if unified and fuse_verify: + # Target-verify runs through the unified_kv decode path. The + # backend writes the current chunk's KV into the ring *before* + # attention (save_kv_cache=True -> store_swa_into_unified ahead + # of runtime.decode), and per-token causal index streams -- built + # once per step in the backend metadata -- keep each draft query + # attending only to positions up to itself. Causal masking among + # the draft tokens comes from those index streams, not from store + # timing. So this path skips only the fused kernel's *own* store + # and returns kv, letting that existing causally-indexed backend + # store run unchanged; we fuse just the norm+RoPE. swa_loc is not + # computed -- it only addresses the kernel store this path drops. + # + # kv is a strided slice of qkv_a and the ring store requires a + # contiguous buffer, so materialise it before the kernel norms + # it in place. The unfused path pays the same copy inside + # _compute_kv_bf16. + kv = kv.contiguous() + swa_cache, swa_loc = None, None + swa_page_size, bf16_store = 1, True + elif unified: + swa_cache = token_to_kv_pool.get_unified_kv(self.layer_id) + # swa_loc is layer-independent; computed once per forward by the + # backend and cached on the metadata (read here by every layer). + swa_loc = attn_backend.get_unified_swa_loc(forward_batch) + swa_page_size, bf16_store = 1, True + else: + swa_cache = token_to_kv_pool.get_swa_raw_buffer(self.layer_id) + swa_loc = attn_backend.get_swa_out_cache_loc(forward_batch) + swa_page_size, bf16_store = ( + token_to_kv_pool.swa_kv_pool.page_size, + False, + ) + + from sglang.kernels.ops.attention.fused_qk_norm_rope_store import ( + fused_qk_norm_rope_swa_store, + ) + + q = fused_qk_norm_rope_swa_store( + q=q, + kv=kv, + q_norm_weight=None, + kv_norm_weight=self.kv_norm.weight, + q_rms_eps=self.eps, + kv_rms_eps=self.eps, + rope_head_dim=self.qk_rope_head_dim, + cos_cache=self.cos_cache, + sin_cache=self.sin_cache, + positions=positions, + swa_cache=swa_cache, + swa_loc=swa_loc, + swa_page_size=swa_page_size, + q_out=q_out, + dtype=x.dtype, + bf16_store=bf16_store, + ) + # On the verify path the kernel normed + RoPE'd kv in place and wrote + # nothing, so hand it back: the caller feeds it to attention as the + # current chunk (attn_k = kv) and save_kv_cache = kv is not None lets + # the backend do its normal causally-indexed store into the ring + # before the decode kernel runs -- exactly as the unfused path did. + if not (unified and fuse_verify): + kv = None + + if not unified and use_cp: + # DSA CP: keep bf16 kv around for the cross-rank all-gather, then + # write to the FlashMLA cache after gather. + kv = self._compute_kv_bf16(x, positions, qkv_a=qkv_a) + kv = cp_materialize_global_token_order( + kv.contiguous(), + forward_batch, + torch.cuda.current_stream(), + ) + elif _is_npu: + q_lora = self.q_norm(q_lora) + q, _ = self.wq_b(q_lora) + q = q.view(-1, self.n_local_heads, self.head_dim) + _dummy = q.new_ones(q.shape[-1]) + q = torch_npu.npu_rms_norm(q, _dummy, self.eps)[0] + + if qkv_a is not None: + kv = qkv_a[..., self.q_lora_rank :] + else: + kv, _ = self.wkv(x) + kv = self.kv_norm(kv) + + cos4, sin4 = self._get_npu_rope_position_cache( + positions, q.dtype, inverse=False + ) + Dsv4NpuRoPE.apply_rotary_mul_inplace( + q, + kv.unsqueeze(1), + cos4, + sin4, + qk_nope_dim=self.qk_nope_head_dim, + ) + attn_backend.store_cache( + layer_id=self.layer_id, + swa_k=kv, + forward_batch=forward_batch, + ) + kv = None + if q_out is not None: + q_out.copy_(q) + else: + q_lora = self.q_norm(q_lora) + q = self._compute_q_b(q_lora, positions, q_out) + if unified: + # unified_kv prefill: keep bf16 kv; the backend writes + # the ring AFTER attention (2-source path). + kv = self._compute_kv_bf16(x_linear, positions, qkv_a=qkv_a) + # HIP/ROCm-only: the unified_kv 2-source prefill path is exclusive + # to DeepseekV4HipRadixBackend. Guard with _is_hip so this CP + # all-gather never enters the NVIDIA (DeepseekV4AttnBackend) path. + if use_cp and _is_hip: + # unified_kv + DSA CP: the 2-source prefill path needs the + # FULL current-chunk KV (extend source + ring write), so + # all-gather the per-rank bf16 KV across the CP group. + comm_stream = getattr( + forward_batch, "_cp_prefetch_comm_stream", None + ) + if comm_stream is not None: + # kv is not read again until this function returns, so the + # indexer + compressor below can run while it gathers. + kv_handle = cp_all_gather_rerange_launch( + kv, self.cp_size, comm_stream, ("kv", self.layer_id) + ) + kv = None + else: + kv = cp_materialize_global_token_order( + kv.contiguous(), + forward_batch, + torch.cuda.current_stream(), + ) + elif use_cp: + # NSA CP: keep bf16 kv around for the cross-rank all-gather, then + # write to the FlashMLA cache after gather. + kv = self._compute_kv_bf16(x_linear, positions, qkv_a=qkv_a) + kv = cp_materialize_global_token_order( + kv.contiguous(), + forward_batch, + torch.cuda.current_stream(), + ) + attn_backend.store_cache( + layer_id=self.layer_id, + swa_k=kv, + forward_batch=forward_batch, + ) + else: + self._compute_kv_to_cache( + x_linear, positions, forward_batch, attn_backend, qkv_a=qkv_a + ) + kv = None + + del qkv_a + + if self.indexer is not None: + self.indexer( + x=x, + q_lora=q_lora, + forward_batch=forward_batch, + attn_backend=attn_backend, + ) + if self.compressor is not None: + attn_backend.forward_core_compressor( + x, + forward_batch, + self.layer_id, + self.compressor, + ) + + if _is_hip and kv_handle is not None: + kv = cp_all_gather_rerange_finish(kv_handle) + + return q, kv + + def forward( + self, + x: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + x_quant=None, + ) -> torch.Tensor: + if not get_attn_tp_context().input_scattered and x.shape[0] == 0: + return x + + attn_backend = get_attn_backend() + if TYPE_CHECKING: + assert isinstance( + attn_backend, + (DeepseekV4AttnBackend, DeepseekV4HipRadixBackend), + ) + + enable_multi_stream = ( + envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() + and self.alt_streams is not None + and get_is_capture_mode() + and ( + is_in_breakable_cuda_graph() + or x.shape[0] <= self._multi_stream_bs_limit + ) + and not (self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch)) + and not (_is_hip and self.compressor is None) + ) or ( + _is_npu + and envs.SGLANG_NPU_USE_MULTI_STREAM.get() + and self.alt_streams is not None + and x.shape[0] <= self._multi_stream_bs_limit + and not forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed() + ) + + tp_slice, q_padded, q_out = slice(None), None, None + if self.attn_tp_size > 1: + # Match the query allocation and sliced sink to this backend. + # FlashInfer SM120 uses native heads; generic FlashMLA keeps padding. + padded_num_heads = self._padded_attn_heads() + # Only [0:n_local_heads] is written below. Uninitialized padded TP + # heads inject NaN into attention on gfx942 (fnuz), so zero-init + # there; other archs tolerate new_empty and skip the per-forward + # memset. + if _is_gfx942_supported: + q_padded = x.new_zeros(x.shape[0], padded_num_heads, self.head_dim) + else: + q_padded = x.new_empty(x.shape[0], padded_num_heads, self.head_dim) + tp_slice = slice(0, self.n_local_heads) + q_out = q_padded[:, tp_slice, :] + attn_sink = self._local_attn_sink() + + if enable_multi_stream: + # Multi-stream path always fuses cache write into the K kernel, + # so the bf16 KV intermediate is gone. + if _is_hip: + q = self._forward_prepare_multi_stream_hip( + x, + positions, + forward_batch, + attn_backend, + q_out, + x_quant=x_quant, + ) + elif _is_npu: + q = self._forward_prepare_multi_stream_npu( + x, + positions, + forward_batch, + attn_backend, + q_out, + x_quant=x_quant, + ) + else: + q = self._forward_prepare_multi_stream( + x, + positions, + forward_batch, + attn_backend, + q_out, + x_quant=x_quant, + ) + kv = None + else: + q, kv = self._forward_prepare( + x, + positions, + forward_batch, + attn_backend, + q_out, + x_quant=x_quant, + ) + + # save_kv_cache = kv is not None selects who writes the ring. When kv is + # None the store was already fused into _forward_prepare* (decode) or + # done inline, so the backend skips its own store_cache; pass `q` as a + # sentinel for the `k is v` assert (attention won't read it once + # save_kv_cache=False). When kv is not None (target-verify, or DSA-CP), + # _forward_prepare* deliberately left the store off and the backend does + # its normal causally-indexed store from attn_k = kv. + attn_k = kv if kv is not None else q + from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( + is_unified_kv_triton, + ) + + if is_unified_kv_triton(): + o = attn_backend.forward( + q=q_out if q_out is not None else q, + k=attn_k, + v=attn_k, + layer=self.attn_mqa, + forward_batch=forward_batch, + compress_ratio=self.compress_ratio, + attn_sink=self.attn_sink, + save_kv_cache=kv is not None, + ) + else: + attn_q = q_padded if q_padded is not None else q + save_kv_cache = False + if forward_batch.forward_mode.is_extend() and is_in_breakable_cuda_graph(): + o = attn_q.new_empty( + (*attn_q.shape[:-1], self.attn_mqa.v_head_dim), + ) + bcg_deepseek_v4_attention_with_output( + attn_q, + attn_k, + o, + self.attn_mqa.layer_id, + self.compress_ratio, + attn_sink, + save_kv_cache, + ) + else: + o = attn_backend.forward( + q=attn_q, + k=attn_k, + v=attn_k, + layer=self.attn_mqa, + forward_batch=forward_batch, + compress_ratio=self.compress_ratio, + attn_sink=attn_sink, + save_kv_cache=save_kv_cache, + ) + o = o[:, tp_slice, :] + if _is_npu: + cos4, sin4 = self._get_npu_rope_position_cache( + positions, o.dtype, inverse=True + ) + Dsv4NpuRoPE.apply_rotary_mul_inplace( + o, + None, + cos4, + sin4, + qk_nope_dim=self.qk_nope_head_dim, + ) + else: + fused_rope_inplace( + o[..., -self.qk_rope_head_dim :], + None, + self.freqs_cis, + positions=positions, + inverse=True, + ) + + o = o.view(o.shape[0], self.n_local_groups, -1) + + if _FP8_WO_A_GEMM: + import deep_gemm + + from sglang.srt.layers import deep_gemm_wrapper + + T, G, D = o.shape + R = self.o_lora_rank + if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0: + # sm100 (Blackwell): ue8m0 scales via the dedicated JIT kernel. + o_fp8, o_s = sglang_per_token_group_quant_fp8_dsv4_wo_a(o) + recipe = (1, 1, 128) + else: + # sm90 (Hopper): fp32 scales. + o_fp8, o_s = sglang_per_token_group_quant_fp8( + o.reshape(T * G, D).contiguous(), + group_size=128, + scale_ue8m0=False, + ) + o_fp8 = o_fp8.view(T, G, D) + o_s = o_s.view(T, G, -1) + recipe = (1, 128, 128) + output = torch.empty(T, G, R, device=o.device, dtype=torch.bfloat16) + deep_gemm.fp8_einsum( + "bhr,hdr->bhd", + (o_fp8, o_s), + (self.wo_a.weight.view(G, R, D), self.wo_a.weight_scale_inv.data), + output, + recipe=recipe, + ) + o = output + else: + wo_a_weight = getattr(self.wo_a, "weight", None) + if wo_a_weight is not None: + wo_a = wo_a_weight.view(self.n_local_groups, self.o_lora_rank, -1) + o = _apply_wo_a_bf16_matmul( + o, wo_a, is_decode=forward_batch.forward_mode.is_decode() + ) + else: + o = _apply_gguf_grouped_wo_a( + o, + self.wo_a.qweight, + self.wo_a.qweight_type.weight_type, + self.o_lora_rank, + ) + + o, _ = self.wo_b(o.flatten(1)) + if self.attn_tp_size > 1 and self.attn_tp_size < get_parallel().tp_size: + o = attn_tp_all_reduce(o) + + return o + + # ---- TBO op decomposition (prefill two-batch-overlap) ---- + def op_attn(self, state): + """Run the attention forward as a single TBO op. + + Consumes the post-input-norm hidden states produced by + ``DeepseekV4DecoderLayer.op_mhc_prepare_attn`` and stores the attention + output for ``op_mhc_post_attn_pre_mlp``. + """ + state.hidden_states_after_attn = self.forward( + x=state.pop("hidden_states_after_input_norm"), + positions=state.positions, + forward_batch=state.forward_batch, + x_quant=state.pop("attn_x_quant"), + ) + + +class DeepseekV4DecoderLayer(nn.Module): + def __init__( + self, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + moe_quant_config_override: Optional[QuantizationConfig] = None, + is_nextn: bool = False, + prefix: str = "", + alt_streams: Optional[List[torch.cuda.Stream]] = None, + compress_ratio_override: Optional[int] = None, + ) -> None: + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.layer_id = layer_id + self.self_attn = self._build_self_attn( + config=config, + layer_id=layer_id, + quant_config=quant_config, + prefix=add_prefix("self_attn", prefix), + alt_streams=alt_streams, + compress_ratio_override=compress_ratio_override, + ) + moe_alt_stream = ( + alt_streams[0] + if ( + alt_streams is not None + and ( + _is_cuda + or envs.SGLANG_ROCM_USE_MULTI_STREAM.get() + or envs.SGLANG_NPU_USE_MULTI_STREAM.get() + ) + ) + else None + ) + self.mlp = deepseek_v2.DeepseekV2MoE( + config=config, + quant_config=moe_quant_config_override or quant_config, + prefix=add_prefix("mlp", prefix), + layer_id=self.layer_id, + alt_stream=moe_alt_stream, + is_nextn=is_nextn, + is_deepseek_v4=True, + ) + + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + self.hc_mult = hc_mult = config.hc_mult + self.hc_sinkhorn_iters = config.hc_sinkhorn_iters + self.hc_eps = config.hc_eps + ( + self.hc_attn_fn, + self.hc_ffn_fn, + self.hc_attn_base, + self.hc_ffn_base, + self.hc_attn_scale, + self.hc_ffn_scale, + ) = make_hc_mixing_params(hc_mult, config.hidden_size) + self.rms_norm_eps = config.rms_norm_eps + self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.use_fused_mhc_post_pre = is_cross_layer_mhc_fusion_enabled() + self._input_layernorm_weight_bf16 = None + self._post_attention_layernorm_weight_bf16 = None + + def _build_self_attn( + self, + *, + config: DeepSeekV4Config, + layer_id: int, + quant_config: Optional[QuantizationConfig], + prefix: str, + alt_streams: Optional[List[torch.cuda.Stream]], + compress_ratio_override: Optional[int], + ) -> nn.Module: + return MQALayer( + config=config, + layer_id=layer_id, + quant_config=quant_config, + prefix=prefix, + alt_streams=alt_streams, + compress_ratio_override=compress_ratio_override, + ) + + def refresh_mhc_norm_weight_cache(self): + # Cache bf16 norm weights so the fused path does not allocate/cast per forward. + self._input_layernorm_weight_bf16 = ( + self.input_layernorm.weight.data.bfloat16().contiguous() + ) + self._post_attention_layernorm_weight_bf16 = ( + self.post_attention_layernorm.weight.data.bfloat16().contiguous() + ) + + def hc_pre( + self, + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + norm: Optional[nn.Module] = None, + forward_batch: Optional[ForwardBatch] = None, + ): + """If *norm* is given and the TileLang path is active, the returned + hidden_states are already post-norm (the norm is fused into the kernel).""" + + @compile_in_capture_mode + def hc_pre_torch_impl(x, hc_fn): + x_flat = x.flatten(1).float() + rsqrt = torch.rsqrt( + x_flat.square().mean(-1, keepdim=True) + self.rms_norm_eps + ) + mixes = (F.linear(x_flat, hc_fn) * rsqrt).unsqueeze(1) + return x_flat, mixes + + shape, dtype = x.size(), x.dtype + + if _is_npu: + return _get_mhc_ops().npu_hc_pre( + x, + hc_fn, + hc_scale, + hc_base, + hc_mult=self.hc_mult, + hc_sinkhorn_iters=self.hc_sinkhorn_iters, + rms_norm_eps=self.rms_norm_eps, + hc_eps=self.hc_eps, + forward_batch=forward_batch, + ) + + if x.shape[0] == 0: + y = torch.empty((0, shape[-1]), dtype=dtype, device=x.device) + post = torch.empty((0, self.hc_mult), dtype=torch.float32, device=x.device) + comb = torch.empty( + (0, self.hc_mult, self.hc_mult), dtype=torch.float32, device=x.device + ) + return y, post, comb, False + + if envs.SGLANG_OPT_USE_FLASHINFER_MHC.get(): + y, post, comb = _flashinfer_hc_pre( + x, + hc_fn, + hc_scale, + hc_base, + rms_eps=self.rms_norm_eps, + hc_eps=self.hc_eps, + sinkhorn_iters=self.hc_sinkhorn_iters, + ) + return y, post, comb, False + + if envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): + from sglang.kernels.ops.layernorm.mhc import mhc_pre + + norm_kwargs = {} + if norm is not None: + norm_kwargs["norm_weight"] = norm.weight.data + norm_kwargs["norm_eps"] = norm.variance_epsilon + + post, comb, y = mhc_pre( + residual=x, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=self.rms_norm_eps, + hc_pre_eps=self.hc_eps, + hc_sinkhorn_eps=self.hc_eps, + hc_post_mult_value=_MHC_POST_MULT_VALUE, + sinkhorn_repeat=self.hc_sinkhorn_iters, + **norm_kwargs, + ) + return y, post.squeeze(-1), comb, norm is not None + + if _is_hip: + from aiter.ops.mhc import mhc_pre + + post, comb, y = mhc_pre( + residual=x, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=self.rms_norm_eps, + hc_pre_eps=self.hc_eps, + hc_sinkhorn_eps=self.hc_eps, + hc_post_mult_value=_MHC_POST_MULT_VALUE, + sinkhorn_repeat=self.hc_sinkhorn_iters, + ) + return y, post.squeeze(-1), comb, False + + if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + from sglang.srt.layers.deep_gemm_wrapper.entrypoint import ( + tf32_hc_prenorm_gemm, + ) + + x_flat = x.flatten(1).bfloat16() + + m, k = x_flat.shape + mix_hc = hc_fn.size(0) + d_out = torch.empty((m, mix_hc), dtype=torch.float, device=x.device) + s_out = torch.empty((m,), dtype=torch.float, device=x.device) + tf32_hc_prenorm_gemm( + x_flat, hc_fn.float().contiguous(), d_out, s_out, num_splits=None + ) + rsqrt = torch.rsqrt(s_out / k + self.rms_norm_eps) + mixes = (d_out * rsqrt.unsqueeze(1)).unsqueeze(1) + else: + x_flat, mixes = hc_pre_torch_impl(x, hc_fn) + + pre, post, comb = _get_mhc_ops().hc_split_sinkhorn( + mixes, + hc_scale, + hc_base, + self.hc_mult, + self.hc_sinkhorn_iters, + self.hc_eps, + ) + # y is the post-norm activation fed into the MoE. Allocate it in the + # symmetric memory pool so the downstream all-reduce uses the low-latency + # NCCL symmetric path: the Triton inplace MoE runner writes the expert + # output back into this buffer, so a symmetric input yields a symmetric + # all-reduce input. Gated by is_allocation_symmetric() (mirrors the + # TileLang path in _mhc_pre_impl / mhc_fused_post_pre). + with use_symmetric_memory( + get_tp_group(), disabled=not is_allocation_symmetric() + ): + y = (pre.squeeze(1).unsqueeze(-1) * x_flat.view(shape)).sum(dim=1).to(dtype) + return y, post.squeeze(1), comb.squeeze(1), False + + def hc_post( + self, + x: torch.Tensor, + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + ): + + if x.shape[0] == 0: + return torch.empty( + (0, self.hc_mult, x.shape[-1]), dtype=x.dtype, device=x.device + ) + + if _is_npu: + return torch.ops.custom.npu_hc_post(x, residual, post, comb) + + if envs.SGLANG_OPT_USE_FLASHINFER_MHC.get(): + from flashinfer.mhc import mhc_post + + return mhc_post(x, residual, post, comb) + + if envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get(): + from sglang.kernels.ops.layernorm.mhc import mhc_post + + return mhc_post(x, residual, post, comb) + + elif _is_hip: + from aiter.ops.mhc import mhc_post + + result = torch.empty_like(residual) + mhc_post(result, x, residual, post, comb) + return result + + assert residual.shape == (x.shape[0], self.hc_mult, x.shape[-1]) + assert post.shape == (x.shape[0], self.hc_mult) + assert comb.shape == (x.shape[0], self.hc_mult, self.hc_mult) + + @compile_in_capture_mode + def hc_post_torch_impl(x, residual, post, comb): + return ( + post.unsqueeze(-1) * x.unsqueeze(1) + + (comb.unsqueeze(-1) * residual.unsqueeze(2)).sum(dim=1) + ).type_as(x) + + return hc_post_torch_impl(x, residual, post, comb) + + def forward( + self, + positions: torch.tensor, + hidden_states: torch.Tensor, + input_ids: torch.Tensor, + forward_batch: ForwardBatch, + input_ids_global: torch.Tensor, + prev_residual: Optional[torch.Tensor] = None, + prev_post: Optional[torch.Tensor] = None, + prev_comb: Optional[torch.Tensor] = None, + ) -> Tuple[ + torch.Tensor, + Optional[torch.Tensor], + Optional[torch.Tensor], + Optional[torch.Tensor], + ]: + use_fused = self.use_fused_mhc_post_pre + + if prev_residual is not None and use_fused: + # Dispatch cascade: aiter HIP (gfx95) -> Triton (gfx95 small-batch + # <=64 tokens, or gfx1250 all sizes) -> TileLang -> None. + input_norm_weight = ( + self._input_layernorm_weight_bf16 + if self._input_layernorm_weight_bf16 is not None + else self.input_layernorm.weight.data + ) + fused = apply_mhc_post_pre_boundary( + hidden_states, + prev_residual, + prev_post, + prev_comb, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + self.hc_mult, + self.rms_norm_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + input_norm_weight, + self.input_layernorm.variance_epsilon, + fn_transpose=False, + ) + if fused is not None: + residual, hidden_states, post, comb, norm_fused = fused + if not norm_fused: + # Triton fused post+pre (gfx95 small-batch or gfx1250) returns + # norm_fused=False — the input layernorm is NOT folded. + # gfx95 takes the fp8-quant path; gfx1250 takes plain layernorm. + if _use_aiter and _is_gfx95_supported: + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None + else: + hidden_states = self.hc_post( + hidden_states, prev_residual, prev_post, prev_comb + ) + residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm=self.input_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + if _use_aiter and _is_gfx95_supported: + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None + else: + residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm=self.input_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + if _use_aiter and _is_gfx95_supported: + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None + + with self.self_attn.maybe_use_decode_attn_tp(forward_batch): + hidden_states = self.self_attn( + x=hidden_states, + positions=positions, + forward_batch=forward_batch, + x_quant=x_quant, + ) + + if use_fused: + post_attn_norm_weight = ( + self._post_attention_layernorm_weight_bf16 + if self._post_attention_layernorm_weight_bf16 is not None + else self.post_attention_layernorm.weight.data + ) + fused = apply_mhc_post_pre_boundary( + hidden_states, + residual, + post, + comb, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + self.hc_mult, + self.rms_norm_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + post_attn_norm_weight, + self.post_attention_layernorm.variance_epsilon, + fn_transpose=True, + ) + if fused is not None: + residual, hidden_states, post, comb, norm_fused = fused + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) + else: + hidden_states = self.hc_post(hidden_states, residual, post, comb) + residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + norm=self.post_attention_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) + else: + hidden_states = self.hc_post(hidden_states, residual, post, comb) + residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + norm=self.post_attention_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) + + hidden_states = self._run_moe_ffn_dp_sync( + hidden_states, + forward_batch, + input_ids=input_ids, + input_ids_global=input_ids_global, + ) + + if not use_fused: + hidden_states = self.hc_post(hidden_states, residual, post, comb) + return hidden_states, None, None, None + + # Return the deferred FFN hc_post state; the next layer consumes it with + # cross-layer fusion, and the final layer is completed in DeepseekV4Model. + return hidden_states, residual, post, comb + + def _run_moe_ffn_dp_sync( + self, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + *, + input_ids: torch.Tensor, + input_ids_global: torch.Tensor, + ) -> torch.Tensor: + _use_cp = self.dsa_enable_prefill_cp and dsa_use_prefill_cp(forward_batch) + _use_tp_moe_gather = ( + not _use_cp + and get_parallel().attn_dp_size > 1 + and get_moe_a2a_backend().is_none() + ) + _use_tp_attn_a2a_scatter = ( + not _use_cp + and get_parallel().attn_tp_size > 1 + and not get_moe_a2a_backend().is_none() + ) + # symmetric gather+scatter for the no-EP TP-MoE dp-attn path: + # all_gatherv gather (in self.mlp's dp_gather) + reduce_scatterv combine. + # The experts ARE TP-sharded by intermediate (moe_tp_size==tp_size), so + # the post-experts reduce is a SUM. reduce_scatterv does that sum+scatter + # in ONE op, REPLACING the MoE-internal post-experts all_reduce — so we + # MUST tell the MoE to skip it (mlp_reduce_scatter=True) or it + # double-reduces. Env-gated via SGLANG_DP_USE_GATHERV, default OFF. + _use_reduce_scatterv = ( + _use_tp_moe_gather + and is_dp_gatherv_active() + and forward_batch.dp_padding_mode is not None + and not forward_batch.dp_padding_mode.is_max_len() + ) + # SGLANG_DP_USE_REDUCE_SCATTER: in the MAX_LEN decode path (equal per-rank + # padding, gatherv inactive, no EP), replace the MoE-internal post-experts + # all_reduce + dp_scatter with an equal-chunk reduce_scatter. On ROCm this + # uses the aiter custom kernel (so BOTH gather and combine are aiter custom), + # elsewhere RCCL reduce_scatter; either way it cuts combine traffic ~2x vs + # all_reduce. tp_size==attn_dp_size required so the global buffer splits + # evenly into per-rank chunks. + _use_reduce_scatter = ( + envs.SGLANG_DP_USE_REDUCE_SCATTER.get() + and _use_tp_moe_gather + and not _use_reduce_scatterv + and not should_use_dp_reduce_scatterv() + and forward_batch.dp_padding_mode is not None + and forward_batch.dp_padding_mode.is_max_len() + and get_parallel().tp_size == get_parallel().attn_dp_size + ) + mlp_reduce_scatter = _use_cp or _use_reduce_scatterv or _use_reduce_scatter + # PoC (SGLANG_DP_SHARED_EXPERT_LOCAL): compute the replicated shared expert + # on LOCAL hidden before the gather and add it back after the combine + # (reduce_scatterv OR dp_scatter), instead of on the gathered global buffer. + # Applies to BOTH prefill and decode: the shared expert is a per-token MLP, + # so computing it on this rank's local tokens (M_local rows) is identical to + # computing it on the gathered global buffer (M_global rows) and keeping the + # local slice -- but costs 1/dp_size the rows. With a replicated (TP1) shared + # expert this cancels the TP1 "full-dim" cost in decode (M_local * dim == + # M_global * dim/tp), so decode no longer pays the ~dp_size x penalty. + _shared_local = None + _do_shared_local = ( + _SHARED_EXPERT_LOCAL + and _use_tp_moe_gather + and getattr(self.mlp, "shared_experts", None) is not None + and getattr(self.mlp, "_shared_expert_tp1", False) + ) + if _use_cp: + moe_a2a_backend = get_moe_a2a_backend() + if moe_a2a_backend.is_none(): + hidden_states = dsa_cp_gather_hidden_states(hidden_states) + else: + assert ( + moe_a2a_backend.is_deepep() + or moe_a2a_backend.is_megamoe() + or moe_a2a_backend.is_mori() + ), ( + "CP requires moe_a2a_backend in ('deepep', 'megamoe', 'mori'), " + f"got {moe_a2a_backend.value!r}." + ) + elif _use_tp_moe_gather: + hidden_states, local_hidden_states = ( + get_global_dp_buffer(get_tp_group()), + hidden_states, + ) + if _do_shared_local and local_hidden_states.shape[0] > 0: + _shared_local = self.mlp._forward_shared_experts(local_hidden_states) + # self_attn has already reduced across attention TP, so these hidden + # states are replicated and must not be summed by a partial gather. + dp_gather_replicate(hidden_states, local_hidden_states, forward_batch) + _a2a_scatter_chunks: Optional[List[torch.Tensor]] = None + if _use_tp_attn_a2a_scatter: + s, r = get_parallel().attn_tp_size, get_parallel().attn_tp_rank + _a2a_scatter_chunks = list(hidden_states.tensor_split(s)) + hidden_states = _a2a_scatter_chunks[r].contiguous() + input_ids = input_ids.tensor_split(s)[r].contiguous() + input_ids_global = input_ids_global.tensor_split(s)[r].contiguous() + # Skip the MoE-internal post-experts all_reduce when we will do the + # reduce via reduce_scatterv/reduce_scatter at the combine below + # (else double-reduce). + with get_forward().scoped(mlp_reduce_scatter=mlp_reduce_scatter): + hidden_states = self.mlp( + hidden_states, + forward_batch, + input_ids=input_ids, + input_ids_global=input_ids_global, + skip_shared_experts=_do_shared_local, + ) + if _use_cp and get_moe_a2a_backend().is_none(): + hidden_states = dsa_cp_reduce_scatter_hidden_states(hidden_states) + elif _use_tp_moe_gather: + hidden_states, global_hidden_states = ( + get_local_dp_buffer(get_tp_group()), + hidden_states, + ) + if should_use_dp_reduce_scatterv() or _use_reduce_scatterv: + # SUM the TP-sharded per-rank partial expert outputs AND scatter + # each rank its own token slice, in one op. Correct because the + # MoE-internal all_reduce was skipped (mlp_reduce_scatter above). + # This is the symmetric inverse of the all_gatherv gather. + get_tp_group().reduce_scatterv( + global_hidden_states, + output=hidden_states, + sizes=get_dp_global_num_tokens(), + ) + elif _use_reduce_scatter: + # Equal-chunk reduce_scatter: SUM the TP-sharded per-rank partial + # expert outputs AND scatter each rank its own (MAX_LEN-padded) + # token chunk in one op (symmetric inverse of the MAX_LEN + # all_gather). Correct because the MoE-internal all_reduce was + # skipped (mlp_reduce_scatter above). dp_reduce_scatter_tensor + # routes to the equal-chunk reduce_scatter_tensor here (its + # variable-length reduce_scatterv branch is gated by + # is_dp_gatherv_active(), which is False under MAX_LEN), which in + # turn uses the aiter custom kernel when it fits (else RCCL). + dp_reduce_scatter_tensor(hidden_states, global_hidden_states) + else: + dp_scatter(hidden_states, global_hidden_states, forward_batch) + # PoC: add the locally-computed shared-expert output to this rank's + # reduce-scattered / dp-scattered local slice (skipped inside self.mlp + # above). Covers both prefill (gatherv) and decode (dp_scatter). + if _shared_local is not None: + n = hidden_states.shape[0] + hidden_states = hidden_states + _shared_local[:n] + if _use_tp_attn_a2a_scatter: + assert _a2a_scatter_chunks is not None + gathered = [torch.empty_like(t) for t in _a2a_scatter_chunks] + attn_tp_all_gather(gathered, hidden_states.contiguous()) + hidden_states = torch.cat(gathered) + return hidden_states + + # ------------------------------------------------------------------ + # TBO op decomposition (prefill two-batch-overlap, EP / mori path) + # + # These mirror the NON-fused branch of ``forward`` (cross-layer mHC + # fusion is disabled under TBO, so every layer is self-contained), split + # into ops so the operations engine can overlap one ubatch's MoE a2a + # dispatch/combine with the other ubatch's attention + expert GEMM. + # The MoE ops themselves (op_gate / op_select_experts / op_dispatch_a/b / + # op_experts / op_combine_a/b / op_shared_experts / op_output) are reused + # as-is from ``self.mlp`` (DeepseekV2MoE) — they decompose ``forward_deepep``. + # ------------------------------------------------------------------ + def op_mhc_prepare_attn( + self, + state, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + residual: Optional[torch.Tensor] = None, + tbo_subbatch_index: Optional[int] = None, + **kwargs, + ): + # Non-fused attention-side mHC pre + input layernorm. + attn_residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm=self.input_layernorm, + forward_batch=forward_batch, + ) + if not norm_fused: + if _use_aiter and (_is_gfx95_supported or _is_gfx1250_supported): + x_quant, hidden_states = _fused_rmsnorm_fp8_quant( + hidden_states, + self.input_layernorm.weight, + self.rms_norm_eps, + ) + else: + hidden_states = self.input_layernorm(hidden_states) + x_quant = None + else: + x_quant = None + + state.attn_residual = attn_residual + state.attn_post = post + state.attn_comb = comb + state.hidden_states_after_input_norm = hidden_states + state.attn_x_quant = x_quant + # mori's op_output slices final_hidden_states[:num_tokens]. + if get_moe_a2a_backend().is_mori(): + state.num_tokens = attn_residual.shape[0] + state.update( + dict( + forward_batch=forward_batch, + positions=positions, + tbo_subbatch_index=tbo_subbatch_index, + ) + ) + + def op_mhc_post_attn_pre_mlp(self, state): + # Close the attention mHC (hc_post), then open the FFN-side mHC pre + + # post-attention layernorm. Produces the 2D MoE input. + # + # Pop each boundary tensor from the state EXACTLY ONCE, up front, and + # reuse the locals for both the fused attempt and the non-fused + # fallback. apply_mhc_post_pre_boundary() returns None when it declines + # to fuse -- most importantly for the 0-token DP two-batch-overlap idle + # ubatch -- in which case control must fall through to the unfused + # hc_post. Popping in the fused call's arguments and again in the + # fallback would double-pop -> KeyError: 'hidden_states_after_attn' on + # every idle DP rank. use_fused_mhc_post_pre is on whenever the aiter + # gfx95 mHC path is available (is_cross_layer_mhc_fusion_enabled), so + # this fallback is reached under DP regardless of the TileLang env. + hidden_states_after_attn = state.pop("hidden_states_after_attn") + attn_residual = state.pop("attn_residual") + attn_post = state.pop("attn_post") + attn_comb = state.pop("attn_comb") + + if self.use_fused_mhc_post_pre: + post_attn_norm_weight = ( + self._post_attention_layernorm_weight_bf16 + if self._post_attention_layernorm_weight_bf16 is not None + else self.post_attention_layernorm.weight.data + ) + fused = apply_mhc_post_pre_boundary( + hidden_states_after_attn, + attn_residual, + attn_post, + attn_comb, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + self.hc_mult, + self.rms_norm_eps, + self.hc_eps, + _MHC_POST_MULT_VALUE, + self.hc_sinkhorn_iters, + post_attn_norm_weight, + self.post_attention_layernorm.variance_epsilon, + fn_transpose=True, + ) + if fused is not None: + ffn_residual, hidden_states, post, comb, norm_fused = fused + if not norm_fused: + # The Triton fused post+pre skips the post-attention + # layernorm (norm_fused=False); apply it before the MoE, + # matching the unfused hc_pre path below. + hidden_states = self.post_attention_layernorm(hidden_states) + state.ffn_residual = ffn_residual + state.ffn_post = post + state.ffn_comb = comb + state.hidden_states_mlp_input = hidden_states + return + + hidden_states = self.hc_post( + hidden_states_after_attn, + attn_residual, + attn_post, + attn_comb, + ) + ffn_residual = hidden_states + hidden_states, post, comb, norm_fused = self.hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + norm=self.post_attention_layernorm, + forward_batch=state.forward_batch, + ) + if not norm_fused: + hidden_states = self.post_attention_layernorm(hidden_states) + state.ffn_residual = ffn_residual + state.ffn_post = post + state.ffn_comb = comb + state.hidden_states_mlp_input = hidden_states + + def op_mhc_postprocess(self, state): + # Close the FFN mHC (hc_post) and emit the next layer's input dict. + hidden_states = self.hc_post( + state.pop("hidden_states_mlp_output"), + state.pop("ffn_residual"), + state.pop("ffn_post"), + state.pop("ffn_comb"), + ) + output = dict( + positions=state.positions, + hidden_states=hidden_states, + # DSV4 non-fused layers carry no residual across layers; the key is + # required by the next layer's op_mhc_prepare_attn (ignored) and by + # _model_forward_tbo_merge_outputs (None -> None). + residual=None, + forward_batch=state.forward_batch, + tbo_subbatch_index=state.tbo_subbatch_index, + ) + state.clear( + expect_keys={ + "positions", + "forward_batch", + "tbo_subbatch_index", + } + ) + return output + + # ------------------------------------------------------------------ + # Non-EP (DP TP-MoE) TBO ops. Overlap the DP all_gatherv (pre-MoE gather) + # + reduce_scatterv (post-MoE combine) with the OTHER ubatch's attn+MoE + # compute. Used when moe_a2a_backend is "none" (DP-attention, TP-MoE) — + # the path ATOM uses for DSV4 (+~7.7% prefill). Replaces the EP mori + # op_dispatch/op_combine. op_mhc_* and op_attn are reused (local hidden). + # ------------------------------------------------------------------ + def op_gather_a(self, state): + # Launch the all_gatherv (local hidden -> global buffer) + the input_ids + # replicate-gather on the shared comm stream; record an event. + fb = state.forward_batch + local = state.pop("hidden_states_mlp_input") # LOCAL [M_local, hidden] + # Shared-expert-local: compute on LOCAL hidden before the gather; added + # back after the combine (same as the non-fused forward). Skipped in the + # global MoE via skip_shared_experts. + do_shared_local = ( + _SHARED_EXPERT_LOCAL + and getattr(self.mlp, "shared_experts", None) is not None + and getattr(self.mlp, "_shared_expert_tp1", False) + ) + state.do_shared_local = do_shared_local + state.shared_local = ( + self.mlp._forward_shared_experts(local) + if (do_shared_local and local.shape[0] > 0) + else None + ) + # Persistent grow-only scratch (keyed per ubatch) instead of a fresh + # torch.empty each layer -> stops the allocator's `reserved` from + # ballooning at large prefill chunks. input_ids_global is gathered ONCE + # per ubatch in _forward_layers_tbo (cached on fb), not here. + sub = state.tbo_subbatch_index + global_rows = get_global_dp_buffer_len() + global_hidden = get_tbo_persistent_buffer( + ("gh", sub), global_rows, local.shape[1], local.dtype, local.device + ) + comm = get_dp_tbo_comm_stream() + compute = torch.cuda.current_stream() + with torch.cuda.stream(comm): + comm.wait_stream(compute) + dp_gather_partial(global_hidden, local, fb) + state.gather_event = _tbo_event(("gather", sub)) + state.gather_event.record(comm) + state.gather_keepalive = local + state.global_hidden = global_hidden + + def op_gather_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("gather_event")) + # Compute now ordered after the gather -> the gather input is safe to + # release (freed on the compute stream, no record_stream deferral). + state.pop("gather_keepalive") + + def op_moe(self, state): + # MoE (gate/topk/experts) on the GLOBAL gathered buffer. mlp_reduce_scatter + # skips the MoE-internal all_reduce (we reduce_scatterv in op_combine). + fb = state.forward_batch + global_hidden = state.pop("global_hidden") + global_ids = fb._tbo_global_input_ids + with get_forward().scoped(mlp_reduce_scatter=True): + state.global_expert_out = self.mlp( + global_hidden, + fb, + input_ids=global_ids, + input_ids_global=global_ids, + skip_shared_experts=state.do_shared_local, + ) + + def op_combine_a(self, state): + # Launch reduce_scatterv (global partial expert sums -> per-rank local) on + # the comm stream; record an event. Symmetric inverse of the all_gatherv. + global_out = state.pop("global_expert_out") + local_out = get_tbo_persistent_buffer( + ("lo", state.tbo_subbatch_index), + get_local_dp_buffer_len(), + global_out.shape[1], + global_out.dtype, + global_out.device, + ) + state.combine_event = dp_reduce_scatterv_async( + local_out, + global_out, + get_dp_global_num_tokens(), + event_key=("combine", state.tbo_subbatch_index), + ) + state.local_out = local_out + # Keep the (variable-size) MoE output alive until op_combine_b waits on + # the combine event (replaces record_stream; avoids reserved churn). + state.combine_keepalive = global_out + + def op_combine_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("combine_event")) + state.pop("combine_keepalive") + hidden = state.pop("local_out") + shared_local = state.pop("shared_local") + state.pop("do_shared_local") + if shared_local is not None: + n = hidden.shape[0] + hidden = hidden + shared_local[:n] + state.hidden_states_mlp_output = hidden + + def _cp_tbo_launch(self, state, x, key, out_rows, collective): + assert _is_hip, "CP+TBO MoE overlap is HIP-only" + x = x.contiguous() + sub = state.tbo_subbatch_index + out = get_tbo_persistent_buffer( + (key, sub), out_rows, x.shape[1], x.dtype, x.device + ) + comm = get_dp_tbo_comm_stream() + comm.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(comm): + collective(out, x) + event = _tbo_event((key, sub)) + event.record(comm) + return out, event, x + + def op_cp_gather_a(self, state): + local = state.pop("hidden_states_mlp_input") + out, event, keepalive = self._cp_tbo_launch( + state, + local, + "cpgh", + local.shape[0] * get_parallel().attn_cp_size, + attn_cp_overlap_all_gather_into_tensor, + ) + state.global_hidden = out + state.cp_gather_event = event + state.cp_gather_keepalive = keepalive + + def op_cp_gather_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("cp_gather_event")) + state.pop("cp_gather_keepalive") + + def op_cp_moe(self, state): + fb = state.forward_batch + global_ids = fb._cp_moe_input_ids + with get_forward().scoped(mlp_reduce_scatter=True): + state.global_expert_out = self.mlp( + state.pop("global_hidden"), + fb, + input_ids=global_ids, + input_ids_global=global_ids, + ) + + def op_cp_combine_a(self, state): + global_out = state.pop("global_expert_out") + out, event, keepalive = self._cp_tbo_launch( + state, + global_out, + "cplo", + global_out.shape[0] // get_parallel().attn_cp_size, + attn_cp_overlap_reduce_scatter_tensor, + ) + state.local_out = out + state.cp_combine_event = event + state.cp_combine_keepalive = keepalive + + def op_cp_combine_b(self, state): + torch.cuda.current_stream().wait_event(state.pop("cp_combine_event")) + state.pop("cp_combine_keepalive") + state.hidden_states_mlp_output = state.pop("local_out") + + +class DeepseekV4Model(nn.Module): + fall_back_to_pt_during_load = False + + def __init__( + self, + config: DeepSeekV4Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.pp_group = get_pp_group() + self.hidden_size = config.hidden_size + if self.pp_group.is_first_rank: + embedding_quant_config = ( + quant_config + if quant_config is not None and quant_config.get_name() == "expert_pack" + else None + ) + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + enable_tp=not is_dp_attention_enabled(), + quant_config=embedding_quant_config, + prefix=add_prefix("embed_tokens", prefix), + ) + else: + self.embed_tokens = PPMissingLayer() + self.rms_norm_eps = config.rms_norm_eps + use_stream_pool = ( + _is_cuda + or ( + _is_hip + and ( + envs.SGLANG_ROCM_USE_MULTI_STREAM.get() + or envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get() + ) + ) + or (_is_npu and envs.SGLANG_NPU_USE_MULTI_STREAM.get()) + ) + device_module = torch.get_device_module() + num_alt_streams = 5 if (_is_cuda or _is_npu) else 2 + self.alt_streams = ( + [device_module.Stream() for _ in range(num_alt_streams)] + if use_stream_pool + else None + ) + self.layers, self.start_layer, self.end_layer = make_layers( + config.num_hidden_layers, + lambda idx, prefix: DeepseekV4DecoderLayer( + config=config, + layer_id=idx, + quant_config=quant_config, + prefix=prefix, + alt_streams=self.alt_streams, + ), + pp_rank=self.pp_group.rank_in_group, + pp_size=self.pp_group.world_size, + prefix=add_prefix("layers", prefix), + ) + if self.pp_group.is_last_rank: + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + else: + self.norm = PPMissingLayer() + self.gemm_output_zero_allocator_size = 0 + self.hc_eps = config.hc_eps + self.hc_mult = hc_mult = config.hc_mult + self.norm_eps = config.rms_norm_eps + if self.pp_group.is_last_rank: + ( + self.hc_head_fn, + self.hc_head_base, + self.hc_head_scale, + ) = make_hc_head_params(hc_mult, config.hidden_size) + + self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + self.use_fused_mhc_post_pre = is_cross_layer_mhc_fusion_enabled() + if self.dsa_enable_prefill_cp: + self.cp_size = get_parallel().attn_cp_size + + self.dspark_layers_to_capture: Optional[List[int]] = None + + def get_input_embeddings(self) -> nn.Module: + return self.embed_tokens + + def hc_head( + self, + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + ): + if x.numel() > 0: + from sglang.kernels.ops.layernorm.mhc_head import fused_hc_head + + return fused_hc_head( + x.contiguous(), + hc_fn, + hc_scale, + hc_base, + norm_eps=self.norm_eps, + hc_eps=self.hc_eps, + ) + return hc_head_torch( + x, + hc_fn, + hc_scale, + hc_base, + norm_eps=self.norm_eps, + hc_eps=self.hc_eps, + ) + + def _cp_children_splittable(self, forward_batch: ForwardBatch) -> bool: + children = forward_batch.tbo_children + if not children: + return False + cp_size = get_parallel().attn_cp_size + for child in children: + if child.batch_size <= 0 or child.extend_seq_lens_cpu is None: + return False + if sum(child.extend_seq_lens_cpu) < cp_size: + return False + return True + + def _can_run_tbo(self, forward_batch: ForwardBatch) -> bool: + """DSV4 prefill-only two-batch-overlap gate. + + TBO batch prep (tbo_split_seq_index / tbo_children) is populated + model-agnostically when --enable-two-batch-overlap is set and the + DP-attention preparer allows it (mori `normal` mode permits prefill + TBO). We additionally restrict to: prefill (EXTEND), single PP, and a + path the DSV4 op strategy implements -- the non-CP path everywhere, plus + the round-robin DSA prefill CP path on HIP. + """ + from sglang.srt.layers.moe import is_tbo_enabled + + if dsa_use_prefill_cp(forward_batch): + path_ok = ( + _is_hip + and not is_cp_v2_active(forward_batch) + and is_dsa_prefill_cp_round_robin_split() + and get_moe_a2a_backend().is_none() + and self._cp_children_splittable(forward_batch) + ) + else: + path_ok = ( + not _is_hip + or not get_moe_a2a_backend().is_none() + or get_parallel().attn_dp_size > 1 + ) + return ( + is_tbo_enabled() + and forward_batch.can_run_tbo + and forward_batch.tbo_children is not None + and forward_batch.global_forward_mode is not None + # MTP target-verify also reports is_extend(); only real prefill + # should enter the prefill TBO strategy. + and forward_batch.global_forward_mode.is_extend_without_speculative() + and path_ok + and self.pp_group.world_size == 1 + ) + + def _forward_layers_tbo( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + from sglang.srt.batch_overlap.operations import execute_overlapped_operations + from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy + from sglang.srt.batch_overlap.two_batch_overlap import ( + _model_forward_filter_inputs, + _model_forward_tbo_merge_outputs, + ) + + if _is_hip and dsa_use_prefill_cp(forward_batch): + return self._forward_layers_tbo_cp( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + + layers = [self.layers[i] for i in range(self.start_layer, self.end_layer)] + operations_strategy = OperationsStrategy.init_new_tbo( + layers, forward_batch.global_forward_mode + ) + + # Split the per-rank batch into the 2 ubatches (token-range slice + pad + # to tbo_padded_len). residual is unused by the DSV4 non-fused layer ops. + inputs_arr = [ + _model_forward_filter_inputs( + hidden_states=hidden_states, + residual=None, + positions=positions, + output_forward_batch=child, + tbo_subbatch_index=idx, + ) + for idx, child in enumerate(forward_batch.tbo_children) + ] + + # Non-EP DP TP-MoE: the per-ubatch DP gather/combine (op_gather/op_combine) + # needs each ubatch's per-rank token counts, but tbo_padded_len is computed + # per-rank locally (not synced). All-gather both ubatches' padded lengths + # once across DP ranks, then populate each child's global_num_tokens + + # global_dp_buffer_len so the gatherv/reduce_scatterv buffers size correctly. + if get_moe_a2a_backend().is_none() and get_parallel().attn_dp_size > 1: + tp_group = get_tp_group() + world = tp_group.world_size + children = forward_batch.tbo_children + local_lens = torch.tensor( + [int(c.tbo_padded_len) for c in children], + dtype=torch.int64, + device=hidden_states.device, + ) + gathered = torch.empty( + (world, local_lens.shape[0]), + dtype=torch.int64, + device=hidden_states.device, + ) + tp_group.all_gather_into_tensor(gathered, local_lens) + gathered_cpu = gathered.tolist() + rank = tp_group.rank_in_group + for idx, child in enumerate(children): + sizes = [gathered_cpu[r][idx] for r in range(world)] + child.global_num_tokens_cpu = sizes + child.global_num_tokens_gpu = gathered[:, idx].contiguous() + child.global_dp_buffer_len = sum(sizes) + # Gather the ubatch's input_ids -> global ONCE here (cached on the + # child) instead of per-layer in op_gather_a. The hash MoE reads + # the SAME global ids every layer, so 61x2 per-layer all_gatherv of + # VARYING size (-> RCCL registers a new internal buffer per size -> + # HSA_STATUS_ERROR_OUT_OF_RESOURCES) collapses to 1 per ubatch. + local_ids = child.input_ids + rows = sizes[rank] + if local_ids.shape[0] < rows: + padded_ids = local_ids.new_zeros((rows,)) + padded_ids[: local_ids.shape[0]] = local_ids + elif local_ids.shape[0] > rows: + padded_ids = local_ids[:rows] + else: + padded_ids = local_ids + gids = torch.empty( + (sum(sizes),), dtype=local_ids.dtype, device=local_ids.device + ) + tp_group.all_gatherv(padded_ids, sizes=sizes, output=gids) + child._tbo_global_input_ids = gids + + outputs_arr = execute_overlapped_operations( + inputs_arr=inputs_arr, + operations_arr=[operations_strategy.operations] * 2, + delta_stages=[0, operations_strategy.tbo_delta_stages], + ) + + hidden_states, _ = _model_forward_tbo_merge_outputs( + outputs_arr[0], outputs_arr[1], hidden_states.shape[0] + ) + return hidden_states + + def _setup_child_cp_metadata(self, child: ForwardBatch, child_backend) -> None: + cp_rank = get_parallel().attn_cp_rank + cp_size = get_parallel().attn_cp_size + child.attn_cp_metadata = prepare_context_parallel_metadata( + len(child.input_ids), + cp_rank, + cp_size, + child.seq_lens_cpu.tolist(), + extend_seqs_len=child.extend_seq_lens_cpu, + ) + if is_dsa_prefill_cp_round_robin_split(): + metadata = child_backend.forward_metadata + core_meta = metadata.core_attn_metadata + core_meta.apply_cp_reindex() + core_meta.init_flashmla_related(is_prefill=True) + if metadata.indexer_metadata is not None: + metadata.indexer_metadata = child_backend.init_forward_metadata_indexer( + core_meta + ) + + def _forward_layers_tbo_cp( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + assert _is_hip, "CP+TBO prefill path is HIP-only" + + from sglang.srt.batch_overlap.operations import execute_overlapped_operations + from sglang.srt.batch_overlap.operations_strategy import OperationsStrategy + from sglang.srt.batch_overlap.two_batch_overlap import ( + _model_forward_filter_inputs, + _model_forward_tbo_merge_outputs, + ) + + original_len = hidden_states.shape[0] + cp_size = get_parallel().attn_cp_size + layers = [self.layers[i] for i in range(self.start_layer, self.end_layer)] + operations_strategy = OperationsStrategy.init_new_tbo( + layers, forward_batch.global_forward_mode, use_cp=True + ) + + attn_backend = get_attn_backend() + children = forward_batch.tbo_children + # Attention-side CP gathers run two-phase (launch early on the comm + # stream / collect right before their consumer). Only the MoE + # collectives are splittable across a YieldOperation, so without this the + # ~2.5 attention-side collectives per layer would stay on the compute + # stream and defeat most of TBO's overlap. + prefetch_comm_stream = get_dp_tbo_comm_stream() + + inputs_arr = [] + for idx, child in enumerate(children): + child_inputs = _model_forward_filter_inputs( + hidden_states=hidden_states, + residual=None, + positions=positions, + output_forward_batch=child, + tbo_subbatch_index=idx, + ) + self._setup_child_cp_metadata(child, attn_backend.children[idx]) + if self.pp_group.is_first_rank: + child_inputs["hidden_states"] = cp_split_and_rebuild_data( + child, child_inputs["hidden_states"] + ) + child_inputs["positions"] = cp_split_and_rebuild_position( + child, child_inputs["positions"] + ) + child._cp_moe_input_ids = cp_round_robin_input_ids(child.input_ids) + child._cp_prefetch_comm_stream = prefetch_comm_stream + inputs_arr.append(child_inputs) + + outputs_arr = execute_overlapped_operations( + inputs_arr=inputs_arr, + operations_arr=[operations_strategy.operations] * 2, + delta_stages=[0, operations_strategy.tbo_delta_stages], + ) + + if self.pp_group.is_last_rank: + for idx, child in enumerate(children): + outputs_arr[idx]["hidden_states"] = cp_all_gather_rerange_output( + outputs_arr[idx]["hidden_states"], + cp_size, + child, + torch.cuda.current_stream(), + ) + hidden_states, _ = _model_forward_tbo_merge_outputs( + outputs_arr[0], outputs_arr[1], original_len + ) + return hidden_states + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor], + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> Union[torch.Tensor, PPProxyTensors]: + cp_v2_active = is_cp_v2_active(forward_batch) + use_prefill_cp = dsa_use_prefill_cp(forward_batch) + if self.pp_group.is_first_rank: + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds + hidden_states = hidden_states.unsqueeze(1).repeat(1, self.hc_mult, 1) + else: + assert pp_proxy_tensors is not None + hidden_states = pp_proxy_tensors["hidden_states"] + # Unflatten 2D PP IPC tensor back to 3D mHC shape. + if hidden_states.ndim == 2: + hidden_states = hidden_states.view( + hidden_states.shape[0], self.hc_mult, self.hidden_size + ) + + if get_parallel().attn_dp_size > 1 and get_moe_a2a_backend().is_none(): + input_ids_global = torch.empty( + (get_global_dp_buffer_len(), 1), + dtype=input_ids.dtype, + device=input_ids.device, + ) + # Token ids are replicated within an attention-TP group. Use replicate + # gather here to avoid summing duplicated ids when attention_tp_size > 1. + # Clone because the MAX_LEN gather may zero its local input in place. + dp_gather_replicate( + input_ids_global, input_ids[:, None].clone(), forward_batch + ) + input_ids_global = input_ids_global.squeeze(-1) + else: + input_ids_global = input_ids + + capture_dspark = self.dspark_layers_to_capture is not None + dspark_aux_hidden_states: List[torch.Tensor] = [] + # DSpark aux capture needs the per-layer eager loop (TBO's overlapped + # execution cannot expose per-layer completed hidden states), so skip + # TBO when capturing -- a perf-only downgrade, not a correctness one. + run_tbo = self._can_run_tbo(forward_batch) and not capture_dspark + if use_prefill_cp and not run_tbo: + if cp_v2_active: + input_ids = cp_round_robin_input_ids_v2(input_ids, forward_batch) + else: + if self.pp_group.is_first_rank: + hidden_states = cp_split_and_rebuild_data( + forward_batch, hidden_states + ) + positions = cp_split_and_rebuild_position(forward_batch, positions) + input_ids = cp_round_robin_input_ids(input_ids) + input_ids_global = input_ids + + # Reset Compressor's per-step freqs_cis cache from any previous step. + for _attr in ("freqs_cis_c4", "freqs_cis_c128"): + if hasattr(forward_batch, _attr): + delattr(forward_batch, _attr) + if run_tbo: + # Two-batch-overlap prefill (EP / mori). Cross-layer mHC fusion is + # disabled here (each layer self-contained), so no trailing hc_post. + hidden_states = self._forward_layers_tbo( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + else: + use_fused = self.use_fused_mhc_post_pre + prev_residual, prev_post, prev_comb = None, None, None + last_layer = None + for i in range(self.start_layer, self.end_layer): + layer = self.layers[i] + last_layer = layer + ctx = ( + nullcontext() + if check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) + else get_global_expert_distribution_recorder().with_current_layer(i) + ) + with ctx: + hidden_states, prev_residual, prev_post, prev_comb = layer( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + input_ids=input_ids, + input_ids_global=input_ids_global, + prev_residual=prev_residual, + prev_post=prev_post, + prev_comb=prev_comb, + ) + if capture_dspark and i in self.dspark_layers_to_capture: + if use_fused: + completed = layer.hc_post( + hidden_states, prev_residual, prev_post, prev_comb + ) + else: + completed = hidden_states + dspark_aux_hidden_states.append(completed.mean(dim=1)) + if use_fused and last_layer is not None: + hidden_states = last_layer.hc_post( + hidden_states, prev_residual, prev_post, prev_comb + ) + + # CP all-gather only on the last PP rank; PP IPC carries CP-split tensors. + if ( + self.pp_group.is_last_rank + and use_prefill_cp + and not cp_v2_active + and not run_tbo + ): + stream = torch.cuda.current_stream() + hidden_states = cp_all_gather_rerange_output( + hidden_states, + self.cp_size, + forward_batch, + stream, + ) + # Gather DSpark aux tensors on the same CP token split. + if capture_dspark: + dspark_aux_hidden_states = [ + cp_all_gather_rerange_output( + aux, self.cp_size, forward_batch, stream + ) + for aux in dspark_aux_hidden_states + ] + + if not self.pp_group.is_last_rank: + # Flatten 3D mHC tensor for PP IPC. + return PPProxyTensors({"hidden_states": hidden_states.flatten(1)}) + + pre_hc_head = hidden_states.flatten(1) + + hidden_states = self.hc_head( + hidden_states, self.hc_head_fn, self.hc_head_scale, self.hc_head_base + ) + hidden_states = self.norm(hidden_states) + + if capture_dspark: + return (hidden_states, pre_hc_head), dspark_aux_hidden_states + + return hidden_states, pre_hc_head + + +class DeepseekV4ForCausalLM(nn.Module): + def __init__( + self, + config: DeepSeekV4Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + # DeepseekV4 enables, by default, the CK w8a8-block GEMM (MLA proj) and the + # batched/contiguous-load rope kernels (faster on gfx95; . + # Module-level toggles default OFF; flipped True here for DSV4 + if _is_hip: + from sglang.kernels.ops.attention.deepseek_v4_rope import set_batched_rope + from sglang.srt.layers.quantization.fp8_utils import set_force_ck_w8a8 + + set_force_ck_w8a8(True) + set_batched_rope(True) + self.config = config + self.tp_size = get_parallel().tp_size + self.quant_config = quant_config + self.determine_num_fused_shared_experts() + self.model = DeepseekV4Model( + config, quant_config, prefix=add_prefix("model", prefix) + ) + self.pp_group = get_pp_group() + if self.pp_group.is_last_rank: + if self.pp_group.world_size == 1 and config.tie_word_embeddings: + self.lm_head = self.model.embed_tokens + else: + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + use_attn_tp_group=get_parallel().enable_dp_lm_head, + ) + else: + self.lm_head = PPMissingLayer() + self.logits_processor = LogitsProcessor(config) + self.capture_aux_hidden_states = False + get_attn_tp_context().init_context(config.q_lora_rank, is_dsa=True) + + self._routed_experts_weights_of_layer = LazyValue( + lambda: { + layer_id: self.model.layers[layer_id].mlp.get_moe_weights() + for layer_id in range(self.model.start_layer, self.model.end_layer) + if isinstance( + self.model.layers[layer_id].mlp, deepseek_v2.DeepseekV2MoE + ) + } + ) + + # Expose start_layer/end_layer for model_runner PP support + self.start_layer = self.model.start_layer + self.end_layer = self.model.end_layer + + self.dsa_enable_prefill_cp = is_dsa_enable_prefill_cp() + if self.dsa_enable_prefill_cp: + self.cp_rank = get_parallel().attn_cp_rank + self.cp_size = get_parallel().attn_cp_size + + # update_weights_from_disk/_tensor/_distributed re-enter load_weights + # mid-serving (RL refit sends many partial batches); the prewarm and + # its barrier must only run on the first (startup) load. + self._mhc_prewarmed_at_load = False + + @property + def routed_experts_weights_of_layer(self): + return self._routed_experts_weights_of_layer.value + + def get_input_embeddings(self) -> nn.Module: + return self.model.get_input_embeddings() + + def set_dspark_layers_to_capture(self, layer_ids: List[int]) -> None: + if not self.pp_group.is_last_rank: + return + if layer_ids is None: + raise ValueError( + "DSPARK requires explicit layer_ids for aux hidden capture." + ) + self.capture_aux_hidden_states = True + self.model.dspark_layers_to_capture = list(layer_ids) + + @classmethod + def shared_experts_fusion_disable_reason(cls, hf_config, quant_config): + """V4 only fuses when explicitly asked to, and then the checkpoint must + carry exactly one shared expert. Asked by the loader before any layer is + built.""" + # Need to disable if quant precision mismatch, even if + # --enforce-shared-experts-fusion is specified + if quant_blocks_shared_experts_fusion(quant_config): + return ( + "Quantization keeps shared experts at a higher precision than the " + "routed experts, so they cannot be fused into the quantized " + "routed-expert path." + ) + if get_parallel().moe_ep_size > 1 and not uses_per_rank_fused_shared_slots(): + return ( + "Expert parallelism keeps only a slice of the routed experts on " + "each rank, so the fused shared expert cannot be appended to the " + "routed weight tensor (only DeepEP/MegaMOE per-rank shared slots " + "support fusion under EP)." + ) + if not get_exec().moe.enforce_shared_experts_fusion: + return "Config does not support fused shared expert(s)." + if hf_config.n_shared_experts != 1: + raise ValueError( + "DeepSeek V4 shared-experts fusion expects exactly one shared " + f"expert, but got n_shared_experts={hf_config.n_shared_experts}." + ) + return None + + def determine_num_fused_shared_experts(self): + # The decision was installed by the loader; this only reads it. + self.num_fused_shared_experts = ( + 0 if is_shared_experts_fusion_disabled() else self.config.n_shared_experts + ) + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> torch.Tensor: + if self.dsa_enable_prefill_cp: + if can_dsa_cp_split(len(input_ids), self.cp_size, True, forward_batch): + forward_batch.attn_cp_metadata = prepare_context_parallel_metadata( + len(input_ids), + self.cp_rank, + self.cp_size, + forward_batch.seq_lens_cpu.tolist(), + extend_seqs_len=forward_batch.extend_seq_lens_cpu, + ) + if is_dsa_prefill_cp_round_robin_split(): + attn_backend = get_attn_backend() + metadata = attn_backend.forward_metadata + core_meta = metadata.core_attn_metadata + core_meta.apply_cp_reindex() + core_meta.init_flashmla_related(is_prefill=True) + if metadata.indexer_metadata is not None: + metadata.indexer_metadata = ( + attn_backend.init_forward_metadata_indexer(core_meta) + ) + + with get_attn_tp_context().maybe_input_scattered(forward_batch): + hidden_states = self.model.forward( + input_ids, positions, forward_batch, input_embeds, pp_proxy_tensors + ) + if not self.pp_group.is_last_rank: + return hidden_states + + aux_hidden_states = None + if self.capture_aux_hidden_states: + hidden_states, aux_hidden_states = hidden_states + hidden_states, pre_hc_head = hidden_states + + return self.logits_processor( + input_ids, + hidden_states, + self.lm_head, + forward_batch, + aux_hidden_states, + hidden_states_before_norm=( + None if aux_hidden_states is not None else pre_hc_head + ), + ) + + def _setup_fp8_wo_a_scales(self, is_nextn: bool) -> None: + from sglang.srt.layers import deep_gemm_wrapper + + if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0: + from deep_gemm import transform_sf_into_required_layout + + if is_nextn: + layers = [self.model.decoder] + else: + layers = [ + self.model.layers[layer_id] + for layer_id in range(self.model.start_layer, self.model.end_layer) + ] + for layer in layers: + attn = layer.self_attn + G = attn.n_local_groups + R = attn.o_lora_rank + D = attn.wo_a.weight.shape[1] + + raw_scale = attn.wo_a.weight_scale_inv.data.view(G, R // 128, D // 128) + if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0: + attn.wo_a.weight_scale_inv.data = transform_sf_into_required_layout( + raw_scale, + mn=R, + k=D, + recipe=(1, 128, 128), + num_groups=G, + is_sfa=False, + ) + attn.wo_a.weight_scale_inv.format_ue8m0 = True + else: + attn.wo_a.weight_scale_inv.data = raw_scale.contiguous() + attn.wo_a.weight_scale_inv.format_ue8m0 = False + + def post_load_weights(self, is_nextn=False, weight_names=None): + if _FP8_WO_A_GEMM: + self._setup_fp8_wo_a_scales(is_nextn) + + if is_nextn: + return + for layer_id in range(self.model.start_layer, self.model.end_layer): + layer = self.model.layers[layer_id] + self_attn = layer.self_attn + if ( + self_attn.compress_ratio in (4, 128) + and not self_attn.compressor.ape_converted + ): + self_attn.compressor.apply_ape_hotfix() + if ( + self_attn.compress_ratio == 4 + and not self_attn.indexer.compressor.ape_converted + ): + self_attn.indexer.compressor.apply_ape_hotfix() + layer.refresh_mhc_norm_weight_cache() + + @staticmethod + def remap_weight_name_to_dpsk_hf_format( + name: str, + is_nextn: bool = False, + num_hidden_layers: Optional[int] = None, + ) -> str: + if name.startswith("embed."): + return "model.embed_tokens." + name.removeprefix("embed.") + if name.startswith("head."): + return "lm_head." + name.removeprefix("head.") + if name == "norm.weight": + return "model.norm.weight" + if name.startswith("hc_head_"): + return "model." + name + + if is_nextn and name.startswith("mtp."): + parts = name.split(".", 2) + if len(parts) >= 3: + rest = parts[2] + nextn_spec_prefixes = [ + "e_proj", + "h_proj", + "emb", + "enorm", + "hnorm", + "norm", + "head", + "hc_head", + ] + is_nextn_spec = any(rest.startswith(p) for p in nextn_spec_prefixes) + if is_nextn_spec: + if rest.startswith("emb.tok_emb"): + rest = rest.replace("emb.tok_emb", "embed_tokens") + elif rest == "norm.weight": + rest = "shared_head.norm.weight" + elif rest.startswith("head."): + rest = "shared_head.head.weight" + elif rest == "e_proj.scale": + rest = "e_proj.weight_scale_inv" + elif rest == "h_proj.scale": + rest = "h_proj.weight_scale_inv" + name = f"model.layers.{num_hidden_layers}." + rest + + if name.startswith("layers."): + name = "model." + name + name = name.replace(".attn.", ".self_attn.") + name = name.replace(".ffn.", ".mlp.") + name = name.replace(".attn_norm.", ".input_layernorm.") + name = name.replace(".ffn_norm.", ".post_attention_layernorm.") + + if "self_attn" in name and name.endswith(".scale"): + name = name.removesuffix(".scale") + ".weight_scale_inv" + + name = name.replace(".gate.tid2eid", ".topk.tid2eid") + name = name.replace(".gate.bias", ".gate.e_score_correction_bias") + name = name.replace(".w1.", ".gate_proj.") + name = name.replace(".w2.", ".down_proj.") + name = name.replace(".w3.", ".up_proj.") + if "mlp" in name and name.endswith(".scale"): + name = name.removesuffix(".scale") + ".weight_scale_inv" + + return name + + def _prewarm_mhc_kernels(self) -> None: + """One-shot MHC JIT prewarm at load time, synced across ranks. + + Runs before any forward so the compile burst stays off the serving + path; the barrier keeps ranks from proceeding while a peer is still + compiling. The early returns below must stay rank-uniform. + """ + if self._mhc_prewarmed_at_load: + return + self._mhc_prewarmed_at_load = True + if _is_npu or not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): + return + layer = next( + (m for m in self.model.layers if isinstance(m, DeepseekV4DecoderLayer)), + None, + ) + if layer is None: + return + + from sglang.kernels.ops.layernorm.mhc import mhc_post, prewarm_mhc_pre + + tic = time.perf_counter() + residual = torch.zeros( + (1, layer.hc_mult, layer.hidden_size), + dtype=torch.bfloat16, + device=layer.hc_attn_fn.device, + ) + prewarm_mhc_pre( + # Template carrying dtype/device; buckets allocate their own sizes. + residual=residual, + fn=layer.hc_attn_fn, + hc_scale=layer.hc_attn_scale, + hc_base=layer.hc_attn_base, + rms_eps=layer.rms_norm_eps, + hc_pre_eps=layer.hc_eps, + hc_sinkhorn_eps=layer.hc_eps, + hc_post_mult_value=_MHC_POST_MULT_VALUE, + sinkhorn_repeat=layer.hc_sinkhorn_iters, + n_splits=1, + n_splits_pre=32, + norm_weight=layer.input_layernorm.weight.data, + norm_eps=layer.input_layernorm.variance_epsilon, + ) + mhc_post( + x=residual.new_zeros((1, layer.hidden_size)), + residual=residual, + post_layer_mix=torch.zeros( + (1, layer.hc_mult, 1), + dtype=torch.float32, + device=residual.device, + ), + comb_res_mix=torch.zeros( + (1, layer.hc_mult, layer.hc_mult), + dtype=torch.float32, + device=residual.device, + ), + ) + torch.cuda.synchronize() + compile_secs = time.perf_counter() - tic + # Runs before init_memory_pool(); don't let transients skew pool sizing. + torch.cuda.empty_cache() + get_tp_group().barrier() + logger.info( + "DeepSeek V4 MHC prewarm at load: compile %.1fs, rank sync +%.1fs", + compile_secs, + time.perf_counter() - tic - compile_secs, + ) + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]], is_nextn=False): + params_dict = dict(self.named_parameters()) + loaded_params: Set[str] = set() + + if is_nextn: + if hasattr(self.config, "num_nextn_predict_layers"): + num_nextn_layers = self.config.num_nextn_predict_layers + assert num_nextn_layers == 1, "Only 1 nextn layer is supported" + nextn_layer_id = ( + 0 + if self.config.num_hidden_layers == 1 + else self.config.num_hidden_layers + ) + else: + raise ValueError("num_nextn_predict_layers is not in the config") + + if not _FP8_WO_A_GEMM: + weights = _prepare_deepseek_v4_weights(weights, self.quant_config) + + stacked_params_mapping = DEEPSEEK_V4_STACKED_PARAMS_MAPPING + + expert_params_mapping = FusedMoE.make_expert_params_mapping( + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts + self.num_fused_shared_experts, + ) + + if is_wint4afp8_or_wint4a16_config(self.quant_config): + expert_params_mapping += FusedMoE.make_expert_input_scale_params_mapping( + num_experts=self.config.n_routed_experts + ) + + cache_compressor_weight = {} + COMPRESSOR_PART = ".compressor.w" + + fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get() + cache_wqkv_a_weight: dict[str, dict[str, torch.Tensor]] = {} + + def auto_weight_loader(module): + return getattr(module, "weight_loader", default_weight_loader) + + if is_nextn: + nextn_layer_prefix = f"model.layers.{nextn_layer_id}" + nextn_spec_weight_names_out_of_layer = [ + "shared_head.norm", + "shared_head.head", + "embed_tokens", + ".e_proj", + "h_proj", + "enorm", + "hnorm", + "hc_head_base", + "hc_head_fn", + "hc_head_scale", + ] + + if self.num_fused_shared_experts > 0: + assert self.num_fused_shared_experts == 1 + log_info_on_rank0(logger, "Shared experts fusion optimization enabled.") + + with concurrent.futures.ThreadPoolExecutor() as executor: + futures = [] + weight_names = [] + for name, loaded_weight in weights: + if ( + _FP8_WO_A_GEMM + and name.endswith(".wo_a.weight") + and loaded_weight.dtype != torch.float8_e4m3fn + ): + raise ValueError( + f"SGLANG_OPT_FP8_WO_A_GEMM is enabled but {name} has " + f"dtype {loaded_weight.dtype}, expected " + "torch.float8_e4m3fn. This checkpoint does not provide " + "a supported fp8-quantized wo_a; rerun with " + "SGLANG_OPT_FP8_WO_A_GEMM=0." + ) + try: + use_async_loading = should_async_load(loaded_weight) + + name = self.remap_weight_name_to_dpsk_hf_format( + name, + is_nextn=is_nextn, + num_hidden_layers=self.config.num_hidden_layers, + ) + + layer_id = get_layer_id(name) + if ( + layer_id is not None + and hasattr(self.model, "start_layer") + and ( + layer_id < self.model.start_layer + or layer_id >= self.model.end_layer + ) + ): + continue + if ( + self.num_fused_shared_experts > 0 + and "mlp.shared_experts" in name + ): + name = name.replace( + "mlp.shared_experts", + f"mlp.experts.{self.config.n_routed_experts}", + ) + + weight_names.append(name) + + if not is_nextn: + if hasattr(self.config, "num_nextn_predict_layers"): + num_nextn_layers = self.config.num_nextn_predict_layers + if num_nextn_layers > 0 and name.startswith("model.layers"): + name_list = name.split(".") + if ( + len(name_list) >= 3 + and int(name_list[2]) + >= self.config.num_hidden_layers + ): + continue + + if name.startswith("mtp"): + continue + else: + if "shared_head.head" in name or "embed_tokens" in name: + continue + + if not name.startswith(nextn_layer_prefix): + continue + + in_decoder = True + for weight_name in nextn_spec_weight_names_out_of_layer: + if weight_name in name: + in_decoder = False + name = name.replace(nextn_layer_prefix, "model") + break + + if in_decoder: + name = name.replace(nextn_layer_prefix, "model.decoder") + + if "rotary_emb.inv_freq" in name: + continue + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + if _is_npu: + name = name.replace("weight_packed", "weight") + if ("mlp.experts." in name) and name not in params_dict: + continue + name = name.replace(weight_name, param_name) + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict and name.startswith("mtp"): + break + param = params_dict[name] + weight_loader = param.weight_loader + maybe_executor_submit( + executor=executor, + futures=futures, + use_async=use_async_loading, + func=weight_loader, + func_args=(param, loaded_weight, shard_id), + ) + loaded_params.add(name) + break + else: + skip_unmaterialized_expert_param = False + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping + if weight_name not in name: + continue + if _is_npu: + name = name.replace("weight_packed", "weight") + resolved_name = name.replace(weight_name, param_name) + if resolved_name not in params_dict: + skip_unmaterialized_expert_param = True + continue + param = params_dict[resolved_name] + weight_loader = param.weight_loader + maybe_executor_submit( + executor=executor, + futures=futures, + use_async=use_async_loading, + func=weight_loader, + func_args=( + param, + loaded_weight, + resolved_name, + ), + func_kwargs={ + "shard_id": shard_id, + "expert_id": expert_id, + }, + ) + loaded_params.add(resolved_name) + break + else: + if skip_unmaterialized_expert_param: + continue + if name.endswith(".bias") and name not in params_dict: + continue + if ( + ".embed_tokens." in name + and not self.pp_group.is_first_rank + ): + continue + if ( + name == "model.norm.weight" + and not self.pp_group.is_last_rank + ): + continue + if ( + name.startswith("model.hc_head_") + or name == "lm_head.weight" + ) and not self.pp_group.is_last_rank: + continue + elif COMPRESSOR_PART in name and ".wkv_gate." not in name: + is_kv = name.endswith(".wkv.weight") + is_wgate = name.endswith(".wgate.weight") + assert is_kv != is_wgate + key = name.rsplit(".", 2)[0] + assert key.endswith(".compressor") + if key not in cache_compressor_weight: + cache_compressor_weight[key] = ( + is_kv, + _clone_if_runai_streamed_tensor(loaded_weight), + ) + else: + assert key in cache_compressor_weight + cached_is_kv, cached_weight = ( + cache_compressor_weight[key] + ) + assert cached_is_kv != is_kv + kv = loaded_weight if is_kv else cached_weight + wgate = loaded_weight if is_wgate else cached_weight + fused_weight = torch.cat([kv, wgate], dim=0) + param_name = key + ".wkv_gate.weight" + param = params_dict[param_name] + weight_loader = auto_weight_loader(param) + maybe_executor_submit( + executor=executor, + futures=futures, + use_async=use_async_loading, + func=weight_loader, + func_args=(param, fused_weight), + ) + loaded_params.add(param_name) + cache_compressor_weight.pop(key) + elif fuse_wqa_wkv and ( + name.endswith(".wq_a.weight") + or name.endswith(".wq_a.weight_scale_inv") + or name.endswith(".wkv.weight") + or name.endswith(".wkv.weight_scale_inv") + or name.endswith(".wq_a.qweight") + or name.endswith(".wkv.qweight") + or name.endswith(".wq_a.qweight_type") + or name.endswith(".wkv.qweight_type") + ): + is_q = ".wq_a." in name + param_name = name.replace( + ".wq_a." if is_q else ".wkv.", ".wqkv_a." + ) + bucket = cache_wqkv_a_weight.setdefault(param_name, {}) + shard_key = "q" if is_q else "kv" + assert ( + shard_key not in bucket + ), f"duplicate shard {shard_key} for {param_name}" + bucket[shard_key] = _clone_if_runai_streamed_tensor( + loaded_weight + ) + if len(bucket) == 2: + fused_weight = _fuse_deepseek_v4_wqkv_a_pair( + param_name, bucket + ) + param = params_dict[param_name] + weight_loader = auto_weight_loader(param) + maybe_executor_submit( + executor=executor, + futures=futures, + use_async=use_async_loading, + func=weight_loader, + func_args=(param, fused_weight), + ) + loaded_params.add(param_name) + cache_wqkv_a_weight.pop(param_name) + else: + if ( + "k_scale" in name or "v_scale" in name + ) and name not in params_dict: + for scale in ["k_scale", "v_scale"]: + if scale in name: + name = name.replace( + f"{scale[0]}_proj", "attn_mqa" + ) + break + if name not in params_dict: + if not name.startswith("mtp"): + logger.warning( + f"{name} not found in params_dict." + ) + continue + param = params_dict[name] + + weight_loader = auto_weight_loader(param) + maybe_executor_submit( + executor=executor, + futures=futures, + use_async=use_async_loading, + func=weight_loader, + func_args=(param, loaded_weight), + ) + loaded_params.add(name) + except Exception as e: + e.add_note(f"{name=} {loaded_weight.shape=}") + raise + + for future in concurrent.futures.as_completed(futures): + future.result() + + assert len(cache_compressor_weight) == 0 + assert len(cache_wqkv_a_weight) == 0, cache_wqkv_a_weight.keys() + unloaded_params = params_dict.keys() - loaded_params + + skipped_checking_patterns = [ + "attn_mqa.k_scale", + "attn_mqa.v_scale", + "blockscale_swizzled", + ] + if not self.pp_group.is_first_rank: + skipped_checking_patterns.append("embed_tokens") + if not self.pp_group.is_last_rank: + skipped_checking_patterns.append("model.norm.") + skipped_checking_patterns.extend(["lm_head", "hc_head_"]) + if is_nextn: + skipped_checking_patterns.extend(["lm_head", "embed_tokens"]) + unloaded_params = { + p + for p in unloaded_params + if all( + skipped_checking_pattern not in p + for skipped_checking_pattern in skipped_checking_patterns + ) + } + if unloaded_params: + logger.warning( + f"Some weights are not initialized from checkpoints: {unloaded_params}" + ) + + self.post_load_weights(is_nextn=is_nextn, weight_names=weight_names) + + if not is_nextn: + self._prewarm_mhc_kernels() + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + del self.lm_head.weight + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + # Hot weight reload (RL workflows). Use the device-agnostic module + # accessor so this works on both CUDA/HIP and NPU. + torch.get_device_module().empty_cache() + torch.get_device_module().synchronize() + + @classmethod + def get_model_config_for_expert_location(cls, config): + return ModelConfigForExpertLocation( + num_layers=config.num_hidden_layers, + num_logical_experts=config.n_routed_experts, + num_groups=None, + ) + + +EntryClass = [DeepseekV4ForCausalLM] + + +def _dequant_fp8(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + from einops import rearrange + + assert ( + weight.dtype == torch.float8_e4m3fn + ), f"expected fp8_e4m3fn, got {weight.dtype}" + assert scale.dtype in ( + torch.float8_e8m0fnu, + torch.float32, + ), f"expected fp8_e8m0fnu or float32, got {scale.dtype}" + + weight_f32 = rearrange( + weight.float(), "(sn bn) (sk bk) -> sn bn sk bk", bn=128, bk=128 + ) + result = rearrange( + weight_f32 * scale.float()[:, None, :, None], "sn bn sk bk -> (sn bn) (sk bk)" + ) + + return result.to(torch.bfloat16) + + +def _clone_if_runai_streamed_tensor(tensor: torch.Tensor) -> torch.Tensor: + if getattr(tensor, RUNAI_STREAMER_TENSOR_ATTR, False): + return tensor.clone().detach() + return tensor + + +def _dequant_fp8_wo_a_streaming( + weights: Iterable[Tuple[str, torch.Tensor]], +) -> Iterable[Tuple[str, torch.Tensor]]: + pending: dict[str, dict[str, torch.Tensor]] = {} + saw_wo_a_scale = False + emitted = False + + for name, tensor in weights: + if name.endswith(".wo_a.weight"): + prefix = name[: -len(".weight")] + bucket = pending.setdefault(prefix, {}) + scale = bucket.pop("scale", None) + if scale is not None: + pending.pop(prefix, None) + emitted = True + yield name, _dequant_fp8(tensor, scale) + else: + bucket["weight"] = _clone_if_runai_streamed_tensor(tensor) + continue + + if name.endswith(".wo_a.scale"): + saw_wo_a_scale = True + prefix = name[: -len(".scale")] + bucket = pending.setdefault(prefix, {}) + weight = bucket.pop("weight", None) + if weight is not None: + pending.pop(prefix, None) + emitted = True + yield prefix + ".weight", _dequant_fp8(weight, tensor) + else: + bucket["scale"] = _clone_if_runai_streamed_tensor(tensor) + continue + + yield name, tensor + + if emitted: + logger.info("Finished streaming dequant fp8 wo_a") + for prefix, bucket in pending.items(): + if "weight" in bucket: + assert not saw_wo_a_scale, f"{prefix}.scale is missing" + yield prefix + ".weight", bucket["weight"] + if "scale" in bucket: + yield prefix + ".scale", bucket["scale"] + + +def _dequant_fp8_wo_a( + weights: Iterable[Tuple[str, torch.Tensor]], +) -> Iterable[Tuple[str, torch.Tensor]]: + weights_dict = dict(weights) + + for name in list(weights_dict.keys()): + if name not in weights_dict: + continue + if not name.endswith(".wo_a.weight"): + continue + scale_name = name.replace(".wo_a.weight", ".wo_a.scale") + assert scale_name in weights_dict + weight = weights_dict.pop(name) + scale = weights_dict.pop(scale_name) + yield name, _dequant_fp8(weight, scale) + + yield from weights_dict.items() + + +def _prepare_deepseek_v4_weights( + weights: Iterable[Tuple[str, torch.Tensor]], + quant_config: Optional[QuantizationConfig], +) -> Iterable[Tuple[str, torch.Tensor]]: + """Keep Expert Pack GGUF weights on the streaming load path.""" + + if quant_config is not None and quant_config.get_name() == "expert_pack": + logger.info("Keep Expert Pack GGUF weights on the streaming load path") + return weights + return _dequant_fp8_wo_a_streaming(weights) + + +def _fuse_deepseek_v4_wqkv_a_pair( + param_name: str, bucket: dict[str, torch.Tensor] +) -> torch.Tensor: + """Fuse Q/KV rows while preserving their common GGUF type scalar.""" + + q = bucket["q"] + kv = bucket["kv"] + if param_name.endswith(".qweight_type"): + if q.numel() != 1 or kv.numel() != 1 or q.item() != kv.item(): + raise ValueError( + f"cannot fuse different GGUF qweight types for {param_name}: " + f"q={q.tolist()} kv={kv.tolist()}" + ) + return q + return torch.cat([q, kv], dim=0) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/manifest.json b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/manifest.json new file mode 100644 index 0000000..34ac278 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/manifest.json @@ -0,0 +1,7 @@ +{ + "baseline_sha256": "76e83c26696a3f3ddb49ff75824b02035a1c57712a4b324e65373bc8fe5f82d3", + "candidate_sha256": "e44b83a6248c06603ef9e3b10d081765c21eef819ec4a4fd00ba4653aa215663", + "status": "kernel_and_e2e_validation_pending", + "source_commit": "0bcd822377da7b5718e674eaf9c870d349424dd1", + "change": "SM120 FlashInfer: native local heads for both q allocation and attention sink; other backends unchanged" +} diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/patch.diff b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/patch.diff new file mode 100644 index 0000000..e220a13 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/native_heads/patch.diff @@ -0,0 +1,43 @@ +--- a/deepseek_v4.py ++++ b/deepseek_v4.py +@@ -757,13 +757,24 @@ + self.register_buffer("freqs_cis", freqs_cis, persistent=False) + self.freqs_cis: torch.Tensor + ++ def _padded_attn_heads(self) -> int: ++ # FlashInfer SM120 specializes native TP-local heads; the generic ++ # FlashMLA path still requires padding to 64 or 128. ++ if ( ++ get_platform().is_sm120 ++ and envs.SGLANG_SM120_FLASHMLA_BACKEND.get() == "flashinfer" ++ and self.n_local_heads in (8, 16, 32, 64, 128) ++ ): ++ return self.n_local_heads ++ return 64 if self.n_local_heads <= 64 else self.n_heads ++ + def _local_attn_sink(self) -> torch.Tensor: + if self.attn_tp_size == 1: + return self.attn_sink + if self._attn_sink_local is None: + rank = self.attn_tp_rank + num_heads = self.n_local_heads +- padded_num_heads = 64 if num_heads <= 64 else self.n_heads ++ padded_num_heads = self._padded_attn_heads() + sink = self.attn_sink.new_zeros(padded_num_heads) + sink[:num_heads] = self.attn_sink[rank * num_heads : (rank + 1) * num_heads] + self._attn_sink_local = sink +@@ -1573,11 +1584,9 @@ + + tp_slice, q_padded, q_out = slice(None), None, None + if self.attn_tp_size > 1: +- # FlashMLA's fp8 sparse decode kernel only specializes h_q for {64, 128}. +- # Pad the per-rank heads to 64 (not the full n_heads) when they fit, to +- # dispatch the cheaper decode::head64 variant; attn_sink is sliced to +- # this rank and padded to match. +- padded_num_heads = 64 if self.n_local_heads <= 64 else self.n_heads ++ # Match the query allocation and sliced sink to this backend. ++ # FlashInfer SM120 uses native heads; generic FlashMLA keeps padding. ++ padded_num_heads = self._padded_attn_heads() + # Only [0:n_local_heads] is written below. Uninitialized padded TP + # heads inject NaN into attention on gfx942 (fnuz), so zero-init + # there; other archs tolerate new_empty and skip the per-forward diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/flash_mla_sm120.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/flash_mla_sm120.py new file mode 100644 index 0000000..f8bf967 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/flash_mla_sm120.py @@ -0,0 +1,641 @@ +"""SM120 FlashMLA sparse decode implementation. + +On SM120 (Blackwell Desktop / RTX PRO 6000) the flash_mla CUDA kernel +is not available, so this module provides alternative implementations: + +- A fused Triton kernel (default, ``SGLANG_SM120_TRITON_FLASHMLA=1``) +- A pure-PyTorch fallback (``SGLANG_SM120_TRITON_FLASHMLA=0``) + +The FP8 KV cache uses a page-internal layout where NOPE+ROPE data has +stride (nope_dim + rope_dim*2) per token, and scales are stored in a +separate region at the end of each page. +""" + +import logging +import math +from typing import Optional + +import torch +import triton +import triton.language as tl + +from sglang.srt.environ import envs +from sglang.srt.utils import is_hip + +logger = logging.getLogger(__name__) +_is_hip = is_hip() + +_GLM_DSA_MODEL_ARCHS = ( + "GlmMoeDsaForCausalLM", + "GlmMoeDsaForCausalLMNextN", +) + +# Page layout constants for DSv4-Flash (MODEL1): +# nope_dim = 448, rope_dim = 64, quantize_block_size = 64 +# nope_rope_stride = 448 + 64*2 = 576 bytes per token +# scale_stride = ceil(448/64) + 1 = 8 bytes per token (7 scales + 1 pad) +# bytes_per_token = 448 + 128 + 8 = 584 +# page_bytes = ceil_div(page_size * 584, 576) * 576 + +_NOPE_DIM = 448 +_ROPE_DIM = 64 +_NOPE_ROPE_STRIDE = _NOPE_DIM + _ROPE_DIM * 2 # 576 +_TILE_SIZE = 64 +_NUM_TILES = _NOPE_DIM // _TILE_SIZE # 7 +_SCALE_STRIDE = _NUM_TILES + 1 # 8 (7 scales + 1 pad) +_D = _NOPE_DIM + _ROPE_DIM # 512 + + +def _gather_and_dequant(k_cache, indices, page_size): + """Gather KV entries from the paged buffer using correct page-internal addressing. + + Args: + k_cache: (num_pages, page_size, 1, bytes_per_token) float8_e4m3fn + Non-contiguous view of the raw page buffer. + indices: (...) int32/int64, token-level indices. -1 = invalid. + page_size: tokens per page (256) + + Returns: + kv: (..., _D) bfloat16, dequantized KV vectors + """ + idx_shape = indices.shape + flat_idx = indices.reshape(-1) # (N,) + N = flat_idx.shape[0] + device = k_cache.device + + # Page-level addressing + page_bytes = k_cache.stride(0) # actual byte stride between pages + pages = flat_idx // page_size + offsets = flat_idx % page_size + + # Clamp invalid indices + safe_pages = pages.clamp(min=0) + safe_offsets = offsets.clamp(min=0) + + # Access raw buffer as uint8 — use as_strided to get full page view + num_pages = k_cache.shape[0] + raw_pages = k_cache.as_strided( + (num_pages, page_bytes), + (page_bytes, 1), + ).view( + torch.uint8 + ) # (num_pages, page_bytes) uint8 + # Note: float8_e4m3fn and uint8 are both 1 byte, view is safe + + # Compute byte offsets within each page + # NOPE: page[safe_page, safe_offset * 576 + 0:448] + # ROPE: page[safe_page, safe_offset * 576 + 448:576] + # SCALES: page[safe_page, page_size * 576 + safe_offset * 8 + 0:7] + + nope_base = safe_offsets * _NOPE_ROPE_STRIDE # (N,) + nope_offsets = nope_base.unsqueeze(-1) + torch.arange( + _NOPE_DIM, device=device, dtype=torch.long + ) # (N, 448) + + rope_base = nope_base + _NOPE_DIM # (N,) + rope_offsets = rope_base.unsqueeze(-1) + torch.arange( + _ROPE_DIM * 2, device=device, dtype=torch.long + ) # (N, 128) + + scale_section_offset = page_size * _NOPE_ROPE_STRIDE # 147456 + scale_base = scale_section_offset + safe_offsets * _SCALE_STRIDE # (N,) + scale_offsets = scale_base.unsqueeze(-1) + torch.arange( + _NUM_TILES, device=device, dtype=torch.long + ) # (N, 7) + + # Gather bytes per page — use advanced indexing + # raw_pages[safe_pages, nope_offsets] → (N, 448) + page_idx_nope = safe_pages.unsqueeze(-1).expand_as(nope_offsets) + nope_bytes = raw_pages[page_idx_nope, nope_offsets] # (N, 448) uint8 + + page_idx_rope = safe_pages.unsqueeze(-1).expand_as(rope_offsets) + rope_bytes = raw_pages[page_idx_rope, rope_offsets] # (N, 128) uint8 + + page_idx_scale = safe_pages.unsqueeze(-1).expand_as(scale_offsets) + scale_bytes = raw_pages[page_idx_scale, scale_offsets] # (N, 7) uint8 + + # Reinterpret dtypes + nope_fp8 = nope_bytes.view(torch.float8_e4m3fn) # (N, 448) + rope_bf16 = rope_bytes.contiguous().view(torch.bfloat16) # (N, 64) + scale_e8m0 = scale_bytes.view(torch.float8_e8m0fnu) # (N, 7) + + # Dequantize: nope_tile * scale_tile → bf16 (vectorized) + result = torch.empty(N, _D, dtype=torch.bfloat16, device=device) + result[:, :_NOPE_DIM] = ( + ( + nope_fp8.view(N, _NUM_TILES, _TILE_SIZE).float() + * scale_e8m0.view(N, _NUM_TILES, 1).float() + ) + .view(N, _NOPE_DIM) + .to(torch.bfloat16) + ) + result[:, _NOPE_DIM:] = rope_bf16 + + return result.reshape(*idx_shape, _D) + + +def _sm120_sparse_decode_fwd( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache=None, + extra_indices=None, + extra_topk_length=None, +): + B, s_q, H_q, D_qk = q.shape + num_pages, page_size, H_k, bpt = k_cache.shape + topk = indices.shape[-1] + + invalid_mask = indices < 0 + safe_indices = indices.clamp(min=0) + + if topk_length is not None: + topk_range = torch.arange(topk, device=topk_length.device).view(1, 1, topk) + invalid_mask = invalid_mask | (topk_range >= topk_length.view(B, 1, 1)) + + # Gather and dequantize using page-aware addressing + gathered_kv = _gather_and_dequant(k_cache, safe_indices, page_size) + + if extra_k_cache is not None and extra_indices is not None: + extra_topk = extra_indices.shape[-1] + extra_page_size = extra_k_cache.shape[1] + extra_invalid = extra_indices < 0 + extra_safe = extra_indices.clamp(min=0) + if extra_topk_length is not None: + extra_range = torch.arange( + extra_topk, device=extra_topk_length.device + ).view(1, 1, extra_topk) + extra_invalid = extra_invalid | ( + extra_range >= extra_topk_length.view(B, 1, 1) + ) + extra_kv = _gather_and_dequant(extra_k_cache, extra_safe, extra_page_size) + gathered_kv = torch.cat([gathered_kv, extra_kv], dim=2) + invalid_mask = torch.cat([invalid_mask, extra_invalid], dim=2) + + gathered_kv[invalid_mask] = 0.0 + + q_f = q.float() + kv_f = gathered_kv.float() + kv_d = kv_f.shape[-1] + if D_qk != kv_d: + q_f = q_f[..., :kv_d] + + scores = torch.einsum("bshd,bstd->bsht", q_f, kv_f) * softmax_scale + scores.masked_fill_(invalid_mask.unsqueeze(2).expand_as(scores), float("-inf")) + + lse = torch.logsumexp(scores, dim=-1) + + if attn_sink is not None: + lse_for_out = torch.logsumexp( + torch.stack([lse, attn_sink.view(1, 1, H_q).expand_as(lse)], dim=0), dim=0 + ) + else: + lse_for_out = lse.clone() + + lonely = lse == float("-inf") + lse_for_out[lonely] = float("inf") + weights = torch.exp(scores - lse_for_out.unsqueeze(-1)) + out = torch.einsum("bsht,bstv->bshv", weights, kv_f[..., :head_dim_v]) + out[lonely.unsqueeze(-1).expand_as(out)] = 0.0 + + return out.to(torch.bfloat16), lse.permute(0, 2, 1) + + +# SM120 FlashMLA: default FlashInfer (CUTLASS SM120 sparse MLA decode). +# Override with SGLANG_SM120_FLASHMLA_BACKEND=triton|torch to force fallback. +_sm120_default_backend = envs.SGLANG_SM120_FLASHMLA_BACKEND.get() + + +def flash_mla_with_kvcache_sm120(**kwargs): + """SM120 FlashMLA sparse decode entry point. + + Dispatches to FlashInfer (default if available), Triton, or PyTorch fallback. + """ + q = kwargs["q"] + k_cache = kwargs["k_cache"] + indices = kwargs["indices"] + topk_length = kwargs.get("topk_length") + attn_sink = kwargs.get("attn_sink") + head_dim_v = kwargs["head_dim_v"] + softmax_scale = kwargs.get("softmax_scale") + if softmax_scale is None: + softmax_scale = q.shape[-1] ** (-0.5) + extra_k_cache = kwargs.get("extra_k_cache") + extra_indices = kwargs.get("extra_indices_in_kvcache") + extra_topk_length = kwargs.get("extra_topk_length") + + if _sm120_default_backend == "flashinfer": + return _flash_mla_flashinfer( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + + if _sm120_default_backend == "triton": + from sglang.kernels.ops.attention.flash_mla_sm120_triton import ( + flash_mla_sparse_decode_triton, + ) + + out, lse = flash_mla_sparse_decode_triton( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + return (out, lse) + + out, lse = _sm120_sparse_decode_fwd( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, + ) + return (out, lse) + + +# --- Page-split utilities: pbs=256 → pbs=64 --- +# SGLang SWA KV cache footer layout per 256-token page: +# [data: 256 * 576 bytes] [scale: 256 * 8 bytes] [padding] +# FlashInfer decode_dsv4 expects per 64-token page: +# [data: 64 * 576 bytes] [scale: 64 * 8 bytes] [padding to 37440] +_PBS_SRC = 256 # SGLang physical page size +_PBS_DST = 64 # FlashInfer page_block_size +_NOPE_ROPE_STRIDE = 576 # bytes per token for nope+rope +_SCALE_STRIDE = 8 # bytes per token for scale (7 + 1 pad) +_BYTES_PER_DST_PAGE = ( + _PBS_DST * _NOPE_ROPE_STRIDE + _PBS_DST * _SCALE_STRIDE +) # 64*576 + 64*8 = 37376 + 512 = 37888 +# Padded to 576 alignment + +_BYTES_PER_DST_PAGE_PADDED = math.ceil(_BYTES_PER_DST_PAGE / 576) * 576 # 37440 + + +@triton.jit +def _page_split_kernel( + src_ptr, + dst_ptr, + N_pages, + src_stride0: tl.constexpr, + dst_stride0: tl.constexpr, + DATA_PER_SUB: tl.constexpr, # 64 * 576 = 36864 + SCALE_PER_SUB: tl.constexpr, # 64 * 8 = 512 + SRC_SCALE_OFF: tl.constexpr, # 256 * 576 = 147456 + DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864 + RATIO: tl.constexpr, # 4 + BLOCK_SIZE: tl.constexpr, + mask_ptr, + HAS_MASK: tl.constexpr, +): + """Fused page-split: copy data+scale for all sub-pages in one kernel. + + When HAS_MASK is set, only pages flagged in ``mask_ptr`` (int8, 1=touched) + are copied; untouched pages are skipped so the kernel no longer rewrites the + entire KV pool every decode step. + """ + pid = tl.program_id(0) + page_idx = pid // RATIO + sub = pid % RATIO + + if page_idx >= N_pages: + return + + if HAS_MASK: + if tl.load(mask_ptr + page_idx) == 0: + return + + src_base = src_ptr + page_idx * src_stride0 + dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0 + + # Copy data region: DATA_PER_SUB bytes from src offset sub*DATA_PER_SUB + data_src_off = sub * DATA_PER_SUB + for start in tl.range(0, DATA_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < DATA_PER_SUB + vals = tl.load(src_base + data_src_off + offs, mask=mask) + tl.store(dst_base + offs, vals, mask=mask) + + # Copy scale region: SCALE_PER_SUB bytes + scale_src_off = SRC_SCALE_OFF + sub * SCALE_PER_SUB + for start in tl.range(0, SCALE_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < SCALE_PER_SUB + vals = tl.load(src_base + scale_src_off + offs, mask=mask) + tl.store(dst_base + DST_SCALE_OFF + offs, vals, mask=mask) + + +@triton.jit +def _page_mark_kernel( + indices_ptr, + mask_ptr, + N_idx, + SRC_PBS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Mark touched source pages (1 byte each) from token-level indices. + + ``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid. + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + idx = tl.load(indices_ptr + offsets, mask=offsets < N_idx, other=-1) + valid = (offsets < N_idx) & (idx >= 0) + page = tl.maximum(idx, 0) // SRC_PBS + tl.store(mask_ptr + page, 1, mask=valid) + + +def _split_kv_pages_to_64( + kv_u8: torch.Tensor, + src_pbs: int, + touched_indices: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Split pbs=N footer-format pages into pbs=64 footer-format pages. + + When ``touched_indices`` (token-level int32 indices into the pbs=src_pbs + SWA pool, -1 = invalid) is provided, only the source pages that actually + contain a referenced token are copied. This avoids rewriting the entire KV + pool on every decode step (only ~2*batch pages are touched vs the full + pool). The output buffer is persistent and reused across steps; untouched + dst pages simply retain their (unreferenced) stale data. + """ + assert src_pbs % _PBS_DST == 0 and src_pbs >= _PBS_DST + if src_pbs == _PBS_DST: + return kv_u8 + + N = kv_u8.shape[0] + ratio = src_pbs // _PBS_DST + num_dst_pages = N * ratio + + from sglang.srt.runtime_context import get_resources + + # Pre-allocated grow-only buffer for page-split output per device. + dev = kv_u8.device + buffers = get_resources().buffers + key = f"flash_mla_sm120_split:{dev}" + buf = buffers.get(key) + if buf is None or buf.shape[0] < num_dst_pages: + # The first allocation can happen under inference mode (autotune), but + # the buffer is written again during CUDA graph capture outside + # inference mode, where an inference tensor cannot be mutated. + with torch.inference_mode(False): + buf = torch.empty( + num_dst_pages, + _BYTES_PER_DST_PAGE_PADDED, + dtype=torch.uint8, + device=dev, + ) + buffers[key] = buf + out = buf[:num_dst_pages] + + # Get raw 2D view of source + src_2d = kv_u8 + if src_2d.ndim == 4: + src_stride0 = src_2d.stride(0) + src_2d = torch.as_strided(src_2d, (N, src_stride0), (src_stride0, 1)) + else: + src_stride0 = src_2d.stride(0) + + use_mask = touched_indices is not None and touched_indices.numel() > 0 + mask_ptr = src_2d # dummy, never dereferenced when HAS_MASK is False + if use_mask: + # Persistent per-device int8 mask, zeroed each call (cheap memset, + # captured cleanly by CUDA graph). 1 = page is referenced this step. + mkey = f"flash_mla_sm120_mask:{dev}" + mbuf = buffers.get(mkey) + if mbuf is None or mbuf.shape[0] < N: + # The first allocation can happen under inference mode (autotune), + # but the buffer is zeroed again later during CUDA graph capture + # outside inference mode -- an inference tensor cannot be mutated + # there, so force a normal tensor. + with torch.inference_mode(False): + mbuf = torch.empty(N, dtype=torch.int8, device=dev) + buffers[mkey] = mbuf + mask = mbuf[:N] + mask.zero_() + idx_flat = touched_indices.reshape(-1).contiguous() + if idx_flat.dtype != torch.int32: + idx_flat = idx_flat.to(torch.int32) + _page_mark_kernel[(triton.cdiv(idx_flat.numel(), 256),)]( + idx_flat, + mask, + idx_flat.numel(), + src_pbs, # SRC_PBS + 256, # vectorized indices per program + ) + mask_ptr = mask + + grid = (N * ratio,) + _page_split_kernel[grid]( + src_2d, + out, + N, + src_stride0, + _BYTES_PER_DST_PAGE_PADDED, + _PBS_DST * _NOPE_ROPE_STRIDE, # DATA_PER_SUB = 36864 + _PBS_DST * _SCALE_STRIDE, # SCALE_PER_SUB = 512 + src_pbs * _NOPE_ROPE_STRIDE, # SRC_SCALE_OFF = 147456 + _PBS_DST * _NOPE_ROPE_STRIDE, # DST_SCALE_OFF = 36864 + ratio, # RATIO = 4 + 1024, # BLOCK_SIZE + mask_ptr, + use_mask, # HAS_MASK + ) + + bpt = _NOPE_ROPE_STRIDE + _SCALE_STRIDE # 584 + return out.as_strided( + (num_dst_pages, _PBS_DST, 1, bpt), + (_BYTES_PER_DST_PAGE_PADDED, bpt, bpt, 1), + ) + + +def _flash_mla_flashinfer( + q, + k_cache, + indices, + topk_length, + attn_sink, + head_dim_v, + softmax_scale, + extra_k_cache, + extra_indices, + extra_topk_length, +): + """FlashInfer SM120 sparse MLA via the paged-attention dispatcher. + + SGLang SWA pool uses page_size=256 (footer format: 256*576 bytes data + 256*8 bytes scale). + FlashInfer decode_dsv4 fast path requires page_block_size=64 (footer: 64*576 + 64*8). + We split 256-token pages into 4 virtual 64-token pages. + Token indices are invariant under page-split (identity mapping). + """ + from flashinfer.mla._sparse_mla_sm120 import ( + _DECODE_MAX_TOKENS as _FI_DECODE_MAX_TOKENS, + ) + from flashinfer.mla._sparse_mla_sm120 import ( + _sparse_mla_sm120_paged_attention, + ) + + B, _, H, D = q.shape # (batch, 1, num_heads, head_dim) + dev = q.device + + # Indices: no remapping needed (page-split preserves token addressing). + idx = indices.squeeze(1) if indices.dim() == 3 else indices + + # --- Page-split: convert pbs=N kv_cache to pbs=64 view --- + # Only the SWA pages actually referenced by `idx` are copied (the rest of + # the persistent dst buffer is left untouched and never read). + kv_u8 = k_cache.view(torch.uint8) if k_cache.dtype != torch.uint8 else k_cache + src_pbs = k_cache.shape[1] if k_cache.ndim >= 3 else _PBS_SRC + kv_64 = ( + _split_kv_pages_to_64(kv_u8, src_pbs, touched_indices=idx) + if src_pbs != _PBS_DST + else kv_u8 + ) + + extra_kv_u8 = ( + extra_k_cache.view(torch.uint8) + if extra_k_cache is not None and extra_k_cache.dtype != torch.uint8 + else extra_k_cache + ) + extra_kv_64 = extra_kv_u8 + + extra_idx = ( + extra_indices.squeeze(1) + if extra_indices is not None and extra_indices.dim() == 3 + else extra_indices + ) + + output = torch.empty(B, H, head_dim_v, dtype=torch.bfloat16, device=dev) + out_lse = torch.empty(B, H, dtype=torch.float32, device=dev) + + # Use split-K for decode-sized batches and paged attention otherwise. + if B <= _FI_DECODE_MAX_TOKENS: + topk = idx.shape[-1] + extra_topk = extra_idx.shape[-1] if extra_idx is not None else 0 + _BI = 64 + num_splits = (topk + _BI - 1) // _BI + ( + (extra_topk + _BI - 1) // _BI if extra_topk > 0 else 0 + ) + mid_out = torch.empty( + B, H, num_splits, head_dim_v, dtype=torch.bfloat16, device=dev + ) + mid_lse = torch.empty(B, H, num_splits, dtype=torch.float32, device=dev) + else: + mid_out = None + mid_lse = None + + _sparse_mla_sm120_paged_attention( + q.squeeze(1) if q.ndim == 4 else q, + kv_64, + idx, + output, + out_lse, + softmax_scale, + d_v=head_dim_v, + topk_length=topk_length, + attn_sink=attn_sink, + extra_kv_cache=extra_kv_64, + extra_indices=extra_idx, + extra_topk_length=extra_topk_length, + mid_out=mid_out, + mid_lse=mid_lse, + ) + + return (output.unsqueeze(1), None) + + +def _validate_flashinfer_sparse_mla_backend( + *, + model_arch: str, + device_sm_major: int, + kv_cache_dtype: torch.dtype, + prefill_impl: str, + decode_impl: str, +) -> bool: + selected = {prefill_impl, decode_impl} + uses_flashinfer_sparse_mla = "flashinfer_sparse_mla" in selected + is_glm_sm12_fp8 = ( + model_arch in _GLM_DSA_MODEL_ARCHS + and device_sm_major == 12 + and kv_cache_dtype == torch.float8_e4m3fn + and not _is_hip + ) + if uses_flashinfer_sparse_mla and not is_glm_sm12_fp8: + raise ValueError( + "flashinfer_sparse_mla supports only GLM DSA with FP8 KV cache " + "on NVIDIA SM120/SM121; " + f"got model_arch={model_arch!r}, sm_major={device_sm_major}, " + f"kv_cache_dtype={kv_cache_dtype}, prefill_impl={prefill_impl!r}, " + f"decode_impl={decode_impl!r}." + ) + if is_glm_sm12_fp8: + unsupported = selected - {"flashinfer_sparse_mla"} + if unsupported: + raise ValueError( + "GLM DSA with FP8 KV cache on NVIDIA SM120/SM121 supports " + "only flashinfer_sparse_mla, " + f"but got {sorted(unsupported)}." + ) + return uses_flashinfer_sparse_mla + + +def flashinfer_sparse_mla_forward( + q: torch.Tensor, + kv_cache: torch.Tensor, + indices: torch.Tensor, + seq_lens: torch.Tensor, + workspace_buffer: torch.Tensor, + *, + page_size: int, + kv_cache_dim: int, + qk_nope_head_dim: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + sm_scale: float, + skip_softmax_threshold_scale_factor: float | None, +) -> torch.Tensor: + """Run FlashInfer's SM120 sparse MLA kernel on SGLang's packed DSA cache.""" + from flashinfer.mla import trtllm_batch_decode_with_kv_cache_mla + + topk = indices.shape[1] + result = trtllm_batch_decode_with_kv_cache_mla( + query=q.unsqueeze(1), + kv_cache=kv_cache.view(torch.uint8) + .view(-1, page_size, kv_cache_dim) + .unsqueeze(1), + workspace_buffer=workspace_buffer, + qk_nope_head_dim=qk_nope_head_dim, + kv_lora_rank=kv_lora_rank, + qk_rope_head_dim=qk_rope_head_dim, + block_tables=indices.unsqueeze(1), + seq_lens=seq_lens, + max_seq_len=topk, + sparse_mla_top_k=topk, + bmm1_scale=float(sm_scale), + bmm2_scale=1.0, + kv_scale_format="arbitrary_fp32", + skip_softmax_threshold_scale_factor=skip_softmax_threshold_scale_factor, + ) + return result.squeeze(1) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/kernels.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/kernels.py new file mode 100644 index 0000000..52a9350 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/kernels.py @@ -0,0 +1,97 @@ +import triton +import triton.language as tl + +@triton.jit +def mark_original( + indices_ptr, + mask_ptr, + N_idx, + SRC_PBS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Mark touched source pages (1 byte each) from token-level indices. + + ``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid. + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ + pid = tl.program_id(0) + if pid >= N_idx: + return + idx = tl.load(indices_ptr + pid) + if idx < 0: + return + page = idx // SRC_PBS + tl.store(mask_ptr + page, 1) + +@triton.jit +def mark_vectorized( + indices_ptr, + mask_ptr, + N_idx, + SRC_PBS: tl.constexpr, + BLOCK: tl.constexpr, +): + """Mark touched source pages (1 byte each) from token-level indices. + + ``indices`` are token indices into the pbs=SRC_PBS SWA pool; -1 = invalid. + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + idx = tl.load(indices_ptr + offsets, mask=offsets < N_idx, other=-1) + valid = (offsets < N_idx) & (idx >= 0) + page = tl.maximum(idx, 0) // SRC_PBS + tl.store(mask_ptr + page, 1, mask=valid) + +@triton.jit +def _page_split_kernel( + src_ptr, + dst_ptr, + N_pages, + src_stride0: tl.constexpr, + dst_stride0: tl.constexpr, + DATA_PER_SUB: tl.constexpr, # 64 * 576 = 36864 + SCALE_PER_SUB: tl.constexpr, # 64 * 8 = 512 + SRC_SCALE_OFF: tl.constexpr, # 256 * 576 = 147456 + DST_SCALE_OFF: tl.constexpr, # 64 * 576 = 36864 + RATIO: tl.constexpr, # 4 + BLOCK_SIZE: tl.constexpr, + mask_ptr, + HAS_MASK: tl.constexpr, +): + """Fused page-split: copy data+scale for all sub-pages in one kernel. + + When HAS_MASK is set, only pages flagged in ``mask_ptr`` (int8, 1=touched) + are copied; untouched pages are skipped so the kernel no longer rewrites the + entire KV pool every decode step. + """ + pid = tl.program_id(0) + page_idx = pid // RATIO + sub = pid % RATIO + + if page_idx >= N_pages: + return + + if HAS_MASK: + if tl.load(mask_ptr + page_idx) == 0: + return + + src_base = src_ptr + page_idx * src_stride0 + dst_base = dst_ptr + (page_idx * RATIO + sub) * dst_stride0 + + # Copy data region: DATA_PER_SUB bytes from src offset sub*DATA_PER_SUB + data_src_off = sub * DATA_PER_SUB + for start in tl.range(0, DATA_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < DATA_PER_SUB + vals = tl.load(src_base + data_src_off + offs, mask=mask) + tl.store(dst_base + offs, vals, mask=mask) + + # Copy scale region: SCALE_PER_SUB bytes + scale_src_off = SRC_SCALE_OFF + sub * SCALE_PER_SUB + for start in tl.range(0, SCALE_PER_SUB, BLOCK_SIZE): + offs = start + tl.arange(0, BLOCK_SIZE) + mask = offs < SCALE_PER_SUB + vals = tl.load(src_base + scale_src_off + offs, mask=mask) + tl.store(dst_base + DST_SCALE_OFF + offs, vals, mask=mask) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/manifest.json b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/manifest.json new file mode 100644 index 0000000..faf2a70 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/manifest.json @@ -0,0 +1,8 @@ +{ + "image": "sha256:d6e7288627be8b02be88e4bba38e73f6d50e2826869f753c13a4c4385ab3eda9", + "source_commit": "0bcd822377da7b5718e674eaf9c870d349424dd1", + "baseline_sha256": "39f0f98151a7cfd750b987d82cf05fafe80e8e972ef53a2b78352ce9b472e9b5", + "candidate_sha256": "83b1cc400593b339d4c241a9062fa1ce022e08e3ff0514a02df80aad0c1ee9e3", + "change": "Vectorize page-mark indices in groups of 256; unchanged same-value stores and page split", + "status": "unvalidated_candidate" +} diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/patch.diff b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/patch.diff new file mode 100644 index 0000000..b9e0194 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_vectorized/patch.diff @@ -0,0 +1,37 @@ +--- a/flash_mla_sm120.py ++++ b/flash_mla_sm120.py +@@ -360,14 +360,11 @@ + Each valid token marks ``mask[token // SRC_PBS] = 1``. Concurrent stores of + the same value 1 are safe (no atomic needed). + """ +- pid = tl.program_id(0) +- if pid >= N_idx: +- return +- idx = tl.load(indices_ptr + pid) +- if idx < 0: +- return +- page = idx // SRC_PBS +- tl.store(mask_ptr + page, 1) ++ offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) ++ idx = tl.load(indices_ptr + offsets, mask=offsets < N_idx, other=-1) ++ valid = (offsets < N_idx) & (idx >= 0) ++ page = tl.maximum(idx, 0) // SRC_PBS ++ tl.store(mask_ptr + page, 1, mask=valid) + + + def _split_kv_pages_to_64( +@@ -441,12 +438,12 @@ + idx_flat = touched_indices.reshape(-1).contiguous() + if idx_flat.dtype != torch.int32: + idx_flat = idx_flat.to(torch.int32) +- _page_mark_kernel[(idx_flat.numel(),)]( ++ _page_mark_kernel[(triton.cdiv(idx_flat.numel(), 256),)]( + idx_flat, + mask, + idx_flat.numel(), + src_pbs, # SRC_PBS +- 1024, # BLOCK (unused, kept for JIT signature) ++ 256, # vectorized indices per program + ) + mask_ptr = mask + diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_verdict_evidence.md b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_verdict_evidence.md new file mode 100644 index 0000000..87038a5 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/page_mark_verdict_evidence.md @@ -0,0 +1,44 @@ +# page-mark 向量化补丁对 GLM-5.3-NVFP4 判死:双证据 + +DSV4 报告的 E2 优化点「`_page_mark_kernel` 标量→向量化(BLOCK=1024)」在 DSV4-Flash 上 +实测 prefill +3.0%。评估迁移到 GLM-5.3-NVFP4 时,静态 + 运行时双证据证明 GLM 路径 +**根本不执行该 kernel**,补丁零增益,实验臂在压测前撤销(省一整轮部署+扫描)。 + +## 证据一(静态,调用链闭环) + +目标文件 `flash_mla_sm120.py`(两镜像 nightly-dev-20260828 / latest md5 完全相同: +`9df4d50d44df9b9f6db36538f2d8b52f`,基础版已归档为本目录 `base_flash_mla_sm120.py`)。 + +全文件仅有的调用点(grep 实证): + +- `:444` `_page_mark_kernel[...]` 唯一发射点,位于 `_split_kv_pages_to_64`(`:373` 定义)内 +- `:515` `_split_kv_pages_to_64(...)` 唯一调用点,位于 `_flash_mla_flashinfer`(`:477` 定义)内 +- `:607` `flashinfer_sparse_mla_forward` —— **GLM(GlmMoeDsaForCausalLM)的入口** + +`flashinfer_sparse_mla_forward`(`:607-644`)函数体直接调用 +`flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla`,入参 +`kv_cache.view(-1, page_size, kv_cache_dim)`、`block_tables=indices.unsqueeze(1)`, +**不经过 `_split_kv_pages_to_64`、不碰 `_page_mark_kernel`**。 + +即:page-mark / page-split 整套机制是 DSV4-Flash 的 SWA(滑动窗口)页分裂路径 +(`_flash_mla_flashinfer`)专用;GLM-5.3 走 sparse-MLA 直接入口,与该路径无交集。 + +## 证据二(运行时,Triton 零启动) + +`_page_mark_kernel` 是 Triton kernel。若 GLM serving 路径会执行它,容器内必然产生 +Triton JIT 缓存或 Triton kernel 加载。实测(60.1,方案 D 配置容器): + +- 容器环境无 `TRITON_CACHE_DIR` 覆盖,检查 `/root/.triton`、`/root/.cache/triton`、`/tmp` + ——GLM serving 全程 **0 个 Triton kernel 编译/加载** +- 对比:同机 DSV4-Flash 容器相同路径下有大量 Triton 缓存 + +结论:GLM-5.3-NVFP4 的 serving 热路径不 launch 任何 Triton kernel, +`_page_mark_kernel` 从未执行 —— 与证据一互相印证。 + +## 归档内容 + +- `base_flash_mla_sm120.py`:两镜像共有的基础模块(md5 9df4d50d…) +- `page_mark_vectorized/`:DSV4 E2 形态的向量化补丁(BLOCK=1024, mask other=-1、 + idx>=0 幂等写 1、grid=cdiv(N,1024)),对 GLM 无效但留作 DSV4 侧资产 +- `native_heads/`:P0 优化点,同理判 N/A —— GLM 已原生直通 flashinfer + (8/16/32-head 模板齐全),`deepseek_v4.py` 补丁对 GLM 无增益未实验 diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/__init__.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/__init__.py new file mode 100644 index 0000000..1defd26 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/__init__.py @@ -0,0 +1,123 @@ +from .cuda_ipc import CudaRTLibrary, create_shared_buffer, free_shared_buffer +from .dlpack_utils import pack_strided_memory +from .mapping import Mapping +from .trtllm_ar import AllReduceFusionOp as AllReduceFusionOp +from .trtllm_ar import AllReduceFusionPattern as AllReduceFusionPattern +from .trtllm_ar import AllReduceStrategyConfig as AllReduceStrategyConfig +from .trtllm_ar import AllReduceStrategyType as AllReduceStrategyType +from .trtllm_ar import QuantizationSFLayout as QuantizationSFLayout +from .trtllm_ar import ( + compute_fp4_swizzled_layout_sf_size as compute_fp4_swizzled_layout_sf_size, +) +from .trtllm_ar import gen_trtllm_comm_module as gen_trtllm_comm_module +from .trtllm_ar import trtllm_allreduce_fusion as trtllm_allreduce_fusion +from .trtllm_ar import ( + trtllm_create_ipc_workspace_for_all_reduce as trtllm_create_ipc_workspace_for_all_reduce, +) +from .trtllm_ar import ( + trtllm_create_ipc_workspace_for_all_reduce_fusion as trtllm_create_ipc_workspace_for_all_reduce_fusion, +) +from .trtllm_ar import trtllm_custom_all_reduce as trtllm_custom_all_reduce +from .trtllm_ar import ( + trtllm_destroy_ipc_workspace_for_all_reduce as trtllm_destroy_ipc_workspace_for_all_reduce, +) +from .trtllm_ar import ( + trtllm_destroy_ipc_workspace_for_all_reduce_fusion as trtllm_destroy_ipc_workspace_for_all_reduce_fusion, +) +from .trtllm_ar import trtllm_lamport_initialize as trtllm_lamport_initialize +from .trtllm_ar import trtllm_lamport_initialize_all as trtllm_lamport_initialize_all +from .trtllm_ar import trtllm_moe_allreduce_fusion as trtllm_moe_allreduce_fusion +from .trtllm_ar import ( + trtllm_moe_finalize_allreduce_fusion as trtllm_moe_finalize_allreduce_fusion, +) +from .vllm_ar import all_reduce as vllm_all_reduce +from .vllm_ar import dispose as vllm_dispose +from .vllm_ar import gen_vllm_comm_module as gen_vllm_comm_module +from .vllm_ar import get_graph_buffer_ipc_meta as vllm_get_graph_buffer_ipc_meta +from .vllm_ar import init_custom_ar as vllm_init_custom_ar +from .vllm_ar import meta_size as vllm_meta_size +from .vllm_ar import register_buffer as vllm_register_buffer +from .vllm_ar import register_graph_buffers as vllm_register_graph_buffers +from .ulysses import UlyssesCommunicator as UlyssesCommunicator +from .ulysses import dispose_ulysses_a2a as dispose_ulysses_a2a +from .ulysses import gen_ulysses_a2a_module as gen_ulysses_a2a_module +from .ulysses import get_ulysses_a2a_module as get_ulysses_a2a_module +from .ulysses import init_ulysses_a2a as init_ulysses_a2a +from .ulysses import ulysses_a2a as ulysses_a2a +from .ulysses_topology import ULYSSES_BACKENDS as ULYSSES_BACKENDS +from .ulysses_topology import UlyssesBackendDecision as UlyssesBackendDecision +from .ulysses_topology import UlyssesBackendError as UlyssesBackendError +from .ulysses_topology import UlyssesRankTopology as UlyssesRankTopology +from .ulysses_topology import decide_ulysses_backend as decide_ulysses_backend +from .ulysses_topology import ( + probe_ulysses_rank_topology as probe_ulysses_rank_topology, +) +from .ulysses_topology import resolve_ulysses_backend as resolve_ulysses_backend + +# Unified AllReduce Fusion API +from .allreduce import AllReduceFusionWorkspace as AllReduceFusionWorkspace +from .trtllm_mnnvl_ar import ( + MNNVLAllReduceFusionWorkspace as MNNVLAllReduceFusionWorkspace, +) +from .allreduce import TRTLLMAllReduceFusionWorkspace as TRTLLMAllReduceFusionWorkspace +from .allreduce import allreduce_fusion as allreduce_fusion +from .allreduce import ( + create_allreduce_fusion_workspace as create_allreduce_fusion_workspace, +) + +# MNNVL A2A (Throughput Backend) +from .trtllm_moe_alltoall import MoeAlltoAll as MoeAlltoAll +from .trtllm_moe_alltoall import moe_a2a_active_rank_mask as moe_a2a_active_rank_mask +from .trtllm_moe_alltoall import moe_a2a_combine as moe_a2a_combine +from .trtllm_moe_alltoall import moe_a2a_dispatch as moe_a2a_dispatch +from .trtllm_moe_alltoall import moe_a2a_initialize as moe_a2a_initialize +from .trtllm_moe_alltoall import ( + moe_a2a_get_workspace_size_per_rank as moe_a2a_get_workspace_size_per_rank, +) +from .trtllm_moe_alltoall import ( + moe_a2a_sanitize_expert_ids as moe_a2a_sanitize_expert_ids, +) +from .trtllm_moe_alltoall import ( + moe_a2a_wrap_payload_tensor_in_workspace as moe_a2a_wrap_payload_tensor_in_workspace, +) + +# DCP A2A (Decode Context Parallel Attention Reduction) +from .dcp_alltoall import decode_cp_a2a_alltoall as decode_cp_a2a_alltoall +from .dcp_alltoall import ( + decode_cp_a2a_allocate_mnnvl_workspace as decode_cp_a2a_allocate_mnnvl_workspace, +) +from .dcp_alltoall import decode_cp_a2a_init_workspace as decode_cp_a2a_init_workspace +from .dcp_alltoall import decode_cp_a2a_workspace_size as decode_cp_a2a_workspace_size + +# from .mnnvl import MnnvlMemory, MnnvlMoe, MoEAlltoallInfo + + +# SSKJ-PIE: PCIe-IPC all-reduce (flashinfer main PR #4393, vendored into 0.6.18) +from .pcie_ipc_ar import PcieIpcAllReduceWorkspace as PcieIpcAllReduceWorkspace +from .pcie_ipc_ar import gen_pcie_ipc_comm_module as gen_pcie_ipc_comm_module +from .pcie_ipc_ar import get_pcie_ipc_comm_module as get_pcie_ipc_comm_module +from .pcie_ipc_policy import IpcLaunchConfig as PcieIpcLaunchConfig +from .pcie_ipc_policy import IpcVariant as PcieIpcVariant +from .pcie_ipc_policy import ( + get_pcie_ipc_launch_config as get_pcie_ipc_launch_config, +) +from .pcie_ipc_topology import ( + probe_pcie_ipc_rank_topology as probe_pcie_ipc_rank_topology, +) +from .pcie_ipc_topology import ( + resolve_pcie_ipc_profile as resolve_pcie_ipc_profile, +) +from .pcie_ipc_tuning import PCIE_IPC_CUSTOM_OP as PCIE_IPC_CUSTOM_OP +from .pcie_ipc_tuning import default_cache_path as pcie_ipc_default_cache_path + + +def __getattr__(name: str): + if name == "all_gather_matmul": + from .all_gather_matmul import all_gather_matmul + + return all_gather_matmul + if name == "quantized_all_reduce": + from .quantized_allreduce import quantized_all_reduce + + return quantized_all_reduce + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_ar.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_ar.py new file mode 100644 index 0000000..39a85ad --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_ar.py @@ -0,0 +1,1006 @@ +""" +Copyright (c) 2026 by FlashInfer team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +import functools +import hashlib +import os +import warnings +from types import SimpleNamespace +from typing import Dict, List, Optional, Sequence, Tuple + +import torch +import torch.distributed as dist +from torch.distributed import ProcessGroup + +from ..api_logging import flashinfer_api + +try: # SSKJ-PIE: trace templates postdate the 0.6.18 release + from ..trace.templates.comm import pcie_ipc_all_reduce_trace +except ImportError: + pcie_ipc_all_reduce_trace = None +from ..jit.comm import gen_pcie_ipc_comm_module +from ..utils import register_custom_op +from .cuda_ipc import create_shared_buffer, free_shared_buffer +from .pcie_ipc_policy import IpcLaunchConfig, get_pcie_ipc_launch_config +from .pcie_ipc_topology import resolve_pcie_ipc_profile +from .pcie_ipc_tuning import ( + PCIE_IPC_CUSTOM_OP, + TUNE_BATCHES, + TUNE_REPEAT, + TUNE_WARMUP, + PcieIpcAllReduceRunner, + cache_covers_workspace, + default_cache_path, + pack_config, + pcie_ipc_tuning_config, + resolve_tuned_config, + tuned_batches_for, + warn_no_tune_group, +) + +_SUPPORTED_WORLD_SIZES = (2, 4, 8) +# Mirrors the launcher, which hard-checks a 2-byte element size: the kernels +# address whole 16-byte packs and are instantiated for half and nv_bfloat16 +# only. Rejecting here turns that into an unsupported shape rather than an +# ICHECK partway through a collective. +_SUPPORTED_DTYPES = (torch.bfloat16, torch.float16) + + +@functools.cache +def get_pcie_ipc_comm_module(): + module = gen_pcie_ipc_comm_module().build_and_load() + + @register_custom_op("flashinfer::pcie_ipc_workspace_size", mutates_args=[]) + def workspace_size( + world_size: int, max_numel: int, elem_size: int, max_blocks: int + ) -> int: + return module.pcie_ipc_workspace_size( + world_size, max_numel, elem_size, max_blocks + ) + + @register_custom_op("flashinfer::pcie_ipc_init", mutates_args=["ipc_ptrs"]) + def init( + ipc_ptrs: List[int], + rank: int, + max_numel: int, + elem_size: int, + max_blocks: int, + ) -> int: + return module.pcie_ipc_init(ipc_ptrs, rank, max_numel, elem_size, max_blocks) + + @register_custom_op("flashinfer::pcie_ipc_dispose", mutates_args=["handle"]) + def dispose(handle: int) -> None: + module.pcie_ipc_dispose(handle) + + @register_custom_op("flashinfer::pcie_ipc_all_reduce", mutates_args=["out"]) + def all_reduce( + handle: int, + inp: torch.Tensor, + out: torch.Tensor, + blocks: int, + threads: int, + variant: int, + enable_pdl: bool, + ) -> None: + module.pcie_ipc_all_reduce( + handle, inp, out, blocks, threads, variant, enable_pdl + ) + + return SimpleNamespace( + workspace_size=workspace_size, + init=init, + dispose=dispose, + all_reduce=all_reduce, + ) + + +class PcieIpcAllReduceWorkspace: + """Shared workspace for the PCIe IPC all-reduce. + + Allocates one slab per rank, shares it over CUDA IPC, and binds it to the + kernels. The workspace is sized once and cannot grow, so ``max_numel`` must + cover the largest collective that will be issued; anything larger must fall + back to another backend. + + This is a **collective**, and an unusually strict one. The kernels spin on + peer flags with no timeout and no metadata exchange, so every rank must + issue the same sequence of calls, with the same shape, dtype and launch + configuration, in the same order. A rank that skips a call, reorders two, + or passes a different explicit ``config`` does not get an error -- the + group hangs, or worse, one rank reads a neighbour's partial sums as if they + were finished. :meth:`launch_config` is a pure function of shape, dtype and + the workspace's own immutable attributes precisely so that every rank + derives the same answer without having to agree on one at runtime; passing + ``config`` explicitly moves that obligation to the caller. + + One workspace serves **one CUDA stream**. Its epoch and arrival counters + assume the calls sharing it are totally ordered, which stream order gives + and concurrent streams do not; the second stream is rejected. Build a + separate workspace per stream. + + Size ``max_numel`` to the real workload rather than to a round number. The + epoch double buffer places its two halves ``world_size * max_numel`` + elements apart, so an oversized workspace spreads them further than the + payload needs and costs measurable time at small batch. The multiplier is + the world size, not 2 -- rounding ``max_numel`` up by 4x at 8 ranks moves + the halves 32x the payload apart. + + Parameters + ---------- + group : ProcessGroup + Process group whose ranks share the workspace. Every rank must build + the workspace with identical arguments. + max_numel : int + Largest element count that will be all-reduced. + dtype : torch.dtype + bfloat16 or float16. Only the element *size* is binding, so one + workspace serves both. + max_blocks : int + Upper bound on the block count any launch may request. Sizes the + barrier and epoch slots. + profile : str, optional + Force the interconnect label (``"rootcplx"`` or ``"pcieswitch"``) + instead of probing for it. The label does not pick a kernel; it + partitions the tune cache so two topologies do not read each other's + measurements. Probing is collective and runs before any allocation. + tune_cache : str, optional + Where tuned configurations are read from at construction and written by + :meth:`tune`. Defaults to ``FLASHINFER_AUTOTUNE_DIR`` (or the workspace + directory). Give the same path to both, or a tuned result will not be + found by the next process. + + Launch configurations start from a seed default that is workable rather + than fast (see :mod:`~flashinfer.comm.pcie_ipc_policy`). Tune once to + replace it with measurements from this machine; the result is persisted and + later processes pick it up when the workspace is built. Tuning never changes + which shapes are supported, only which kernel a supported shape runs. + + Examples + -------- + >>> ws = PcieIpcAllReduceWorkspace(group=tp_group, max_numel=max_tokens * hidden) + >>> if ws.supports(x): + ... out = ws.all_reduce(x) + >>> ws.destroy() + + Tuning, once per machine: + + >>> ws.tune([hidden]) # collective; every rank calls it + """ + + def __init__( + self, + group: ProcessGroup, + max_numel: int, + dtype: torch.dtype = torch.bfloat16, + max_blocks: int = 128, + profile: Optional[str] = None, + tune_batches: Sequence[int] = TUNE_BATCHES, + tune_cache: Optional[str] = None, + ) -> None: + # Construction is a staged transaction. Every rank must execute the same + # sequence of collectives, so a rank that finds a problem does NOT raise + # where it finds it -- it records an outcome and raises only at the next + # gather, together with everyone else. Raising early would leave the + # peers blocked in a collective that their partner has already left. + self._ipc_ptrs: Optional[List[int]] = None + self._handle: Optional[int] = None + self.group = group + self.rank = dist.get_rank(group=group) + self.world_size = dist.get_world_size(group=group) + self.device = torch.device("cuda", torch.cuda.current_device()) + # Bound on first executing use; see _check_stream. + self._stream: Optional[torch.cuda.Stream] = None + self.elem_size = 0 + self.max_numel = max_numel + self.max_blocks = max_blocks + self.profile = "" + self.profile_reason = "" + # Resolved launch configurations, keyed exactly. Consulted before any + # AutoTuner call because even a pure cache lookup there takes a global + # lock, which is real overhead at this operator's scale. + self._tuned: Dict[Tuple[int, int, torch.dtype], IpcLaunchConfig] = {} + self._runner: Optional[PcieIpcAllReduceRunner] = None + self._tune_group: Optional[ProcessGroup] = None + self._tune_batches = tuple(int(b) for b in tune_batches) + self._tune_cache = tune_cache or default_cache_path(self.world_size) + self._tune_cache_exists = False + self._tuned_configs_loaded = False + self._warned_untuned = False + + # --- stage 1: local validation, encoded rather than raised ----------- + error: Optional[str] = None + if self.world_size not in _SUPPORTED_WORLD_SIZES: + error = ( + f"world size {self.world_size} unsupported; " + f"expected one of {_SUPPORTED_WORLD_SIZES}" + ) + elif dtype not in _SUPPORTED_DTYPES: + error = f"dtype {dtype} unsupported; expected one of {_SUPPORTED_DTYPES}" + else: + self.elem_size = torch.empty((), dtype=dtype).element_size() + pack_elems = 16 // self.elem_size + if max_numel <= 0 or max_numel % pack_elems != 0: + # The kernels address the scratch in 16-byte packs, so a + # capacity that is not a whole number of packs is rejected by + # the launcher on every call. Catch it here instead. + error = ( + f"max_numel must be a positive multiple of {pack_elems} " + f"for {dtype}, got {max_numel}" + ) + elif max_blocks <= 0: + error = f"max_blocks must be positive, got {max_blocks}" + + # Layout must be identical on every rank, or one of them reads a peer + # slab at the wrong offsets. Gather the config alongside the outcome so + # a single collective settles both. + local = { + "error": error, + "max_numel": max_numel, + "elem_size": self.elem_size, + "max_blocks": max_blocks, + "profile": profile, + # Buckets pick which shape a tuned entry is reused for, so ranks + # that disagree would resolve different configurations. + "tune_batches": self._tune_batches, + "tune_cache": self._tune_cache, + } + self._joint_check(local, "validating arguments") + + # --- stage 2: topology, then module + workspace size ----------------- + # Both before any allocation, so an unsupported topology or a failed + # JIT build costs nothing to unwind. + try: + decision = resolve_pcie_ipc_profile(group, requested=profile) + self.profile = decision.profile + self.profile_reason = decision.reason + module = get_pcie_ipc_comm_module() + nbytes = module.workspace_size( + self.world_size, max_numel, self.elem_size, max_blocks + ) + except Exception as e: # noqa: BLE001 - re-raised jointly below + nbytes = 0 + self._joint_check({"error": f"{type(e).__name__}: {e}"}, "preparing") + raise # unreachable: _joint_check raises on every rank + self._joint_check({"error": None}, "preparing") + + # --- stage 3: allocate and share, then bind -------------------------- + # NOTE: create_shared_buffer() runs its own all_gather_object and + # barrier internally. A failure *inside* it leaves the group in a state + # this constructor cannot repair; that is a property of the shared + # helper, not something worked around here. + self._ipc_ptrs = create_shared_buffer(nbytes, group=group) + bind_error: Optional[str] = None + try: + self._handle = module.init( + self._ipc_ptrs, self.rank, max_numel, self.elem_size, max_blocks + ) + # init() zeroes this rank's slab; no peer may push into it until + # every rank has done so. + torch.cuda.synchronize(self.device) + except Exception as e: # noqa: BLE001 - re-raised jointly below + bind_error = f"{type(e).__name__}: {e}" + + # Whether to tear down is a group decision: the cleanup itself contains + # barriers, so one rank must never enter it alone. + try: + self._joint_check({"error": bind_error}, "binding the workspace") + except Exception: + self.destroy() + raise + + # --- stage 4: tuned configurations, if any have been persisted ------- + # Loaded once, here, and never reloaded: a rank that picks up a file + # update its peers have not seen would choose a different kernel, and + # the group hangs rather than erroring. + try: + self._init_tuning() + except Exception: + self.destroy() + raise + dist.barrier(group=group) + + def _joint_check(self, local: dict, what: str) -> None: + """Gather per-rank outcomes and fail the whole group, or none of it. + + Raises the same error on every rank, so the caller can rely on all + ranks taking the same branch afterwards. + """ + gathered: List[Optional[dict]] = [None] * self.world_size + dist.all_gather_object(gathered, local, group=self.group) + entries = [g for g in gathered if g is not None] + + failed = {i: g["error"] for i, g in enumerate(entries) if g.get("error")} + if failed: + raise ValueError(f"pcie ipc workspace failed while {what}: {failed}") + + mismatched = { + key: [g[key] for g in entries] + for key in local + if key != "error" and len({repr(g[key]) for g in entries}) > 1 + } + if mismatched: + raise ValueError( + "every rank must build the workspace with identical arguments, " + f"but these differ across the group: {mismatched}" + ) + + @property + def handle(self) -> int: + if self._handle is None: + raise RuntimeError("workspace has been destroyed") + return self._handle + + def _check_stream(self) -> None: + """Bind the workspace to one stream, and reject use from another. + + The workspace carries mutable protocol state -- the epoch that selects + which half of the scratch a call stages through, and the arrival + counter that commits it. Both are advanced by the kernels themselves + and are only well defined if the calls that share this workspace are + totally ordered. Stream order gives that; two streams do not, and + concurrent calls would interleave their epoch reads and commits and + silently corrupt each other. + + Capture is exempt: `torch.cuda.graph` records on a side stream but + nothing executes, and the captured nodes form a linear chain that + replays in order. Replaying such a graph concurrently with other calls + on the same workspace is still unsupported and cannot be checked from + here. + """ + if torch.cuda.is_current_stream_capturing(): + return + current = torch.cuda.current_stream(self.device) + if self._stream is None: + self._stream = current + elif current != self._stream: + raise RuntimeError( + "this workspace is already bound to " + f"{self._stream}, but all_reduce was called on {current}. " + "One workspace serves one stream: its epoch and arrival " + "counters assume the calls sharing it are totally ordered. " + "Build a second workspace for the second stream." + ) + + def rebind_stream(self) -> None: + """Allow the next call to come from a different stream. + + The workspace rejects a second stream because it cannot tell "used + sequentially from another stream" from "used concurrently", and only + the latter is unsafe. A caller that knows the previous stream's work + has completed -- because it synchronized, or recorded and waited on an + event -- can say so here and move the binding. + + This is an assertion by the caller, not a check: calling it without + actually ordering the two streams reintroduces the corruption it exists + to prevent. + """ + self._stream = None + + def launch_config(self, inp: torch.Tensor) -> Optional[IpcLaunchConfig]: + """Seed launch configuration for ``inp``, or ``None`` if unsupported. + + The seed is a default, not a measurement -- see + :mod:`~flashinfer.comm.pcie_ipc_policy`. :meth:`tuned_launch_config` + is what returns a measured answer once :meth:`tune` has run. + + Depends only on shape, dtype and the workspace's own immutable + attributes, never on rank-local state: every rank must reach the same + answer or the collective deadlocks. + + Raises + ------ + ValueError + If ``inp`` is not on the workspace's device. This is deliberately + not reported as "unsupported": a caller checking :meth:`supports` + reads ``False`` as "use another backend", so answering ``False`` + here would turn a local bug into a silent fallback on one rank -- + and one rank taking a different branch hangs the rest. + """ + # Checked before the workspace state so the diagnosis is the same + # whether or not the workspace is still alive. + if inp.device != self.device: + raise ValueError( + f"input is on {inp.device} but the workspace was built on {self.device}" + ) + if self._handle is None: + return None + if inp.dtype not in _SUPPORTED_DTYPES: + return None + if inp.element_size() != self.elem_size: + return None + if not inp.is_contiguous() or inp.dim() == 0: + return None + numel = inp.numel() + if numel > self.max_numel: + return None + return get_pcie_ipc_launch_config( + self.world_size, numel, self.elem_size, self.max_blocks + ) + + def supports(self, inp: torch.Tensor) -> bool: + """Whether the kernels can run ``inp`` at all. + + A capability question -- dtype, contiguity, workspace capacity, and + enough payload for the reduce-scatter to give every rank a share. It + does not mean the shape has been measured on this machine; call + :meth:`tune` for that. + + Raises the same way :meth:`launch_config` does on a device mismatch -- + that is a caller bug, not an unsupported shape. + + Autotuning never changes this answer: it only picks a faster + configuration for a shape that is already supported. + """ + return self.launch_config(inp) is not None + + def _init_tuning(self) -> None: + """Build the runner and load any persisted configurations. Collective.""" + self._runner = PcieIpcAllReduceRunner(self) + path = self._tune_cache + exists = os.path.isfile(path) + # Whether the file is there has to be a group fact before anyone acts + # on it: half a group running tuned configurations and half running the + # seed is a hang, not a slowdown. + self._joint_check({"error": None, "cache": exists}, "checking the tune cache") + self._tune_cache_exists = exists + if exists: + from ..autotuner import AutoTuner + + AutoTuner.get().load_configs(path) + # Settled against the loaded keys, where the answer is known, rather + # than inferred from a miss later. + self._tuned_configs_loaded = exists and cache_covers_workspace( + self.world_size, self.profile, self.max_blocks, self.max_numel + ) + self._joint_check( + { + "error": None, + "digest": self._cache_digest(), + "covers": self._tuned_configs_loaded, + }, + "loading the tune cache", + ) + + def _warn_if_untuned(self) -> None: + """Say once that this workspace resolved to seed configurations. + + Two causes with different fixes, so two messages: a machine nobody + tuned, or a cache keyed for a different workspace (see + :func:`~flashinfer.comm.pcie_ipc_tuning.cache_covers_workspace`). + + On the cold path only, so the steady state is untouched: a serving loop + reaches this at most once per distinct shape, and the flag makes it once + per workspace. Warning here rather than in ``__init__`` keeps it tied to + actually using the kernels, not to building a workspace the caller may + never route to. + """ + if self._tuned_configs_loaded or self._warned_untuned: + return + self._warned_untuned = True + if self._tune_cache_exists: + warnings.warn( + f"PCIe IPC all-reduce loaded {self._tune_cache} but it holds no " + f"entry for this workspace ({self.world_size} ranks, " + f"max_numel={self.max_numel}, max_blocks={self.max_blocks}, " + f"profile={self.profile}); it was tuned for a different one, so " + "every shape falls back to a seed configuration. Re-tune with " + "this workspace's parameters, or build it with the ones the " + "cache was written for.", + UserWarning, + stacklevel=4, + ) + else: + warnings.warn( + "PCIe IPC all-reduce is running seed launch configurations: " + f"nothing has been tuned for {self.world_size} ranks on this " + f"machine ({self._tune_cache} does not exist). The seed picks a " + "workable kernel, not a fast one. Call workspace.tune([hidden]) " + "once per machine; the result is persisted and later processes " + "pick it up.", + UserWarning, + stacklevel=4, + ) + + def _cache_digest(self) -> str: + """Fingerprint of the tuned entries this rank will actually use. + + ``load_configs`` silently drops entries whose metadata does not match + the machine, so "we all read the same file" is not the same as "we all + hold the same table". + """ + from ..autotuner import AutoTuner + + prefix = f"('{PCIE_IPC_CUSTOM_OP}'" + tuner = AutoTuner.get() + entries = sorted( + (key, repr(value)) + for key, value in tuner._file_configs.items() + if key.startswith(prefix) + ) + return hashlib.sha256(repr(entries).encode()).hexdigest()[:16] + + def tuned_launch_config(self, inp: torch.Tensor) -> Optional[IpcLaunchConfig]: + """Launch configuration for ``inp``, measured if one has been persisted. + + Admission is asked first and is final: a shape the kernels cannot run + returns ``None`` here too, whatever the cache holds. + + Inside an ``autotune(True)`` context this runs the search; outside one + it is a lookup. Same split as the other tunable ops in this library. + """ + seed = self.launch_config(inp) + if seed is None: + return None + from ..autotuner import AutoTuner + + tuner = AutoTuner.get() + key = (inp.numel(), inp.shape[-1], inp.dtype) + # The hot cache is skipped while tuning, so a search that has more + # shapes to cover is not short-circuited by an earlier answer. + if not tuner.is_tuning_mode: + cached = self._tuned.get(key) + if cached is not None: + return cached + # Resolving is collective and reads the verdict back to the host, so it + # cannot happen inside a graph capture. Say that here: the CUDA-level + # failure is "Cannot copy between CPU and CUDA tensors during CUDA + # graph capture", which names neither this workspace nor the fix. + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"the launch configuration for shape {tuple(inp.shape)} dtype " + f"{inp.dtype} has not been resolved yet, and resolving it " + "inside a CUDA graph capture is not possible: the ranks agree " + "on it with a collective whose result is read back on the " + "host. Call workspace.prepare() with every shape you intend to " + "capture -- after tune(), which clears this cache -- or pass " + "config= explicitly at the call site." + ) + config = self._resolve_tuned(inp, seed, tuner) + self._tuned[key] = config + return config + + def _resolve_tuned( + self, inp: torch.Tensor, seed: IpcLaunchConfig, tuner + ) -> IpcLaunchConfig: + """Cold path: search or look up, then make the group agree.""" + hidden = inp.shape[-1] + batch = inp.numel() // hidden + tuning_config = pcie_ipc_tuning_config(self._tune_batches) + can_profile = tuner.is_tuning_mode and self._runner.can_profile(inp.device) + if can_profile: + _, tactic = tuner.choose_one( + PCIE_IPC_CUSTOM_OP, [self._runner], tuning_config, [inp] + ) + else: + # An enclosing autotune context may belong to another operator and + # may have replaced the global file cache. Without a matching + # distributed tune group, profiling this collective is unsafe. + # Restore the workspace's explicit tune cache and perform a lookup + # with its own bucket policy instead of retaining the seed tactic. + if tuner.is_tuning_mode: + # The runner would have said this had it been asked for + # candidates; short-circuiting before ``choose_one`` is what + # would otherwise swallow it. It is the actionable half of the + # diagnosis -- the caller is tuning, so "go tune" is not. + warn_no_tune_group(stacklevel=4) + if self._tune_cache_exists: + tuner.load_configs(self._tune_cache) + _, _, tactic, _ = tuner.search_cache( + PCIE_IPC_CUSTOM_OP, + [self._runner], + ((batch, hidden),), + tuning_config, + inputs=[inp], + ) + config = resolve_tuned_config(seed, tactic, self.world_size, self.max_blocks) + + # Unconditional, even when the cache missed and `config is seed`. The + # ranks would otherwise have to agree on whether to run this collective + # before running it, and disagreeing about that is the hang it exists + # to prevent. It costs one small reduction per distinct shape. + packed = pack_config(config) + bounds = torch.tensor([packed, -packed], dtype=torch.int64, device=self.device) + dist.all_reduce(bounds, op=dist.ReduceOp.MAX, group=self.group) + if int(bounds[0]) != -int(bounds[1]): + # Fall back rather than raise: the seed is a pure function, so it + # is agreed by construction and the group stays alive. + warnings.warn( + "ranks resolved different tuned configurations for shape " + f"{tuple(inp.shape)}; falling back to the seed configuration. " + "The tune cache is inconsistent across ranks -- delete " + f"{self._tune_cache} and re-tune.", + RuntimeWarning, + stacklevel=3, + ) + return seed + if not tuner.is_tuning_mode: + self._warn_if_untuned() + return config + + @flashinfer_api(trace=pcie_ipc_all_reduce_trace) + def all_reduce( + self, + inp: torch.Tensor, + *, + out: Optional[torch.Tensor] = None, + config: Optional[IpcLaunchConfig] = None, + enable_pdl: bool = False, + ) -> torch.Tensor: + """Out-of-place all-reduce. + + Parameters + ---------- + inp : torch.Tensor + Contiguous CUDA tensor whose byte size is a multiple of 16. + out : torch.Tensor, optional + Destination. Allocated when omitted. + config : IpcLaunchConfig, optional + Launch geometry and kernel selection. Resolved from the tune cache + or the seed when omitted; pass one explicitly only to benchmark or + to reach a kernel neither would choose. Ranks that disagree on it + hang -- see the collective contract in the class docstring. + enable_pdl : bool + Programmatic dependent launch. **Currently rejected.** The TP8 + block kernel triggers launch completion before it writes its + island ack and barrier flag, so a dependent kernel could start + while this call's protocol state is still being written. + + Returns + ------- + torch.Tensor + The reduced tensor. + + Raises + ------ + ValueError + If the kernels cannot run this shape. Check :meth:`supports` first + and fall back to another backend. + """ + if config is None: + config = self.tuned_launch_config(inp) + if config is None: + raise ValueError( + f"unsupported shape {tuple(inp.shape)} dtype {inp.dtype} " + f"at {self.world_size} ranks; check supports() first" + ) + self._check_stream() + # Raise rather than fall back: a device mismatch is a caller bug, and + # silently opting this rank out would hang every other rank. + if inp.device != self.device: + raise ValueError( + f"input is on {inp.device} but the workspace was built on {self.device}" + ) + if out is None: + out = torch.empty_like(inp) + elif out.device != self.device: + raise ValueError( + f"output is on {out.device} but the workspace was built on " + f"{self.device}" + ) + self._launch(inp, out, config, enable_pdl) + return out + + def _launch( + self, + inp: torch.Tensor, + out: torch.Tensor, + config: IpcLaunchConfig, + enable_pdl: bool = False, + ) -> None: + """Issue one collective with an explicit configuration. + + The launch without the admission, device and stream checks around it. + Callers that have already done those -- the tuner, which sweeps many + configurations over one validated pair of buffers -- use this so the + checks do not run once per candidate. + """ + get_pcie_ipc_comm_module().all_reduce( + self.handle, + inp, + out, + config.blocks, + config.threads, + int(config.variant), + enable_pdl, + ) + + def prepare( + self, + shapes: Sequence[Tuple[int, int]], + *, + dtype: torch.dtype = torch.bfloat16, + ) -> Dict[Tuple[int, int], Optional[IpcLaunchConfig]]: + """Resolve the launch configuration for each shape now. **Collective.** + + Resolution is lazy by default: the first call at a given shape looks the + configuration up and then makes the group agree on it, which costs one + small reduction whose verdict is read back on the host. That is fine in + eager mode and impossible inside a CUDA graph capture, so a shape first + used inside a capture fails to capture. + + This moves that work to a point the caller chooses. Nothing else + changes -- the same lookup, the same agreement, the same number of + collectives -- and afterwards every listed shape is served from the + in-process cache, so a capture of it touches no collective at all. + + Call it **after** :meth:`tune`, which clears that cache, and list every + shape that will be captured: shapes left out are still resolved lazily + and still cannot be captured. Serving frameworks pad the batch to the + bucket they capture, so list the padded sizes, not the real ones. + + Parameters + ---------- + shapes : sequence of (batch, hidden) + Shapes to resolve. Every rank must pass the same list in the same + order -- resolution is collective, so a rank with a different list + deadlocks the group rather than disagreeing. + dtype : torch.dtype + Which of the two supported dtypes to resolve for. The cache is + keyed by dtype, so resolve each one that will be used. + + Returns + ------- + dict + ``{(batch, hidden): config}``, with ``None`` for shapes the kernels + do not support -- those fall back to another backend at call time + and never reach a capture. + """ + shapes = [(int(batch), int(hidden)) for batch, hidden in shapes] + # Same reasoning as tune(): the loop below issues one collective per + # shape, so a rank with a different list hangs rather than disagrees. + self._joint_check( + {"error": None, "shapes": shapes, "dtype": str(dtype)}, + "preparing launch configurations", + ) + resolved: Dict[Tuple[int, int], Optional[IpcLaunchConfig]] = {} + for batch, hidden in shapes: + probe = torch.empty((batch, hidden), dtype=dtype, device=self.device) + resolved[(batch, hidden)] = self.tuned_launch_config(probe) + return resolved + + def tune( + self, + hiddens: Sequence[int], + *, + dtype: torch.dtype = torch.bfloat16, + cache: Optional[str] = None, + tune_group=None, + warmup: int = TUNE_WARMUP, + repeat: int = TUNE_REPEAT, + ) -> Dict[Tuple[int, int], IpcLaunchConfig]: + """Measure the launch configuration for every tuned shape. Collective. + + A convenience wrapper around the library's usual tuning idiom:: + + with flashinfer.autotune(True, cache=path): + for batch in batches: + ws.all_reduce(sample(batch)) + + which also works, and does the same thing. This adds what a collective + needs on top of it: a gloo subgroup for the timing reduction so every + rank picks the same kernel, longer timing runs than the library default + (the library defaults resolve too little at this scale), a check that + every rank agrees on the arguments, and a single writer for the result + file. + + Every rank must call this with identical arguments, and clocks should be + pinned first (``nvidia-smi -lgc``): boost drift is larger than the + differences being ranked. + + Parameters + ---------- + hiddens : Sequence[int] + Hidden sizes to tune -- the ones this job will actually run. There + is no default: admission does not constrain the hidden size, so + there is no finite set to enumerate, and guessing would quietly tune + a shape nobody uses. + + The **batch** dimension is not here. It comes from ``tune_batches`` + on the constructor, because the buckets have to be the same on the + tuning side and the lookup side, which makes them a property of the + workspace rather than of one call. + dtype : torch.dtype + Which of the two supported dtypes to measure. Both are 2 bytes so + the traffic is identical, but they take different conversion paths. + cache : str, optional + Where to persist results. Defaults to the workspace's + ``tune_cache``, which is also where the next process reads them. + tune_group : ProcessGroup, optional + Group used to reduce per-candidate timings so every rank picks the + same winner. Built here as a gloo subgroup when the workspace spans + the default process group; must be supplied otherwise, because + ``new_group`` is collective over the *default* group and building + one here would hang a job whose workspace is a strict subgroup. + warmup, repeat : int + Untimed and timed iterations per candidate. The library defaults + time too short a span to resolve candidates for a collective this + fast, so these default higher. + + Returns + ------- + dict + ``{(hidden, batch): config}`` for every shape that was measured, so + the caller can see what tuning actually covered and what it chose. + + Raises + ------ + ValueError + If none of ``hiddens`` yields a shape the kernels admit -- otherwise + the call is a silent no-op. + """ + from ..autotuner import ( + AutoTuner, + autotune, + get_autotune_process_group, + set_autotune_process_group, + ) + + hiddens = tuple(int(h) for h in hiddens) + path = cache or self._tune_cache + # Everything the collective profiling contract requires to match, in + # one gather. A blocklist set on one rank alone silently shortens that + # rank's candidate list, and the timing reduction then deadlocks on the + # first divergence. + self._joint_check( + { + "error": None, + "hiddens": hiddens, + "dtype": str(dtype), + "cache": path, + "warmup": warmup, + "repeat": repeat, + "tune_batches": self._tune_batches, + "blocklist": os.environ.get("FLASHINFER_TACTICS_BLOCKLIST", ""), + "digest": self._cache_digest(), + }, + "starting a tuning run", + ) + + if tune_group is None: + tune_group = self._make_tune_group() + elif dist.get_world_size(tune_group) != self.world_size: + raise ValueError( + f"tune_group spans {dist.get_world_size(tune_group)} ranks but " + f"the workspace spans {self.world_size}" + ) + + tuner = AutoTuner.get() + previous_group = get_autotune_process_group() + previous_counts = (tuner.warmup, tuner.repeat) + set_autotune_process_group(tune_group) + # The library defaults time too short a span to resolve candidates at + # this operator's scale. + tuner.warmup, tuner.repeat = warmup, repeat + covered: List[Tuple[int, int]] = [] + skipped: List[int] = [] + try: + for hidden in hiddens: + batches = [ + b + for b in tuned_batches_for( + hidden, self._tune_batches, self.max_numel + ) + if self.launch_config( + torch.empty((b, hidden), dtype=dtype, device=self.device) + ) + is not None + ] + if not batches: + # Recorded rather than skipped silently: the call would + # otherwise return cleanly having measured nothing. + skipped.append(hidden) + continue + torch.cuda.synchronize(self.device) + self.rebind_stream() + with autotune(True, tuning_buckets=tuple(batches), round_up=False): + for batch in batches: + inp = torch.randint( + 0, + 16, + (batch, hidden), + dtype=torch.int32, + device=self.device, + ).to(dtype) + tuner.choose_one( + PCIE_IPC_CUSTOM_OP, + [self._runner], + pcie_ipc_tuning_config(self._tune_batches), + [inp], + ) + covered.append((hidden, batch)) + finally: + tuner.warmup, tuner.repeat = previous_counts + # Restore rather than clear: a caller may be tuning something else + # around this. + set_autotune_process_group(previous_group) + + if skipped: + message = ( + f"tune() measured nothing for hidden {skipped} at " + f"{self.world_size} ranks: the kernels do not support those " + "shapes, and tuning does not widen what is supported." + ) + if not covered: + raise ValueError(message) + warnings.warn(message, RuntimeWarning, stacklevel=2) + + # Winners live in the in-memory cache now, so drop anything this + # workspace resolved from the seed. + self._tuned.clear() + self._tuned_configs_loaded = True + dist.barrier(group=self.group) + if self.rank == 0: + os.makedirs(os.path.dirname(path) or ".", exist_ok=True) + tuner.save_configs(path) + # Nobody leaves before the file is on disk: a peer that rebuilt its + # workspace first would load a half-written table. + dist.barrier(group=self.group) + return { + (hidden, batch): self.tuned_launch_config( + torch.empty((batch, hidden), dtype=dtype, device=self.device) + ) + for hidden, batch in covered + } + + def _make_tune_group(self): + """A gloo subgroup for reducing candidate timings. + + gloo because the reduction carries one float64 and an NCCL collective + immediately after a spin-waiting IPC kernel is exactly the interference + a timing loop does not want. + """ + if self._tune_group is not None: + return self._tune_group + ranks = dist.get_process_group_ranks(self.group) + if len(ranks) != dist.get_world_size(): + raise ValueError( + "tune() cannot build its own reduction group for a workspace " + "that spans a strict subgroup: new_group() is collective over " + "the default process group, so every process would have to " + "call it. Pass tune_group= built by all ranks instead." + ) + self._tune_group = dist.new_group(ranks=ranks, backend="gloo") + return self._tune_group + + def destroy(self) -> None: + """Release the handle and the shared slab. + + Collective: every rank must call this, and the peer unmapping is + separated from the free by a barrier inside ``free_shared_buffer``. + """ + if self._handle is not None: + # all_reduce() launches asynchronously, so a collective may still be + # running or spinning on this slab. free_shared_buffer() unmaps the + # peers, and unmapping memory a live kernel is still touching is a + # use-after-free -- wait for the device before tearing anything + # down. This is the conservative choice; a stream-scoped wait would + # need the workspace to track every stream it has been used on. + torch.cuda.synchronize(self.device) + get_pcie_ipc_comm_module().dispose(self._handle) + self._handle = None + if self._ipc_ptrs is not None: + free_shared_buffer(self._ipc_ptrs, group=self.group) + self._ipc_ptrs = None + if self._tune_group is not None: + dist.destroy_process_group(self._tune_group) + self._tune_group = None + self._tuned.clear() + + def __enter__(self) -> "PcieIpcAllReduceWorkspace": + return self + + def __exit__(self, *exc_info) -> None: + self.destroy() diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_policy.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_policy.py new file mode 100644 index 0000000..da181fb --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_policy.py @@ -0,0 +1,188 @@ +""" +Copyright (c) 2026 by FlashInfer team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. + +Launch configurations for the PCIe IPC all-reduce. + +Two layers, with very different standing. + +**Admission** (:func:`_admits`) is a capability question: which shapes the +kernels can run at all. It is not a performance judgement, and tuning cannot +change it. + +**The seed** (:func:`_seed`) is a *default*, not a measurement. It picks the +one side of the one crossover that ports between machines -- push straight to +every peer while the payload is small, reduce-scatter/all-gather once it is +not -- and nothing finer. + +Thresholds fitted per batch on one machine do not survive the trip to another, +so only the shape of the answer lives here; the numbers come from +:meth:`~flashinfer.comm.PcieIpcAllReduceWorkspace.tune`, which measures them +where they will run. Running untuned is warned about once per workspace. +""" + +from dataclasses import dataclass, replace +from enum import IntEnum +from functools import lru_cache +from typing import Optional + + +# Block counts above this are never useful on either fabric and the workspace +# is sized for it. +MAX_BLOCKS = 128 + + +class IpcVariant(IntEnum): + """Which kernel to launch; mirrors ``fi::Variant`` in the header. + + Values cross the FFI boundary as integers, so they are append-only. + ``FLAT_STAGED`` is accepted at world size 8 only -- at 4 it would name the + same kernel as ``STAGED``, and at 2 there is no staged-vs-flat distinction. + """ + + UNSTAGED = 0 + STAGED = 1 + STAGED_RING = 2 + FLAT_STAGED = 3 + + +@dataclass(frozen=True) +class IpcLaunchConfig: + blocks: int + threads: int + variant: IpcVariant + + +# Payload above which reduce-scatter/all-gather beats pushing to every peer. +# Keyed on bytes, not tokens: the crossover trades bytes moved against barrier +# latency, and only that ratio ports between fabrics. +_SEED_STAGE_BYTES = 32 * 1024 + +# The neighbour-ordered kernel has one outbound stream per rank whatever the +# grid, so extra blocks pay only once there are bytes enough to keep the link +# busy. Capped low: with no switch-local peer, concurrent transfers collapse +# rather than add. +_SEED_RING_BYTES_PER_BLOCK = 256 * 1024 +_SEED_RING_MAX_BLOCKS = 4 + + +def _admits(world_size: int, numel: int, elem_size: int) -> bool: + """Whether the kernels can run this shape at all. + + Independent of which kernel is chosen: tuning launches every variant on + whatever shape is admitted, so a precondition that held only for the + variant the seed happens to pick would still deadlock the group under + :meth:`~flashinfer.comm.PcieIpcAllReduceWorkspace.tune`. + """ + if world_size not in (2, 4, 8): + return False + pack_elems = 16 // elem_size + # Matches the launcher's own check; the kernels address whole 16-byte packs. + if numel % pack_elems != 0: + return False + # Reduce-scatter gives each rank num_packs // world_size packs. Below one + # pack per rank that split degenerates onto a single owner: correct, but it + # leaves the other ranks idle, and a payload that small is better served by + # another backend than by an IPC collective. + return numel >= pack_elems * world_size + + +def _seed( + world_size: int, numel: int, elem_size: int, max_blocks: int +) -> IpcLaunchConfig: + """Default configuration for a shape nothing has measured yet.""" + payload = numel * elem_size + + if world_size == 2: + # Staging moves the same bytes it would have pushed, so there is no + # crossover here and no second branch to justify. + return IpcLaunchConfig(min(16, max_blocks), 128, IpcVariant.UNSTAGED) + + if payload >= _SEED_STAGE_BYTES: + # Neighbour-ordered rather than all-to-all: with no switch-local peer, + # simultaneous writes to every peer collapse, and the penalty grows with + # the payload. Picking wrong on this arm is unbounded rather than merely + # slow, which is why the threshold sits low. + blocks = max( + 1, + min( + _SEED_RING_MAX_BLOCKS, + payload // _SEED_RING_BYTES_PER_BLOCK, + max_blocks, + ), + ) + return IpcLaunchConfig(blocks, 256, IpcVariant.STAGED_RING) + + if world_size == 4: + # Staging always cuts egress at four ranks and the all-to-all form pays + # only two barriers, so the one-shot push is never the answer. One + # block, because its grid multiplies the concurrency that collapses. + return IpcLaunchConfig(1, 256, IpcVariant.STAGED) + + # Eight ranks below the crossover: the island-partitioned push has no + # barriers, which is what the staged path's six island barriers must beat. + return IpcLaunchConfig(min(16, max_blocks), 256, IpcVariant.UNSTAGED) + + +def _is_launchable(world_size: int, config: IpcLaunchConfig, max_blocks: int) -> bool: + """Reject configurations the kernels cannot accept. + + A violation here degrades to "unsupported shape" and a caller fallback, + which is far better than reaching the kernel and failing a hard check. + """ + if not 0 < config.blocks <= max_blocks: + return False + if not world_size <= config.threads <= 1024: + return False + # One configuration must name exactly one kernel, so the pairs the header + # does not dispatch are rejected rather than aliased onto a neighbour. + if world_size == 2 and config.variant not in ( + IpcVariant.UNSTAGED, + IpcVariant.STAGED, + ): + return False + if config.variant == IpcVariant.FLAT_STAGED and world_size != 8: + return False + # The block-partitioned TP8 kernel derives its chunk from blockIdx.x & 3. + # Every other kernel uses flat grid-stride loops. + if ( + world_size == 8 + and config.variant == IpcVariant.STAGED + and config.blocks % 4 != 0 + ): + return False + return True + + +@lru_cache(maxsize=None) +def get_pcie_ipc_launch_config( + world_size: int, + numel: int, + elem_size: int, + max_blocks: int = MAX_BLOCKS, +) -> Optional[IpcLaunchConfig]: + """Launch configuration for one shape, or ``None`` when unsupported. + + ``None`` means the kernels cannot run the shape, so the caller must use + another backend. It never means "untuned": an untuned shape gets the seed. + + Depends only on its arguments, so every rank in a group reaches the same + answer -- a prerequisite, since a rank that opts out while its peers opt in + deadlocks the collective. + """ + if not _admits(world_size, numel, elem_size): + return None + config = _seed(world_size, numel, elem_size, max_blocks) + config = replace(config, blocks=min(config.blocks, max_blocks)) + return config if _is_launchable(world_size, config, max_blocks) else None diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_topology.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_topology.py new file mode 100644 index 0000000..1bb869a --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_topology.py @@ -0,0 +1,210 @@ +""" +Copyright (c) 2026 by FlashInfer team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +import socket +from dataclasses import dataclass, field +from typing import Dict, List, Optional + +import torch +import torch.distributed as dist +from torch.distributed import ProcessGroup + +# Which fabric the group is on. The distinction is the interconnect, not the +# GPU: the same card behaves differently depending on whether its NUMA island +# contains a PCIe switch. It selects no kernel -- it keys the tune cache, so two +# topologies on one machine do not read each other's measurements. +PROFILE_ROOTCPLX = "rootcplx-noswitch" +PROFILE_SWITCHPAIR = "pcieswitch-pairs" +PCIE_IPC_PROFILES = (PROFILE_ROOTCPLX, PROFILE_SWITCHPAIR) + +_PROFILE_ALIASES = { + "rootcplx": PROFILE_ROOTCPLX, + "rootcplx-noswitch": PROFILE_ROOTCPLX, + "pcieswitch": PROFILE_SWITCHPAIR, + "pcieswitch-pairs": PROFILE_SWITCHPAIR, +} + + +@dataclass +class PcieIpcRankTopology: + """Per-rank probe result, exchanged across the group. + + ``peer_switch_local`` is keyed by the *peer GPU's UUID* so the decision + layer can join results across ranks regardless of each process's + ``CUDA_VISIBLE_DEVICES`` ordering, and so the probe only ever describes + GPUs this rank can actually see. + """ + + rank: int + hostname: str = "" + device_index: int = -1 + device_uuid: str = "" + peer_switch_local: Dict[str, bool] = field(default_factory=dict) + pair_errors: Dict[str, str] = field(default_factory=dict) + probe_error: Optional[str] = None + + +@dataclass(frozen=True) +class PcieIpcProfileDecision: + profile: str + reason: str + + +def probe_pcie_ipc_rank_topology( + rank: int, device: Optional[torch.device] = None +) -> PcieIpcRankTopology: + """Probe whether this rank's GPU shares a PCIe switch with any peer. + + Never raises: any failure is recorded in ``probe_error`` and the decision + layer treats an unknown topology conservatively. + + Only the GPU this rank owns is probed against the other visible GPUs, so a + job pinned to a subset of the machine describes that subset rather than the + whole host. That matters on a mixed box where one island sits behind a + switch and another does not. + """ + topo = PcieIpcRankTopology(rank=rank) + try: + topo.hostname = socket.gethostname() + parsed = ( + torch.device("cuda", torch.cuda.current_device()) + if device is None + else torch.device(device) + ) + if parsed.type != "cuda": + raise ValueError(f"probe requires a CUDA device, got {parsed!r}") + device_index = ( + parsed.index if parsed.index is not None else torch.cuda.current_device() + ) + topo.device_index = device_index + + import pynvml + + pynvml.nvmlInit() + try: + + def _uuid(idx: int) -> str: + props = torch.cuda.get_device_properties(idx) + uuid = getattr(props, "uuid", None) + if uuid is None: + raise RuntimeError( + "torch.cuda.get_device_properties(...).uuid unavailable; " + "cannot establish physical GPU identity" + ) + return f"GPU-{uuid}" + + def _handle(idx: int): + return pynvml.nvmlDeviceGetHandleByUUID(_uuid(idx).encode()) + + topo.device_uuid = _uuid(device_index) + my_handle = _handle(device_index) + # NVML_TOPOLOGY_HOSTBRIDGE is the first level that leaves the switch + # fabric, so anything strictly below it means the pair talks through + # a PCIe switch without reaching the host bridge. + hostbridge = pynvml.NVML_TOPOLOGY_HOSTBRIDGE + for peer in range(torch.cuda.device_count()): + if peer == device_index: + continue + peer_uuid = _uuid(peer) + try: + level = pynvml.nvmlDeviceGetTopologyCommonAncestor( + my_handle, _handle(peer) + ) + topo.peer_switch_local[peer_uuid] = level < hostbridge + except pynvml.NVMLError as pair_err: + topo.pair_errors[peer_uuid] = str(pair_err) + finally: + pynvml.nvmlShutdown() + except Exception as e: # noqa: BLE001 - any probe failure => conservative fallback + topo.probe_error = f"{type(e).__name__}: {e}" + return topo + + +def decide_pcie_ipc_profile( + requested: Optional[str], topologies: List[PcieIpcRankTopology] +) -> PcieIpcProfileDecision: + """Pick the fabric label from the gathered probes. Pure function. + + An explicit ``requested`` profile always wins. Otherwise the group is + switch-paired only if some rank positively observed a switch-local peer; + anything unknown or unprobeable falls back to ``rootcplx-noswitch``. + Guessing wrong costs a tune cache keyed on the other fabric, so the label + that claims less is the safe default. + """ + # The intra-node constraint is checked first: CUDA IPC cannot cross hosts, + # so an explicit profile must not be able to wave it through. + hosts = {t.hostname for t in topologies if t.hostname} + if len(hosts) > 1: + raise ValueError( + f"pcie ipc all-reduce is intra-node only, but the group spans {sorted(hosts)}" + ) + + if requested is not None: + key = requested.strip().lower() + if key not in _PROFILE_ALIASES: + raise ValueError( + f"unknown pcie ipc profile {requested!r}; " + f"expected one of {sorted(_PROFILE_ALIASES)}" + ) + return PcieIpcProfileDecision(_PROFILE_ALIASES[key], "requested explicitly") + + failed = [t.rank for t in topologies if t.probe_error] + if failed: + return PcieIpcProfileDecision( + PROFILE_ROOTCPLX, f"probe failed on ranks {failed}; assuming no switch pair" + ) + + # Only pairs where BOTH endpoints belong to this group count. The probe + # walks every GPU the process can see, which for a subgroup is a superset: + # a switch-local pair outside the group says nothing about how the group's + # own ranks talk to each other. + members = {t.device_uuid for t in topologies if t.device_uuid} + for t in topologies: + for peer_uuid, switch_local in t.peer_switch_local.items(): + if switch_local and peer_uuid in members: + return PcieIpcProfileDecision( + PROFILE_SWITCHPAIR, + f"rank {t.rank} shares a PCIe switch with group member {peer_uuid}", + ) + + partial = [t.rank for t in topologies if any(u in members for u in t.pair_errors)] + if partial: + return PcieIpcProfileDecision( + PROFILE_ROOTCPLX, + f"some in-group pairs unprobeable on ranks {partial}; " + "assuming no switch pair", + ) + return PcieIpcProfileDecision(PROFILE_ROOTCPLX, "no switch-local pair observed") + + +def resolve_pcie_ipc_profile( + group: ProcessGroup, + requested: Optional[str] = None, + device: Optional[torch.device] = None, +) -> PcieIpcProfileDecision: + """Probe every rank and agree on one profile. + + Collective. Runs before any workspace allocation or JIT build so an + unsupported topology costs nothing, and gathers the per-rank probes so + every rank reaches the same decision from the same evidence. + """ + rank = dist.get_rank(group=group) + local = probe_pcie_ipc_rank_topology(rank, device=device) + gathered: List[Optional[PcieIpcRankTopology]] = [None] * dist.get_world_size( + group=group + ) + dist.all_gather_object(gathered, local, group=group) + return decide_pcie_ipc_profile(requested, [t for t in gathered if t is not None]) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_tuning.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_tuning.py new file mode 100644 index 0000000..ba4bf66 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/comm/pcie_ipc_tuning.py @@ -0,0 +1,530 @@ +""" +Copyright (c) 2026 by FlashInfer team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. + +Autotuning for the PCIe IPC all-reduce. + +The seed in :mod:`~flashinfer.comm.pcie_ipc_policy` is a default, not a +measurement: one crossover, and no constants fitted to any machine. This module +measures the same choice, over the launch configurations the dispatch can +actually reach. + +Two properties of the surrounding code shape everything here: + +**The autotuner never looks at a kernel's output**, and this kernel family's +characteristic failure is wrong *and* fast. So every candidate is verified +against a reference before it is timed, and the verdict is reduced across the +group -- see :meth:`PcieIpcAllReduceRunner.get_valid_tactics`. + +**Every wait in the kernels is an unbounded spin.** Ranks that disagree on the +launch configuration, or that issue different numbers of calls, hang rather +than raise. So the candidate list is a pure function of group-identical +arguments, the verification verdict is reduced before it is used, and the +resolved configuration is checked for group agreement before it is cached. + +The policy module keeps three jobs here: admission decides which shapes are +supported at all, and the seed is both tactic ``-1`` and the fallback whenever a +tuned answer cannot be used. +""" + +import os +import warnings +from functools import lru_cache +from typing import Dict, List, Optional, Sequence, Tuple + +import torch +import torch.distributed as dist + +from ..autotuner import ( + DynamicTensorSpec, + TunableRunner, + TuningConfig, + make_bucket_mapper, +) +from .pcie_ipc_policy import ( + MAX_BLOCKS, + IpcLaunchConfig, + IpcVariant, + _is_launchable, +) + +# Baked into every persisted cache key, so renaming it silently invalidates +# every cache file rather than mis-resolving one. +PCIE_IPC_CUSTOM_OP = "flashinfer::pcie_ipc_all_reduce" + +# Bump when a variant's meaning, the scratch-region assignment, or the +# candidate encoding changes. The autotuner's own metadata records library and +# driver versions but nothing about this op, and a dev checkout does not move +# the FlashInfer version. +PCIE_IPC_TUNE_VERSION = 1 + +# Not all powers of two: the extra entries are block counts the search selected +# on real hardware, and it cannot converge on a configuration its own grid +# cannot name. +TUNE_BLOCKS: Tuple[int, ...] = (1, 2, 4, 8, 12, 16, 32, 64, 96, 128) +TUNE_THREADS: Tuple[int, ...] = (64, 128, 256, 512, 1024) + +# Batch buckets. Floor semantics, so a bucket is always a batch the tuner +# actually measured. Matches the benchmark's default sweep. +TUNE_BATCHES: Tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128) + +# Higher than the library defaults, which time too short a span to resolve +# candidates for a collective of this scale. +TUNE_WARMUP = 10 +TUNE_REPEAT = 50 + +# Reference tactic. The autotuner reserves -1 for "the fallback that implements +# any shape"; here that is the policy module's seed configuration. +TABLE_TACTIC = -1 + +# Inputs are drawn from [0, INIT_MAX_VALUE) so the group sum stays integral and +# exactly representable, which is what lets verification use a zero tolerance +# despite the kernels summing in a different order than NCCL. +INIT_MAX_VALUE = 16 + + +def candidate_tactics( + world_size: int, + max_blocks: int = MAX_BLOCKS, + blocks: Tuple[int, ...] = TUNE_BLOCKS, + threads: Tuple[int, ...] = TUNE_THREADS, +) -> Tuple[Tuple[int, int, int], ...]: + """Every launch configuration the dispatch can reach, as tactics. + + A pure function of group-identical arguments, so every rank derives the + same list in the same order -- which the autotuner's collective profiling + requires and cannot check. + """ + return _candidate_tactics_cached(world_size, max_blocks, blocks, threads) + + +@lru_cache(maxsize=None) +def _candidate_tactics_cached(world_size, max_blocks, blocks, threads): + out = [] + for variant in IpcVariant: + for b in blocks: + for t in threads: + if _is_launchable( + world_size, IpcLaunchConfig(b, t, variant), max_blocks + ): + out.append((int(variant), b, t)) + return tuple(out) + + +def config_to_tactic(config: IpcLaunchConfig) -> Tuple[int, int, int]: + """Encode a configuration as a tactic. + + Plain ints, because a tactic has to survive a JSON round-trip: the + autotuner writes ``[0, 32, 128]`` and reads back ``(0, 32, 128)``. + Self-describing rather than an index into :func:`candidate_tactics`, so + editing the grid cannot repoint a persisted entry at a different kernel. + """ + return (int(config.variant), int(config.blocks), int(config.threads)) + + +def tactic_to_config(tactic: Sequence[int]) -> IpcLaunchConfig: + """Decode a tactic. Raises ``ValueError`` on anything malformed.""" + if len(tactic) != 3: + raise ValueError(f"expected a 3-element tactic, got {tactic!r}") + variant, blocks, threads = (int(v) for v in tactic) + try: + return IpcLaunchConfig(blocks, threads, IpcVariant(variant)) + except ValueError as exc: + raise ValueError(f"tactic {tactic!r} names no variant: {exc}") from exc + + +def cache_covers_workspace( + world_size: int, profile: str, max_blocks: int, max_numel: int +) -> bool: + """Whether the loaded cache holds any entry written for this workspace. + + ``max_numel`` is part of the key, so a workspace sized differently from the + tuned one misses every entry at once rather than a few -- a configuration + mistake rather than an untuned shape, and no single lookup can tell those + apart, since the seed is a valid answer either way. + + dtype is not compared: one workspace serves both 2-byte dtypes and each gets + its own entries, so a match would be required for a cache that covers the + workspace perfectly well in the dtype the caller is not using. + + Scanned rather than parsed for the same reason + :meth:`PcieIpcAllReduceWorkspace._cache_digest` scans -- the key format + belongs to the autotuner. + """ + from ..autotuner import AutoTuner + + prefix = f"('{PCIE_IPC_CUSTOM_OP}'" + # cache_key_extras up to the dtype, with the closing paren traded for the + # separator that must follow it. + head = ( + PCIE_IPC_TUNE_VERSION, + int(world_size), + str(profile), + int(max_blocks), + int(max_numel), + ) + needle = repr(head)[:-1] + ", " + return any( + key.startswith(prefix) and needle in key + for key in AutoTuner.get()._file_configs + ) + + +def resolve_tuned_config( + table_config: IpcLaunchConfig, + tactic, + world_size: int, + max_blocks: int, +) -> IpcLaunchConfig: + """Turn a tactic into a configuration, falling back to the seed. + + The autotuner does not check that a cached tactic can implement the shape + it is being reused for, so a cache written against a larger ``max_blocks`` + would otherwise reach the launcher's hard checks and raise on every rank in + the middle of a collective. + """ + if tactic is None or tactic == TABLE_TACTIC: + return table_config + try: + config = tactic_to_config(tactic) + except (TypeError, ValueError): + return table_config + if not _is_launchable(world_size, config, max_blocks): + return table_config + return config + + +def small_int_initializer( + shapes: Tuple[int, ...], dtype: torch.dtype, device: torch.device +) -> torch.Tensor: + """Synthesize profiling inputs that can be compared at zero tolerance. + + The autotuner's default fills tensors with ``rand() * 10 - 5``, which no + reference can be compared against exactly. Small integers keep the group + sum exact in both supported dtypes, so verification uses ``torch.equal`` + and cannot mistake a reduction-order difference for a protocol bug. Zero is + in the range on purpose: the sentinel kernels rewrite real zeros in the + payload, and that path should be exercised. + """ + return torch.randint( + 0, INIT_MAX_VALUE, shapes, device=device, dtype=torch.int32 + ).to(dtype) + + +@lru_cache(maxsize=None) +def pcie_ipc_tuning_config(batches: Tuple[int, ...] = TUNE_BATCHES) -> TuningConfig: + """Tuning configuration for one bucket set. + + Cached so that the serving-side cache lookup and the tuning-side search + share one object: the bucket mapper has to be identity-stable or the + autotuner's profile lookup degenerates. + + Only the batch dimension is dynamic. Hidden stays static so it lands + verbatim in the cache key -- the configuration follows the payload in bytes, + and a bucketed hidden would silently reuse another payload's answer. + """ + return TuningConfig( + dynamic_tensor_specs=( + DynamicTensorSpec( + input_idx=(0,), + dim_idx=(0,), + gen_tuning_buckets=batches, + map_to_tuning_buckets=make_bucket_mapper(batches, round_map=False), + ), + ), + tensor_initializers=((0, small_int_initializer),), + # Capture is required, not preferred. Without it the profiler issues + # each iteration separately, the host cannot keep up with a collective + # this short, and the span it times is dominated by launch gaps -- it + # would rank host overhead rather than kernels. + use_cold_l2_cache=False, + use_cuda_graph=True, + ) + + +def default_cache_path(world_size: int) -> str: + """Where tuned configurations are persisted. + + World size is in the filename as well as the cache key so that a TP4 and a + TP8 job on the same host never contend for one file. + """ + import pathlib + + override = os.getenv("FLASHINFER_AUTOTUNE_DIR") + if override: + base = pathlib.Path(override) + else: + from ..jit.env import FLASHINFER_WORKSPACE_DIR + + base = FLASHINFER_WORKSPACE_DIR / "autotune" + return str(base / f"pcie_ipc_all_reduce_ws{world_size}.json") + + +def cache_key_extras( + world_size: int, + profile: str, + max_blocks: int, + max_numel: int, + dtype: torch.dtype, +) -> Tuple: + """Everything the autotuner's own cache key leaves out. + + That key is only the bucketed input shapes, so without these a TP4 and a + TP8 entry at the same shape would collide, a configuration tuned on one + fabric would be reused on the other, and a cache written for one workspace + size would be applied to another. ``max_numel`` matters because the epoch + double buffer places its halves ``world_size * max_numel`` apart, so the + best block count genuinely depends on it. + + Every field is a workspace immutable or the input dtype, which is what the + autotuner requires: the tuple must come out the same for the caller's real + tensors and for the ones it synthesizes. + """ + return ( + PCIE_IPC_TUNE_VERSION, + int(world_size), + str(profile), + int(max_blocks), + int(max_numel), + str(dtype), + ) + + +def pack_config(config: IpcLaunchConfig) -> int: + """Pack a configuration into one integer for a cross-rank comparison.""" + return ( + (int(config.variant) << 32) | (int(config.blocks) << 16) | int(config.threads) + ) + + +def reduce_verdict(wrong: "torch.Tensor", group) -> "torch.Tensor": + """Make every rank agree on which candidates computed the wrong answer. + + ``MAX`` over a per-candidate "was wrong" flag, which is the same decision + as ``MIN`` over "was right": one rank seeing a mismatch condemns the + candidate everywhere. A rank-local verdict would let ranks profile + different candidate sets, and the autotuner's timing reduction then + deadlocks on the first divergence. + + Corruption is not necessarily uniform across ranks: the cross-island race + this protocol can produce leaves some of them clean, so a rank-local verdict + can miss it entirely. + + Factored out so a test can assert the operator without a GPU. + """ + dist.all_reduce(wrong, op=dist.ReduceOp.MAX, group=group) + return wrong + + +def tuned_batches_for( + hidden: int, batches: Tuple[int, ...], max_numel: int +) -> Tuple[int, ...]: + """Drop buckets that would exceed the workspace at this hidden size.""" + return tuple(b for b in batches if b * hidden <= max_numel) + + +def warn_no_tune_group(stacklevel: int = 2) -> None: + """Say why a tuning session left this collective untuned. + + Raised from two places that reach the same dead end -- the runner, when the + autotuner does ask it for candidates, and the workspace, when it declines + to ask at all. The generic "nothing is tuned" advice does not fit here: the + caller *is* tuning, so telling them to tune is a dead end. What they are + missing is the reduction group, and that is what this names. + """ + warnings.warn( + "PCIe IPC all-reduce skipped autotuning: no matching " + "autotune process group is installed on every rank. Call " + "PcieIpcAllReduceWorkspace.tune(), or install one with " + "set_autotune_process_group() before entering autotune().", + RuntimeWarning, + stacklevel=stacklevel, + ) + + +class PcieIpcAllReduceRunner(TunableRunner): + """Adapts the all-reduce to the autotuner, and screens candidates first. + + One instance per workspace, built once and kept: the autotuner puts + ``hash(runner)`` in its in-memory cache key, so a fresh instance per call + would miss every entry and re-tune. + """ + + def __init__(self, workspace) -> None: + # A weak-ish coupling on purpose: the runner needs the raw launch and + # the group, not the public API, whose admission checks would run once + # per candidate and whose tracing decorator would recurse. + self._ws = workspace + # Named to end in _cache so the base __hash__ would skip it even if the + # override below is ever removed. + self._buf_cache: Dict[Tuple[Tuple[int, ...], torch.dtype], torch.Tensor] = {} + + def __hash__(self) -> int: + # Everything that changes what this runner does, and nothing that + # changes per call. The base implementation hashes __dict__ values and + # would fold in the workspace object's identity, which differs between + # processes and would defeat the persisted cache. + ws = self._ws + return hash( + ( + type(self).__name__, + PCIE_IPC_TUNE_VERSION, + ws.world_size, + ws.profile, + ws.max_blocks, + ws.max_numel, + ) + ) + + def get_cache_key_extras(self, inputs) -> Tuple: + ws = self._ws + return cache_key_extras( + ws.world_size, ws.profile, ws.max_blocks, ws.max_numel, inputs[0].dtype + ) + + def _output_for(self, inp: torch.Tensor) -> torch.Tensor: + key = (tuple(inp.shape), inp.dtype) + out = self._buf_cache.get(key) + if out is None: + out = torch.empty_like(inp) + self._buf_cache[key] = out + return out + + def _table_config(self, inp: torch.Tensor) -> Optional[IpcLaunchConfig]: + return self._ws.launch_config(inp) + + def can_profile(self, device) -> bool: + """Whether a real search is safe, as a group decision. + + Reduced rather than read locally because the answer decides how many + times each rank enters the profiler. Ranks that search different numbers + of candidates do not disagree, they deadlock. + """ + from ..autotuner import get_autotune_process_group + + group = get_autotune_process_group() + ok = group is not None and dist.get_world_size(group) == self._ws.world_size + flag = torch.tensor([1 if ok else 0], dtype=torch.int32, device=device) + dist.all_reduce(flag, op=dist.ReduceOp.MIN, group=self._ws.group) + return bool(flag.item()) + + def get_valid_tactics(self, inputs, profile) -> List: + """Candidates that computed the right answer, in a group-agreed order. + + This is the gate the autotuner does not have. It selects by ``argmin`` + on wall time and never inspects an output, while this kernel family's + characteristic failure -- a sentinel poll returning stale data rather + than waiting -- is wrong *and* fast. Screening here rather than during + profiling keeps the verdict's collective out of the timed window, and + costs one launch per candidate on a cache miss only. + + Cardinality is the hazard. Every rank must issue exactly these launches + in exactly this order; an early return between the first launch and the + verdict reduction leaves peers spinning inside a kernel this rank never + issued, with no timeout. Hence: barrier first, every buffer allocated + before the loop, and a loop body that does not allocate, synchronise + with the host, or branch. + """ + inp = inputs[0] + ws = self._ws + table_config = self._table_config(inp) + if table_config is None: + # The autotuner is being asked about a shape the kernels cannot run + # at all. Nothing to choose between; the caller falls back. + return [TABLE_TACTIC] + + if not self.can_profile(inp.device): + # Tuning mode is process-global, so this op can be swept by a + # caller that only meant to tune its GEMMs. Without a reduction + # over the candidate timings the ranks would argmin independently + # and pick different kernels, which this protocol does not survive. + # Offering only the seed degrades that into a no-op. + warn_no_tune_group(stacklevel=3) + return [TABLE_TACTIC] + + tactics = candidate_tactics(ws.world_size, ws.max_blocks) + configs = [table_config] + [tactic_to_config(t) for t in tactics] + + ref = inp.clone() + dist.all_reduce(ref, group=ws.group) + out = self._output_for(inp) + wrong = torch.zeros(len(configs), dtype=torch.int32, device=inp.device) + + dist.barrier(group=ws.group) + for i, config in enumerate(configs): + # A kernel that leaves part of the payload unwritten would + # otherwise show the previous candidate's correct result. + out.fill_(float("nan")) + ws._launch(inp, out, config) + wrong[i] = torch.ne(out, ref).any() + reduce_verdict(wrong, ws.group) + + verdict = wrong.tolist() + if verdict[0]: + # The seed computing the wrong answer is not something to route + # around: it is what every untuned shape and every cache miss falls + # back to. The verdict is group-wide, so + # every rank raises together and the group unwinds cleanly. + raise RuntimeError( + "the seed configuration for shape " + f"{tuple(inp.shape)} ({table_config}) does not match a " + "reference all-reduce; refusing to tune on top of it" + ) + survivors = zip(tactics, verdict[1:], strict=True) + return [TABLE_TACTIC] + [t for t, bad in survivors if not bad] + + def forward( + self, inputs, tactic=TABLE_TACTIC, do_preparation: bool = False, **kwargs + ): + inp = inputs[0] + out = self._output_for(inp) + if do_preparation: + # Buffer now allocated; launching here would make the call counts + # depend on whether the autotuner decided to prepare. + return out + table_config = self._table_config(inp) + if table_config is None: + raise RuntimeError( + f"shape {tuple(inp.shape)} is not one the kernels support; " + "the tuner must not have been asked about it" + ) + config = resolve_tuned_config( + table_config, tactic, self._ws.world_size, self._ws.max_blocks + ) + self._ws._launch(inp, out, config) + return out + + +__all__ = [ + "PCIE_IPC_CUSTOM_OP", + "PcieIpcAllReduceRunner", + "PCIE_IPC_TUNE_VERSION", + "TABLE_TACTIC", + "TUNE_BATCHES", + "TUNE_BLOCKS", + "TUNE_REPEAT", + "TUNE_THREADS", + "TUNE_WARMUP", + "cache_key_extras", + "candidate_tactics", + "config_to_tactic", + "default_cache_path", + "pack_config", + "pcie_ipc_tuning_config", + "reduce_verdict", + "resolve_tuned_config", + "small_int_initializer", + "tactic_to_config", + "tuned_batches_for", +] diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/csrc/pcie_ipc_all_reduce.cu b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/csrc/pcie_ipc_all_reduce.cu new file mode 100644 index 0000000..bdbf102 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/csrc/pcie_ipc_all_reduce.cu @@ -0,0 +1,206 @@ +/* + * Copyright (c) 2026 by FlashInfer team. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#include + +#include + +#include "flashinfer/comm/pcie_ipc_all_reduce.cuh" +#include "tvm_ffi_utils.h" + +namespace fi = flashinfer::comm::pcie_ipc; + +using tvm::ffi::Array; + +// Opaque handle, matching the fptr_t convention used by the other custom +// all-reduce bindings in this directory. +using fptr_t = int64_t; +static_assert(sizeof(void*) == sizeof(fptr_t)); + +namespace { + +// Everything the launcher needs that does not change between calls. The +// workspace itself is owned by the caller (see pcie_ipc_all_reduce.cuh). +struct PcieIpcHandle { + fi::PeerViews views; + fi::WorkspaceLayout layout; + int rank; + int world_size; + int max_blocks; + int64_t max_numel; + int elem_size; +}; + +} // namespace + +/*! + * \brief Bytes each rank must allocate and share over CUDA IPC. + * + * The caller passes the result to create_shared_buffer() and hands the + * resulting pointer array to pcie_ipc_init(). + */ +int64_t pcie_ipc_workspace_size(int64_t world_size, int64_t max_numel, int64_t elem_size, + int64_t max_blocks) { + TVM_FFI_ICHECK(world_size == 2 || world_size == 4 || world_size == 8) + << "pcie ipc all-reduce supports world_size 2, 4 or 8, got " << world_size; + TVM_FFI_ICHECK_GT(max_numel, 0) << "max_numel must be positive"; + TVM_FFI_ICHECK_EQ(elem_size, 2) + << "only 2-byte dtypes (bfloat16, float16) are supported, got elem_size " << elem_size; + TVM_FFI_ICHECK_GT(max_blocks, 0) << "max_blocks must be positive"; + return fi::workspace_size(static_cast(world_size), max_numel, static_cast(elem_size), + static_cast(max_blocks)); +} + +/*! + * \brief Bind an already-shared workspace and return an opaque handle. + * + * \param ipc_ptrs Peer pointers; entry i must address rank i's slab. + * + * The slab is zeroed here because the sentinel protocol reads +0.0 as "not yet + * written". The caller MUST barrier after this returns and before the first + * collective: a peer that starts pushing into this slab before we zero it + * would lose its payload. + */ +fptr_t pcie_ipc_init(Array ipc_ptrs, int64_t rank, int64_t max_numel, int64_t elem_size, + int64_t max_blocks) { + const int world_size = static_cast(ipc_ptrs.size()); + TVM_FFI_ICHECK(world_size == 2 || world_size == 4 || world_size == 8) + << "pcie ipc all-reduce supports world_size 2, 4 or 8, got " << world_size; + TVM_FFI_ICHECK(rank >= 0 && rank < world_size) << "rank " << rank << " out of range"; + TVM_FFI_ICHECK_EQ(elem_size, 2) + << "only 2-byte dtypes (bfloat16, float16) are supported, got elem_size " << elem_size; + TVM_FFI_ICHECK_GT(max_blocks, 0) << "max_blocks must be positive"; + + int64_t ptrs[fi::kMaxWorldSize]; + for (int i = 0; i < world_size; ++i) { + TVM_FFI_ICHECK_NE(ipc_ptrs[i], 0) << "ipc_ptrs[" << i << "] is null"; + ptrs[i] = ipc_ptrs[i]; + } + + auto* handle = new PcieIpcHandle(); + handle->layout = fi::compute_workspace_layout(world_size, max_numel, static_cast(elem_size), + static_cast(max_blocks)); + handle->views = fi::make_peer_views(ptrs, world_size, static_cast(rank), handle->layout); + handle->rank = static_cast(rank); + handle->world_size = world_size; + handle->max_blocks = static_cast(max_blocks); + handle->max_numel = max_numel; + handle->elem_size = static_cast(elem_size); + + cudaError_t err = cudaMemset(reinterpret_cast(ptrs[rank]), 0, handle->layout.total_bytes); + if (err != cudaSuccess) { + delete handle; + TVM_FFI_LOG_AND_THROW(RuntimeError) + << "failed to zero the pcie ipc workspace: " << cudaGetErrorString(err); + } + return reinterpret_cast(handle); +} + +void pcie_ipc_dispose(fptr_t handle) { delete reinterpret_cast(handle); } + +/*! + * \brief Out-of-place all-reduce over the shared workspace. + * + * \param blocks,threads,variant Launch configuration chosen by the caller; + * \c variant is a fi::Variant and the (world_size, variant) pairs that + * dispatch are listed in pcie_ipc_all_reduce.cuh. + */ +void pcie_ipc_all_reduce(fptr_t handle, TensorView inp, TensorView out, int64_t blocks, + int64_t threads, int64_t variant, bool enable_pdl) { + auto* h = reinterpret_cast(handle); + ffi::CUDADeviceGuard device_guard(inp.device().device_id); + auto stream = get_stream(inp.device()); + + TVM_FFI_ICHECK(inp.IsContiguous() && out.IsContiguous()) << "input and output must be contiguous"; + TVM_FFI_ICHECK_EQ(encode_dlpack_dtype(inp.dtype()), encode_dlpack_dtype(out.dtype())) + << "input and output dtype must match"; + TVM_FFI_ICHECK_EQ(inp.numel(), out.numel()) << "input and output must have the same size"; + + const int64_t numel = inp.numel(); + const int64_t elem_size = get_element_size(inp); + TVM_FFI_ICHECK_EQ(elem_size, h->elem_size) + << "dtype element size " << elem_size << " does not match the workspace's " << h->elem_size; + TVM_FFI_ICHECK_LE(static_cast(numel * elem_size), h->layout.max_payload_bytes) + << "payload exceeds the workspace capacity"; + + const int64_t pack_elems = 16 / elem_size; + TVM_FFI_ICHECK_EQ(numel % pack_elems, 0) + << "numel must be divisible by the 16-byte pack width (" << pack_elems << ")"; + TVM_FFI_ICHECK_EQ(h->max_numel % pack_elems, 0) + << "max_numel must be divisible by the 16-byte pack width"; + TVM_FFI_ICHECK(blocks > 0 && blocks <= h->max_blocks) + << "blocks must be in (0, " << h->max_blocks << "], got " << blocks; + TVM_FFI_ICHECK(threads > 0 && threads <= 1024) << "threads must be in (0, 1024], got " << threads; + // Every barrier signals from threadIdx.x < world_size, so a narrower block + // leaves some peers with nobody to signal them and the collective hangs. + TVM_FFI_ICHECK_GE(threads, h->world_size) + << "threads must be at least world_size (" << h->world_size << "), got " << threads; + // Refused rather than silently wrong: ipc_topo_rsag8_block_param_kernel + // triggers launch completion before island_owner_ack and its barrier flag + // store, so a dependent kernel can start while this call's phase-4 state is + // still being written. Re-enabling needs that release moved past both stores, + // an audit of the other six, and an SM90 regression. + TVM_FFI_ICHECK(!enable_pdl) + << "enable_pdl is not supported yet: in the TP8 block kernel the launch-completion " + "trigger precedes the island ack and barrier flag stores"; + TVM_FFI_ICHECK(variant >= 0 && variant < fi::kVariantCount) + << "variant must be in [0, " << fi::kVariantCount << "), got " << variant; + const auto algo = static_cast(variant); + // Reject rather than silently alias, so one configuration always names one + // kernel. + TVM_FFI_ICHECK( + !(h->world_size == 2 && algo != fi::Variant::kUnstaged && algo != fi::Variant::kStaged)) + << "world_size 2 accepts only kUnstaged and kStaged, got variant " << variant; + TVM_FFI_ICHECK(!(algo == fi::Variant::kFlatStaged && h->world_size != 8)) + << "kFlatStaged is world_size 8 only, got " << h->world_size; + // Only the block-partitioned TP8 kernel needs this: it derives its chunk + // from blockIdx.x & 3. Every other kernel uses flat grid-stride loops and + // accepts any block count. + if (h->world_size == 8 && algo == fi::Variant::kStaged) { + TVM_FFI_ICHECK_EQ(blocks % 4, 0) + << "the TP8 topology kernel requires blocks divisible by 4, got " << blocks; + } + + cudaError_t err = cudaSuccess; + switch (encode_dlpack_dtype(out.dtype())) { + case bfloat16_code: + err = fi::all_reduce(static_cast(inp.data_ptr()), + static_cast(out.data_ptr()), numel, h->views, + h->rank, h->world_size, h->max_blocks, h->max_numel, + static_cast(blocks), static_cast(threads), algo, + enable_pdl, stream); + break; + case float16_code: + err = fi::all_reduce( + static_cast(inp.data_ptr()), static_cast(out.data_ptr()), numel, + h->views, h->rank, h->world_size, h->max_blocks, h->max_numel, static_cast(blocks), + static_cast(threads), algo, enable_pdl, stream); + break; + default: + // The kernel templates carry a generic path, but only the two 2-byte + // dtypes are instantiated and measured. + TVM_FFI_LOG_AND_THROW(NotImplementedError) + << "pcie ipc all-reduce supports bfloat16 and float16 only"; + } + if (err != cudaSuccess) { + TVM_FFI_LOG_AND_THROW(RuntimeError) + << "pcie ipc all-reduce launch failed: " << cudaGetErrorString(err); + } +} + +TVM_FFI_DLL_EXPORT_TYPED_FUNC(pcie_ipc_workspace_size, pcie_ipc_workspace_size); +TVM_FFI_DLL_EXPORT_TYPED_FUNC(pcie_ipc_init, pcie_ipc_init); +TVM_FFI_DLL_EXPORT_TYPED_FUNC(pcie_ipc_dispose, pcie_ipc_dispose); +TVM_FFI_DLL_EXPORT_TYPED_FUNC(pcie_ipc_all_reduce, pcie_ipc_all_reduce); diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/include/flashinfer/comm/pcie_ipc_all_reduce.cuh b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/include/flashinfer/comm/pcie_ipc_all_reduce.cuh new file mode 100644 index 0000000..0f989d8 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/data/include/flashinfer/comm/pcie_ipc_all_reduce.cuh @@ -0,0 +1,2247 @@ +/* + * Copyright (c) 2026 by FlashInfer team. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef FLASHINFER_COMM_PCIE_IPC_ALL_REDUCE_CUH_ +#define FLASHINFER_COMM_PCIE_IPC_ALL_REDUCE_CUH_ + +// Custom all-reduce for intra-node PCIe machines without NVLink. +// +// Every peer transfer on such a machine crosses the CPU root complex, where +// all-to-all writes collapse to a fraction of what the same kernel achieves +// when each rank writes to a single destination. The kernels here therefore +// stage their pushes so that at any instant each rank has exactly one outbound +// and one inbound stream, and the 8-rank path keeps a 4+4 island decomposition +// so the scarce cross-socket links carry the minimum traffic. See the PR +// description for the bandwidth measurements this is derived from; they are a +// property of the machine, not of the code. +// +// All state lives in a caller-owned workspace shared over CUDA IPC; see +// compute_workspace_layout() for the byte layout and make_peer_views() for the +// per-region pointers. The caller owns the allocation because tearing it down +// needs a collective barrier between "every rank unmaps its peers" and "every +// rank frees its own slab", which a destructor cannot express. + +#include +#include +#include + +#include +#include +#include + +namespace flashinfer { +namespace comm { +namespace pcie_ipc { + +constexpr int kMaxWorldSize = 8; +constexpr int kSignalPhases = 8; + +// Which kernel all_reduce() launches, together with world_size. Values are +// part of the FFI signature, so they are explicit and append-only. +// +// kFlatStaged is accepted at world_size 8 only: at 4 it would name the same +// kernel as kStaged, and at 2 there is no staged-vs-flat distinction. +enum class Variant : int { + kUnstaged = 0, // push to every peer at once + kStaged = 1, // staged pushes; island-decomposed at world_size 8 + kStagedRing = 2, // staged pushes in neighbour order; world_size 4 and 8 + kFlatStaged = 3, // staged pushes without the island decomposition +}; + +constexpr int kVariantCount = 4; + +// Which staging area a kernel uses. The kernels come in two protocol families +// and a region may hold only one of them. +// +// Sentinel kernels poll for +0.0 meaning "not yet written", sanitise real zeros +// out of the payload, and store +0.0 back once a poll succeeds. Barrier kernels +// are content-blind: they publish raw payload and leave it there. Nothing else +// sweeps the workspace -- the host zeroes it once at init and never again. +// +// So a sentinel kernel landing on a barrier kernel's leftovers reads stale +// payload, and its all-gather poll, which watches a single slot, exits on it +// immediately: wrong output, not a hang. The epoch double buffer does not +// substitute for this -- it guarantees the other half is quiescent, not clean. +// +// At world_size 8 that puts the two topology kernels in kBlock and both +// sentinel kernels in kPack. +enum class ScratchRegion : int { kBlock = 0, kPack = 1 }; + +template +struct alignas(sizeof(T) * N) Vec { + T data[N]; +}; + +template +struct PackTraits { + static constexpr int kPackElems = 16 / sizeof(T); + using Pack = Vec; +}; + +template +__device__ __forceinline__ void pdl_grid_sync_const() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (Enabled) { + cudaGridDependencySynchronize(); + } +#endif +} + +template +__device__ __forceinline__ void pdl_grid_release_const() { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 900 + if constexpr (Enabled) { + __syncthreads(); + __threadfence(); + if (threadIdx.x == 0) { + cudaTriggerProgrammaticLaunchCompletion(); + } + } +#endif +} + +__device__ __forceinline__ void store_release_i32(int32_t* addr, int32_t value) { +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700 + asm volatile("st.release.sys.global.u32 [%1], %0;" ::"r"(value), "l"(addr)); +#else + asm volatile("membar.sys; st.volatile.global.u32 [%1], %0;" ::"r"(value), "l"(addr)); +#endif +} + +__device__ __forceinline__ int32_t load_acquire_i32(int32_t* addr) { + int32_t value; +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 700 + asm volatile("ld.acquire.sys.global.u32 %0, [%1];" : "=r"(value) : "l"(addr)); +#else + asm volatile("ld.volatile.global.u32 %0, [%1]; membar.gl;" : "=r"(value) : "l"(addr)); +#endif + return value; +} + +// Has the peer reached generation `expected`? +// +// Generations are a free-running counter, so a plain `observed < expected` +// breaks the first time it wraps: the slot still holds the old maximum, the +// new generation has wrapped to the minimum, and the comparison lets the +// barrier through before the peer has arrived. Switching to unsigned does not +// fix it either -- it just moves the break to UINT_MAX -> 0. +// +// Compare on the circle instead: reinterpret the difference as a signed +// distance modulo 2^32. `observed - expected >= 0` is then true exactly when +// the peer is at or past `expected`, for any pair within 2^31 generations of +// each other -- which is always, since a rank advances one generation per +// call and its peers are at most one call behind. +// constexpr so the boundary behaviour can be pinned at compile time, below. +// Those static_asserts are the whole defence against this being "simplified" +// back to a plain `<`: that version is correct for two billion calls and then +// releases every barrier a generation early, which no test anyone would +// actually run is going to catch. +__host__ __device__ __forceinline__ constexpr bool generation_reached(int32_t observed, + int32_t expected) { + return static_cast(static_cast(observed) - static_cast(expected)) >= + 0; +} + +static_assert(generation_reached(5, 5), "a peer at the expected generation has arrived"); +static_assert(generation_reached(6, 5), "a peer past the expected generation has arrived"); +static_assert(!generation_reached(4, 5), "a peer one generation behind has not arrived"); +// The wrap that motivated this function. Signed `observed < expected` reads +// INT32_MAX < INT32_MIN as false and lets the barrier through. +static_assert(generation_reached(INT32_MIN, INT32_MAX), + "the generation after INT32_MAX has arrived"); +static_assert(!generation_reached(INT32_MAX, INT32_MIN), + "the generation before the wrap has not arrived"); +// Unsigned `<` fixes the pair above but breaks this one, at UINT32_MAX -> 0. +// Only the modular comparison gets both. +static_assert(generation_reached(0, -1), "0 is one generation past -1"); +static_assert(!generation_reached(-1, 0), "-1 is one generation before 0"); + +__device__ __forceinline__ void store_volatile_i32(int32_t* addr, int32_t value) { + asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(value), "l"(addr)); +} + +__device__ __forceinline__ int32_t load_volatile_i32(int32_t* addr) { + int32_t value; + asm volatile("ld.volatile.global.u32 %0, [%1];" : "=r"(value) : "l"(addr)); + return value; +} + +__device__ __forceinline__ float to_float(float x) { return x; } +__device__ __forceinline__ float to_float(half x) { return __half2float(x); } +__device__ __forceinline__ float to_float(nv_bfloat16 x) { return __bfloat162float(x); } + +template +__device__ __forceinline__ T from_float(float x); + +template <> +__device__ __forceinline__ float from_float(float x) { + return x; +} + +template <> +__device__ __forceinline__ half from_float(float x) { + return __float2half(x); +} + +template <> +__device__ __forceinline__ nv_bfloat16 from_float(float x) { + return __float2bfloat16(x); +} + +__device__ __forceinline__ uint32_t add_half2_u32(uint32_t a, uint32_t b) { + auto ah = *reinterpret_cast(&a); + auto bh = *reinterpret_cast(&b); + half2 out = __hadd2(ah, bh); + return *reinterpret_cast(&out); +} + +__device__ __forceinline__ uint32_t add_bfloat162_u32(uint32_t a, uint32_t b) { + auto ah = *reinterpret_cast<__nv_bfloat162*>(&a); + auto bh = *reinterpret_cast<__nv_bfloat162*>(&b); + __nv_bfloat162 out = __hadd2(ah, bh); + return *reinterpret_cast(&out); +} + +template +__device__ __forceinline__ uint4 packed_add_u4(uint4 a, uint4 b) { + if constexpr (std::is_same_v) { + a.x = add_half2_u32(a.x, b.x); + a.y = add_half2_u32(a.y, b.y); + a.z = add_half2_u32(a.z, b.z); + a.w = add_half2_u32(a.w, b.w); + } else { + static_assert(std::is_same_v); + a.x = add_bfloat162_u32(a.x, b.x); + a.y = add_bfloat162_u32(a.y, b.y); + a.z = add_bfloat162_u32(a.z, b.z); + a.w = add_bfloat162_u32(a.w, b.w); + } + return a; +} + +template +__device__ __forceinline__ float2 lane_to_float2(uint32_t lane); + +template <> +__device__ __forceinline__ float2 lane_to_float2(uint32_t lane) { + auto value = *reinterpret_cast(&lane); + return __half22float2(value); +} + +template <> +__device__ __forceinline__ float2 lane_to_float2(uint32_t lane) { + auto value = *reinterpret_cast<__nv_bfloat162*>(&lane); + return __bfloat1622float2(value); +} + +template +__device__ __forceinline__ uint32_t float2_to_lane(float2 value); + +template <> +__device__ __forceinline__ uint32_t float2_to_lane(float2 value) { + half2 out = __float22half2_rn(value); + return *reinterpret_cast(&out); +} + +template <> +__device__ __forceinline__ uint32_t float2_to_lane(float2 value) { + __nv_bfloat162 out = __float22bfloat162_rn(value); + return *reinterpret_cast(&out); +} + +template +__device__ __forceinline__ uint4 reduce_u4_fp32(uint4 const (&values)[WorldSize]) { + float2 acc0 = {0.0f, 0.0f}; + float2 acc1 = {0.0f, 0.0f}; + float2 acc2 = {0.0f, 0.0f}; + float2 acc3 = {0.0f, 0.0f}; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + float2 v0 = lane_to_float2(values[peer].x); + float2 v1 = lane_to_float2(values[peer].y); + float2 v2 = lane_to_float2(values[peer].z); + float2 v3 = lane_to_float2(values[peer].w); + acc0.x += v0.x; + acc0.y += v0.y; + acc1.x += v1.x; + acc1.y += v1.y; + acc2.x += v2.x; + acc2.y += v2.y; + acc3.x += v3.x; + acc3.y += v3.y; + } + uint4 out; + out.x = float2_to_lane(acc0); + out.y = float2_to_lane(acc1); + out.z = float2_to_lane(acc2); + out.w = float2_to_lane(acc3); + return out; +} + +template +struct ZeroBits; + +template <> +struct ZeroBits { + using Raw = uint16_t; + static constexpr Raw kPos = 0x0000u; + static constexpr Raw kNeg = 0x8000u; +}; + +template <> +struct ZeroBits { + using Raw = uint16_t; + static constexpr Raw kPos = 0x0000u; + static constexpr Raw kNeg = 0x8000u; +}; + +template <> +struct ZeroBits { + using Raw = uint32_t; + static constexpr Raw kPos = 0x00000000u; + static constexpr Raw kNeg = 0x80000000u; +}; + +template +__device__ __forceinline__ void clear_pos_zero(T& value) { + using Bits = ZeroBits; + using Raw = typename Bits::Raw; + Raw* raw = reinterpret_cast(&value); + if (*raw == Bits::kPos) { + *raw = Bits::kNeg; + } +} + +template +__device__ __forceinline__ bool is_pos_zero(T value) { + using Bits = ZeroBits; + using Raw = typename Bits::Raw; + Raw raw = *reinterpret_cast(&value); + return raw == Bits::kPos; +} + +template +__device__ __forceinline__ T pos_zero() { + using Bits = ZeroBits; + using Raw = typename Bits::Raw; + Raw raw = Bits::kPos; + return *reinterpret_cast(&raw); +} + +template +__device__ __forceinline__ typename PackTraits::Pack load_pack_volatile( + typename PackTraits::Pack const* base, int idx) { + uint4 raw; + auto const* addr = reinterpret_cast(base + idx); + asm volatile("ld.volatile.global.v4.b32 {%0, %1, %2, %3}, [%4];" + : "=r"(raw.x), "=r"(raw.y), "=r"(raw.z), "=r"(raw.w) + : "l"(addr)); + return *reinterpret_cast::Pack*>(&raw); +} + +template +__device__ __forceinline__ void store_pack_volatile(typename PackTraits::Pack* base, int idx, + typename PackTraits::Pack value) { + uint4 raw = *reinterpret_cast(&value); + auto* addr = reinterpret_cast(base + idx); + asm volatile("st.volatile.global.v4.b32 [%4], {%0, %1, %2, %3};" ::"r"(raw.x), "r"(raw.y), + "r"(raw.z), "r"(raw.w), "l"(addr)); +} + +template +__device__ __forceinline__ void clear_pos_zero_pack(typename PackTraits::Pack& pack) { +#pragma unroll + for (int i = 0; i < PackTraits::kPackElems; ++i) { + clear_pos_zero(pack.data[i]); + } +} + +template +__device__ __forceinline__ bool has_pos_zero_pack(typename PackTraits::Pack const& pack) { + bool has_zero = false; +#pragma unroll + for (int i = 0; i < PackTraits::kPackElems; ++i) { + has_zero |= is_pos_zero(pack.data[i]); + } + return has_zero; +} + +template +__device__ __forceinline__ typename PackTraits::Pack zero_pack() { + typename PackTraits::Pack pack; +#pragma unroll + for (int i = 0; i < PackTraits::kPackElems; ++i) { + pack.data[i] = pos_zero(); + } + return pack; +} + +__device__ __forceinline__ uint32_t clear_pos_zero_u16x2(uint32_t raw) { + uint32_t lo = raw & 0xffffu; + uint32_t hi = raw & 0xffff0000u; + if (lo == 0u) { + lo = 0x8000u; + } + if (hi == 0u) { + hi = 0x80000000u; + } + return hi | lo; +} + +__device__ __forceinline__ bool has_pos_zero_u16x2(uint32_t raw) { + return (raw & 0xffffu) == 0u || (raw & 0xffff0000u) == 0u; +} + +__device__ __forceinline__ uint4 clear_pos_zero_u4_16(uint4 value) { + value.x = clear_pos_zero_u16x2(value.x); + value.y = clear_pos_zero_u16x2(value.y); + value.z = clear_pos_zero_u16x2(value.z); + value.w = clear_pos_zero_u16x2(value.w); + return value; +} + +__device__ __forceinline__ bool has_pos_zero_u4_16(uint4 value) { + return has_pos_zero_u16x2(value.x) || has_pos_zero_u16x2(value.y) || + has_pos_zero_u16x2(value.z) || has_pos_zero_u16x2(value.w); +} + +__device__ __forceinline__ uint4 load_u4_volatile(uint4 const* base, int idx) { + uint4 value; + auto const* addr = base + idx; + asm volatile("ld.volatile.global.v4.b32 {%0, %1, %2, %3}, [%4];" + : "=r"(value.x), "=r"(value.y), "=r"(value.z), "=r"(value.w) + : "l"(addr)); + return value; +} + +__device__ __forceinline__ void store_u4_volatile(uint4* base, int idx, uint4 value) { + auto* addr = base + idx; + asm volatile("st.volatile.global.v4.b32 [%4], {%0, %1, %2, %3};" ::"r"(value.x), "r"(value.y), + "r"(value.z), "r"(value.w), "l"(addr)); +} + +__device__ __forceinline__ int phase_offset(int phase, int block, int peer, int max_blocks, + int world_size) { + // Epoch slots occupy [0, max_blocks). Barrier slots start after that + // dedicated prefix so phase 0 can never corrupt an epoch. + return max_blocks + phase * max_blocks * world_size + block * world_size + peer; +} + +__device__ __forceinline__ int flag_offset(int block, int max_blocks, int world_size) { + return max_blocks + kSignalPhases * max_blocks * world_size + block; +} + +// Call-level double-buffer state: the epoch at [0] and the arrival counter at +// [1], placed past the flag region so no existing offset moves. +// +// ONE PAIR PER SCRATCH REGION. The hard constraint runs one way only: +// +// **Kernels writing the same region MUST share a counter.** TP4 alternates +// ipc_rsag_push and ipc_rsag_ring with the payload and both write the block +// region at the same addresses, so private counters would let a ring call on +// half 0 be followed immediately by a push call that also reads 0 -- back to +// back, with nothing in between to drain the first. +// +// Kernels in *different* regions need not be: any intervening collective +// already drains the previous one. They are kept apart anyway, because binding +// the state on the same host line that picks the region makes the two +// impossible to get out of step. +// +// Rank-local, and indexed by nothing, so every CTA of a launch agrees on the +// half regardless of gridDim. Per-block parity cannot: it counts how many times +// *that block* has run, so one change in gridDim desynchronises the block ranges +// permanently and a block picks a half another block is still using. +// +// Consistency across ranks comes from the same SPMD argument that already +// backs the per-block barrier flags: every rank runs the same sequence of +// collectives, so every rank is on the same call parity. +__host__ __device__ __forceinline__ int scratch_state_offset(int max_blocks, int world_size, + ScratchRegion region) { + return max_blocks + kSignalPhases * max_blocks * world_size + max_blocks + + 2 * static_cast(region); +} + +// Debug only: pin every call to half 0, i.e. the pre-double-buffer behaviour. +// Kept because the cross-island race is invisible without a way to build the +// broken protocol on demand -- it is what proves a repro actually has power, +// and it isolates the cost of the double buffer in a single benchmark session. +// Never define this in a shipping build. +// Debug only: remove the leading CTA barrier from the three signalling helpers. +// The resulting build is incorrect -- a warp can announce "my stage is done" +// while its siblings are still writing. Never define this in a shipping build. +#ifndef FLASHINFER_PCIE_IPC_DEBUG_NO_BARRIER_ENTRY_SYNC +#define FLASHINFER_PCIE_IPC_DEBUG_NO_BARRIER_ENTRY_SYNC 0 +#endif + +#ifndef FLASHINFER_PCIE_IPC_DEBUG_NO_BLOCK_EPOCH +#define FLASHINFER_PCIE_IPC_DEBUG_NO_BLOCK_EPOCH 0 +#endif + +// Debug only: restore the per-block epoch parity these kernels used before the +// call-level counter replaced it. This is the negative control for the +// grid-change regression -- without it, a passing test cannot distinguish "the +// fix works" from "the sequence never opened the window". Distinct from +// NO_BLOCK_EPOCH above, which removes double buffering entirely; this one keeps +// two halves and only breaks the agreement about which half a call is on. +// Never define this in a shipping build. +#ifndef FLASHINFER_PCIE_IPC_DEBUG_PER_BLOCK_EPOCH +#define FLASHINFER_PCIE_IPC_DEBUG_PER_BLOCK_EPOCH 0 +#endif + +// Read this call's epoch and advance it, both at kernel entry. +// +// Advancing at entry rather than at exit is what keeps this affordable. The +// flip only has to follow every block's *read*, not every block's work: the +// state is rank-local (views.self_signal), so its only reader is the next +// kernel on this stream, and that cannot start until this launch has fully +// retired. The last block to arrive therefore knows every peer block has +// already read, and can flip immediately. +// +// Committing at exit instead would put a tail __syncthreads() plus an L2 atomic +// round trip on the block-retirement critical path. It would also be fragile: +// an early return added after the arrival would freeze the counter below +// gridDim.x - 1 and silently pin the epoch. +// +// Electing the last arrival deliberately does not depend on block scheduling +// order or on the whole grid being resident. "Block 0 flips" would need every +// block launched before block 0 reaches this point, which CUDA does not +// promise, and a grid-wide spin barrier deadlocks once gridDim exceeds +// occupancy. +// +// At gridDim.x == 1 there is nothing to elect, so the atomic is skipped; the +// end state is identical (counter 0, epoch flipped). +// Exit half of the PER_BLOCK_EPOCH debug build; compiles to nothing otherwise. +// The flip sits at the exit because that is where the implementation it rebuilds +// put it, and entry-versus-exit changes the cost. +__device__ __forceinline__ void debug_commit_per_block_epoch(int32_t* per_block_slot, int epoch) { +#if FLASHINFER_PCIE_IPC_DEBUG_PER_BLOCK_EPOCH + // Bare -- no __syncthreads(), matching the implementation this rebuilds. + if (threadIdx.x == 0) { + store_volatile_i32(per_block_slot, epoch ^ 1); + } +#else + (void)per_block_slot; + (void)epoch; +#endif +} + +__device__ __forceinline__ int advance_scratch_epoch(int32_t* state, int32_t* per_block_slot) { +#if FLASHINFER_PCIE_IPC_DEBUG_NO_BLOCK_EPOCH + (void)state; + (void)per_block_slot; + return 0; +#elif FLASHINFER_PCIE_IPC_DEBUG_PER_BLOCK_EPOCH + // Read only; the flip is issued at kernel exit by + // debug_commit_per_block_epoch(). + (void)state; + return load_volatile_i32(per_block_slot) & 1; +#else + (void)per_block_slot; + const int epoch = load_volatile_i32(state) & 1; + // Every thread must have read before this block announces its arrival. + __syncthreads(); + if (threadIdx.x == 0) { + // Device scope, not system scope. This state is rank-local -- its only + // reader is this rank's next kernel on this stream -- and that reader wants + // nothing but the value itself, so there is nothing for a release to order. + // st.release.sys would flush this thread's writes system-wide over PCIe on + // the entry critical path. + if (gridDim.x == 1) { + // Sole CTA: trivially the last arrival, and the counter is already 0. + store_volatile_i32(state, epoch ^ 1); + } else if (atomicAdd(state + 1, 1) == static_cast(gridDim.x) - 1) { + store_volatile_i32(state + 1, 0); + store_volatile_i32(state, epoch ^ 1); + } + } + return epoch; +#endif +} + +__device__ __forceinline__ void block_barrier(uint64_t const* signal_ptrs, int rank, int world_size, + int max_blocks, int phase, int flag) { + // Publishing a signal means "this CTA's writes for the previous stage are + // done", so every thread must have finished them before the signalling + // threads announce it. The call sites use __threadfence_system(), which + // orders only the *calling* thread's accesses -- it does not wait for the + // rest of the CTA, and cannot help at all where the previous stage was a load + // loop. Only a CTA barrier establishes that. +#if !FLASHINFER_PCIE_IPC_DEBUG_NO_BARRIER_ENTRY_SYNC + __syncthreads(); +#endif + int block = blockIdx.x; + int32_t* self = reinterpret_cast(signal_ptrs[rank]); + if (threadIdx.x < world_size) { + int peer = threadIdx.x; + int32_t* peer_signal = reinterpret_cast(signal_ptrs[peer]); + store_release_i32(peer_signal + phase_offset(phase, block, rank, max_blocks, world_size), flag); + int32_t* self_slot = self + phase_offset(phase, block, peer, max_blocks, world_size); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + __syncthreads(); +} + +__device__ __forceinline__ void block_barrier_mask(uint64_t const* signal_ptrs, int rank, + int world_size, int max_blocks, int phase, + int flag, uint32_t participant_mask) { + // Entry barrier: see block_barrier above. +#if !FLASHINFER_PCIE_IPC_DEBUG_NO_BARRIER_ENTRY_SYNC + __syncthreads(); +#endif + if ((participant_mask & (1u << rank)) == 0u) { + __syncthreads(); + return; + } + int block = blockIdx.x; + int32_t* self = reinterpret_cast(signal_ptrs[rank]); + if (threadIdx.x < world_size) { + int peer = threadIdx.x; + if ((participant_mask & (1u << peer)) != 0u) { + int32_t* peer_signal = reinterpret_cast(signal_ptrs[peer]); + store_release_i32(peer_signal + phase_offset(phase, block, rank, max_blocks, world_size), + flag); + int32_t* self_slot = self + phase_offset(phase, block, peer, max_blocks, world_size); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + } + __syncthreads(); +} + +__device__ __forceinline__ void island_owner_gather(uint64_t const* signal_ptrs, int rank, int base, + int owner, int max_blocks, int phase, + int flag) { + // Entry barrier: see block_barrier above. +#if !FLASHINFER_PCIE_IPC_DEBUG_NO_BARRIER_ENTRY_SYNC + __syncthreads(); +#endif + int block = blockIdx.x; + int32_t* owner_signal = reinterpret_cast(signal_ptrs[owner]); + if (threadIdx.x == 0) { + store_release_i32(owner_signal + phase_offset(phase, block, rank, max_blocks, 8), flag); + } + if (rank == owner && threadIdx.x < 4) { + int peer = base + threadIdx.x; + int32_t* self_slot = owner_signal + phase_offset(phase, block, peer, max_blocks, 8); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + __syncthreads(); +} + +__device__ __forceinline__ void owner_pair_barrier(uint64_t const* signal_ptrs, int rank, int owner, + int cross_owner, int max_blocks, int phase, + int flag) { + __syncthreads(); + if (rank == owner && threadIdx.x == 0) { + int block = blockIdx.x; + int32_t* cross_signal = reinterpret_cast(signal_ptrs[cross_owner]); + store_release_i32(cross_signal + phase_offset(phase, block, rank, max_blocks, 8), flag); + int32_t* self_signal = reinterpret_cast(signal_ptrs[rank]); + int32_t* self_slot = self_signal + phase_offset(phase, block, cross_owner, max_blocks, 8); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + __syncthreads(); +} + +__device__ __forceinline__ void island_owner_ready(uint64_t const* signal_ptrs, int rank, int base, + int owner, int max_blocks, int phase, int flag) { + __syncthreads(); + int block = blockIdx.x; + if (rank == owner) { + if (threadIdx.x < 4) { + int peer = base + threadIdx.x; + int32_t* peer_signal = reinterpret_cast(signal_ptrs[peer]); + store_release_i32(peer_signal + phase_offset(phase, block, owner, max_blocks, 8), flag); + } + } else if (threadIdx.x == 0) { + int32_t* self_signal = reinterpret_cast(signal_ptrs[rank]); + int32_t* self_slot = self_signal + phase_offset(phase, block, owner, max_blocks, 8); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + __syncthreads(); +} + +__device__ __forceinline__ void island_owner_ack(uint64_t const* signal_ptrs, int rank, int base, + int owner, int max_blocks, int phase, int flag) { + __syncthreads(); + int block = blockIdx.x; + int32_t* owner_signal = reinterpret_cast(signal_ptrs[owner]); + if (rank != owner && threadIdx.x == 0) { + store_release_i32(owner_signal + phase_offset(phase, block, rank, max_blocks, 8), flag); + } + if (rank == owner && threadIdx.x < 4) { + int peer = base + threadIdx.x; + if (peer != owner) { + int32_t* self_slot = owner_signal + phase_offset(phase, block, peer, max_blocks, 8); + while (!generation_reached(load_acquire_i32(self_slot), flag)) { + } + } + } + __syncthreads(); +} + +template +__device__ __forceinline__ typename PackTraits::Pack add_pack(typename PackTraits::Pack a, + typename PackTraits::Pack b) { + using Traits = PackTraits; + using Pack = typename Traits::Pack; + if constexpr (std::is_same_v || std::is_same_v) { + uint4 av = *reinterpret_cast(&a); + uint4 bv = *reinterpret_cast(&b); + uint4 out = packed_add_u4(av, bv); + return *reinterpret_cast(&out); + } + + Pack out; +#pragma unroll + for (int i = 0; i < Traits::kPackElems; ++i) { + out.data[i] = from_float(to_float(a.data[i]) + to_float(b.data[i])); + } + return out; +} + +template +__device__ __forceinline__ typename PackTraits::Pack reduce_loaded_packs( + typename PackTraits::Pack const (&values)[WorldSize]) { + using Pack = typename PackTraits::Pack; + if constexpr (std::is_same_v || std::is_same_v) { + uint4 acc = *reinterpret_cast(&values[0]); +#pragma unroll + for (int peer = 1; peer < WorldSize; ++peer) { + uint4 next = *reinterpret_cast(&values[peer]); + acc = packed_add_u4(acc, next); + } + return *reinterpret_cast(&acc); + } else { + Pack acc = values[0]; +#pragma unroll + for (int peer = 1; peer < WorldSize; ++peer) { + acc = add_pack(acc, values[peer]); + } + return acc; + } +} + +// Debug-only hook for reproducing the cross-island scratch race. +// +// The hazard needs the SLOW island to still be reading the cross slot while +// the FAST island's next call overwrites it. Delaying a whole kernel launch +// from the host cannot produce that: the pair barrier releases both islands +// together, after which the reader reaches its cross read almost immediately +// while the writer still has several phases to go. The stall has to be here, +// between the pair rendezvous and the cross read. +// +// Enabled only when FLASHINFER_PCIE_IPC_DEBUG_CROSS_STALL_NS is defined, and +// only on the island selected by ..._STALL_ISLAND. Never define these in a +// shipping build. +#ifndef FLASHINFER_PCIE_IPC_DEBUG_CROSS_STALL_NS +#define FLASHINFER_PCIE_IPC_DEBUG_CROSS_STALL_NS 0 +#endif +#ifndef FLASHINFER_PCIE_IPC_DEBUG_STALL_ISLAND +#define FLASHINFER_PCIE_IPC_DEBUG_STALL_ISLAND 0 +#endif + +__device__ __forceinline__ void debug_cross_read_stall(int rank) { +#if FLASHINFER_PCIE_IPC_DEBUG_CROSS_STALL_NS > 0 + const int island = rank < 4 ? 0 : 1; + if (island == FLASHINFER_PCIE_IPC_DEBUG_STALL_ISLAND) { + __nanosleep(FLASHINFER_PCIE_IPC_DEBUG_CROSS_STALL_NS); + } + __syncthreads(); +#else + (void)rank; +#endif +} + +template +struct PushOneshotParamData { + uint64_t tmp_ptrs[kMaxWorldSize]; + uint64_t signal_ptrs[kMaxWorldSize]; + T const* input; + T* output; + int32_t* epoch_slots; + // Call-level double-buffer state for this launch's scratch region, bound + // host-side so the state and the region cannot disagree. See + // scratch_state_offset(). + int32_t* scratch_state; + int num_packs; + int rank_stride_packs; + int epoch_stride_packs; + int rank; + int max_blocks; +}; + +template +struct IpcTp2RemotePushData { + uint64_t tmp_ptrs[2]; + T const* input; + T* output; + int32_t* epoch_slots; + // Call-level double-buffer state for this launch's scratch region, bound + // host-side so the state and the region cannot disagree. See + // scratch_state_offset(). + int32_t* scratch_state; + int num_packs; + // Half-size of the epoch double buffer, in packs. Derived from max_numel, + // NOT from this call's num_packs: the two epoch halves must sit at fixed + // addresses. If they moved with the payload, a rank that finished a large + // collective and flipped its epoch would start writing a small one inside + // the region a lagging peer is still draining -- corrupting it, or having + // the peer's reset wipe the just-published data so the poll never ends. + // Every other v2 kernel already derives its stage offset this way. + int rank_stride_packs; + int rank; +}; + +template +__global__ __launch_bounds__(1024, 1) void ipc_tp2_remote_push_kernel( + const IpcTp2RemotePushData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + pdl_grid_sync_const(); + + int peer = params.rank ^ 1; + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + int epoch = advance_scratch_epoch(params.scratch_state, epoch_slot); + int stage_offset = epoch * 2 * params.rank_stride_packs; + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int stride = gridDim.x * blockDim.x; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + int peer_write_base = stage_offset + params.rank * params.num_packs; + int local_poll_base = stage_offset + peer * params.num_packs; + uint4 reset = {0u, 0u, 0u, 0u}; + if constexpr (Stream) { + for (int idx = tid; idx < params.num_packs; idx += stride) { + uint4 local_value = input[idx]; + uint4 publish_value = clear_pos_zero_u4_16(local_value); + store_u4_volatile(peer_buffer, peer_write_base + idx, publish_value); + uint4 peer_value; + while (true) { + peer_value = load_u4_volatile(local_buffer, local_poll_base + idx); + if (!has_pos_zero_u4_16(peer_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = packed_add_u4(local_value, peer_value); + store_u4_volatile(local_buffer, local_poll_base + idx, reset); + } + } else { + for (int idx = tid; idx < params.num_packs; idx += stride) { + uint4 value = clear_pos_zero_u4_16(input[idx]); + store_u4_volatile(peer_buffer, peer_write_base + idx, value); + } + for (int idx = tid; idx < params.num_packs; idx += stride) { + uint4 peer_value; + while (true) { + peer_value = load_u4_volatile(local_buffer, local_poll_base + idx); + if (!has_pos_zero_u4_16(peer_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = packed_add_u4(input[idx], peer_value); + store_u4_volatile(local_buffer, local_poll_base + idx, reset); + } + } + } else { + Pack const* input = reinterpret_cast(params.input); + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + int peer_write_base = stage_offset + params.rank * params.num_packs; + int local_poll_base = stage_offset + peer * params.num_packs; + Pack reset = zero_pack(); + if constexpr (Stream) { + for (int idx = tid; idx < params.num_packs; idx += stride) { + Pack local_value = input[idx]; + Pack publish_value = local_value; + clear_pos_zero_pack(publish_value); + store_pack_volatile(peer_buffer, peer_write_base + idx, publish_value); + Pack peer_value; + while (true) { + peer_value = load_pack_volatile(local_buffer, local_poll_base + idx); + if (!has_pos_zero_pack(peer_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = add_pack(local_value, peer_value); + store_pack_volatile(local_buffer, local_poll_base + idx, reset); + } + } else { + for (int idx = tid; idx < params.num_packs; idx += stride) { + Pack value = input[idx]; + clear_pos_zero_pack(value); + store_pack_volatile(peer_buffer, peer_write_base + idx, value); + } + for (int idx = tid; idx < params.num_packs; idx += stride) { + Pack peer_value; + while (true) { + peer_value = load_pack_volatile(local_buffer, local_poll_base + idx); + if (!has_pos_zero_pack(peer_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = add_pack(input[idx], peer_value); + store_pack_volatile(local_buffer, local_poll_base + idx, reset); + } + } + } + + debug_commit_per_block_epoch(epoch_slot, epoch); + pdl_grid_release_const(); +} + +// Owner of a pack under the reduce-scatter split. Must agree with the explicit +// chunk ranges the same kernels walk, which give the remainder to the last rank +// -- so `part == 0` (fewer packs than ranks) gives it the whole payload too. +// Writing to one owner and polling another spins forever: no timeout here. +template +__device__ __forceinline__ int rsag_owner_for_pack(int idx, int part) { + int owner = part > 0 ? idx / part : WorldSize - 1; + return owner < WorldSize ? owner : WorldSize - 1; +} + +// Staged (neighbour-ordered) RS/AG push. +// +// ipc_rsag_push_param_kernel below has every rank writing to all WorldSize-1 +// peers at the same time. Where every peer transfer crosses the CPU root +// complex that pattern runs far below what the same kernel-issued writes reach +// when a rank has a single outbound destination -- the collective is limited by +// the shape of the traffic, not the amount. +// +// This variant keeps the algorithm and the total bytes identical and only +// reorders the pushes: each push phase is split into WorldSize-1 passes, and in +// pass p every rank writes solely to peer (rank + 1 + p) % WorldSize. That map +// is a permutation, so during a pass every GPU has exactly one outbound stream +// and every GPU is the target of exactly one -- the pattern the fabric +// sustains. A barrier between passes keeps the ranks in the same pass, since +// without it they drift and the passes overlap back into all-to-all. +// +// The cost is 2*(WorldSize-1) barriers per collective, which is why the policy +// only selects this variant once the payload is large enough to pay for them. +// Keep the established phase numbering for this kernel. Epoch slots and +// barrier phases have separate ranges in the signal region (see +// phase_offset), so phase 1 is no longer needed for alias avoidance. +constexpr int kRingRsPhase0 = 1; + +template +__global__ __launch_bounds__(1024, 1) void ipc_rsag_ring_push_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + // RS uses phases [1, WorldSize-1], AG uses [WorldSize, 2*WorldSize-2]. + static_assert(2 * WorldSize - 2 < kSignalPhases, + "ring push needs 2*(WorldSize-1) barrier phases"); + pdl_grid_sync_const(); + + int32_t* self_signal = reinterpret_cast(params.signal_ptrs[params.rank]); + int flag = static_cast( + static_cast( + load_acquire_i32(self_signal + flag_offset(blockIdx.x, params.max_blocks, WorldSize))) + + 1u); + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + int epoch = advance_scratch_epoch(params.scratch_state, epoch_slot); + int stage_offset = epoch * params.epoch_stride_packs; + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int stride = gridDim.x * blockDim.x; + int part = params.num_packs / WorldSize; + int my_start = params.rank * part; + int my_end = (params.rank == WorldSize - 1) ? params.num_packs : my_start + part; + int my_slot = stage_offset + params.rank * params.rank_stride_packs; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + uint4 reset = {0u, 0u, 0u, 0u}; + + // Reduce-scatter, own chunk: stays on this GPU, so it costs no fabric time + // and does not need a pass of its own. + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_u4_volatile(local_buffer, my_slot + idx, clear_pos_zero_u4_16(input[idx])); + } + + // Reduce-scatter, staged: one destination per pass. + for (int p = 0; p < WorldSize - 1; ++p) { + int target = (params.rank + 1 + p) % WorldSize; + int t_start = target * part; + int t_end = (target == WorldSize - 1) ? params.num_packs : t_start + part; + auto* target_buffer = reinterpret_cast(params.tmp_ptrs[target]); + for (int idx = t_start + tid; idx < t_end; idx += stride) { + store_u4_volatile(target_buffer, my_slot + idx, clear_pos_zero_u4_16(input[idx])); + } + block_barrier(params.signal_ptrs, params.rank, WorldSize, params.max_blocks, + kRingRsPhase0 + p, flag); + } + + // Owner reduce: every contribution is now in this rank's own buffer, so + // this phase touches local memory only. The reduced value is stashed back + // into this rank's own slot for the all-gather passes to re-read. + for (int idx = my_start + tid; idx < my_end; idx += stride) { + uint4 values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_u4_volatile(local_buffer, offset); + waiting |= has_pos_zero_u4_16(values[peer]); + } + if (!waiting) { + break; + } + } + uint4 acc = values[0]; +#pragma unroll + for (int peer = 1; peer < WorldSize; ++peer) { + acc = packed_add_u4(acc, values[peer]); + } + reinterpret_cast(params.output)[idx] = acc; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + if (peer != params.rank) { + store_u4_volatile(local_buffer, stage_offset + peer * params.rank_stride_packs + idx, + reset); + } + } + store_u4_volatile(local_buffer, my_slot + idx, clear_pos_zero_u4_16(acc)); + } + + // All-gather, staged: same permutation schedule as the reduce-scatter. + for (int p = 0; p < WorldSize - 1; ++p) { + int target = (params.rank + 1 + p) % WorldSize; + auto* target_buffer = reinterpret_cast(params.tmp_ptrs[target]); + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_u4_volatile(target_buffer, my_slot + idx, + load_u4_volatile(local_buffer, my_slot + idx)); + } + block_barrier(params.signal_ptrs, params.rank, WorldSize, params.max_blocks, WorldSize + p, + flag); + } + + // Consume the chunks owned by others, then clear this rank's own slot so + // the sentinel state is clean for the epoch that reuses it. + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + if (owner == params.rank) { + continue; + } + int offset = stage_offset + owner * params.rank_stride_packs + idx; + uint4 value; + while (true) { + value = load_u4_volatile(local_buffer, offset); + if (!has_pos_zero_u4_16(value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = value; + store_u4_volatile(local_buffer, offset, reset); + } + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_u4_volatile(local_buffer, my_slot + idx, reset); + } + } else { + Pack const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + Pack reset = zero_pack(); + + for (int idx = my_start + tid; idx < my_end; idx += stride) { + Pack value = input[idx]; + clear_pos_zero_pack(value); + store_pack_volatile(local_buffer, my_slot + idx, value); + } + for (int p = 0; p < WorldSize - 1; ++p) { + int target = (params.rank + 1 + p) % WorldSize; + int t_start = target * part; + int t_end = (target == WorldSize - 1) ? params.num_packs : t_start + part; + auto* target_buffer = reinterpret_cast(params.tmp_ptrs[target]); + for (int idx = t_start + tid; idx < t_end; idx += stride) { + Pack value = input[idx]; + clear_pos_zero_pack(value); + store_pack_volatile(target_buffer, my_slot + idx, value); + } + block_barrier(params.signal_ptrs, params.rank, WorldSize, params.max_blocks, + kRingRsPhase0 + p, flag); + } + for (int idx = my_start + tid; idx < my_end; idx += stride) { + Pack values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_pack_volatile(local_buffer, offset); + waiting |= has_pos_zero_pack(values[peer]); + } + if (!waiting) { + break; + } + } + Pack acc = reduce_loaded_packs(values); + reinterpret_cast(params.output)[idx] = acc; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + if (peer != params.rank) { + store_pack_volatile(local_buffer, stage_offset + peer * params.rank_stride_packs + idx, + reset); + } + } + Pack publish = acc; + clear_pos_zero_pack(publish); + store_pack_volatile(local_buffer, my_slot + idx, publish); + } + for (int p = 0; p < WorldSize - 1; ++p) { + int target = (params.rank + 1 + p) % WorldSize; + auto* target_buffer = reinterpret_cast(params.tmp_ptrs[target]); + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_pack_volatile(target_buffer, my_slot + idx, + load_pack_volatile(local_buffer, my_slot + idx)); + } + block_barrier(params.signal_ptrs, params.rank, WorldSize, params.max_blocks, WorldSize + p, + flag); + } + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + if (owner == params.rank) { + continue; + } + int offset = stage_offset + owner * params.rank_stride_packs + idx; + Pack value; + while (true) { + value = load_pack_volatile(local_buffer, offset); + if (!has_pos_zero_pack(value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = value; + store_pack_volatile(local_buffer, offset, reset); + } + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_pack_volatile(local_buffer, my_slot + idx, reset); + } + } + + if (threadIdx.x == 0) { + store_release_i32(self_signal + flag_offset(blockIdx.x, params.max_blocks, WorldSize), flag); + } + debug_commit_per_block_epoch(epoch_slot, epoch); + pdl_grid_release_const(); +} + +template +__global__ __launch_bounds__(1024, 1) void ipc_rsag_push_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + pdl_grid_sync_const(); + + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + int epoch = advance_scratch_epoch(params.scratch_state, epoch_slot); + int stage_offset = epoch * params.epoch_stride_packs; + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int stride = gridDim.x * blockDim.x; + int part = params.num_packs / WorldSize; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + uint4 reset = {0u, 0u, 0u, 0u}; + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + int offset = stage_offset + params.rank * params.rank_stride_packs + idx; + uint4 value = clear_pos_zero_u4_16(input[idx]); + store_u4_volatile(owner_buffer, offset, value); + } + + int start = params.rank * part; + int end = (params.rank == WorldSize - 1) ? params.num_packs : start + part; + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + for (int idx = start + tid; idx < end; idx += stride) { + uint4 values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_u4_volatile(local_buffer, offset); + waiting |= has_pos_zero_u4_16(values[peer]); + } + if (!waiting) { + break; + } + } + uint4 acc = values[0]; +#pragma unroll + for (int peer = 1; peer < WorldSize; ++peer) { + acc = packed_add_u4(acc, values[peer]); + } + reinterpret_cast(params.output)[idx] = acc; + uint4 publish = clear_pos_zero_u4_16(acc); +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + store_u4_volatile(local_buffer, offset, reset); + if (peer != params.rank) { + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(peer_buffer, final_offset, publish); + } + } + } + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + if (owner == params.rank) { + continue; + } + int offset = stage_offset + owner * params.rank_stride_packs + idx; + uint4 value; + while (true) { + value = load_u4_volatile(local_buffer, offset); + if (!has_pos_zero_u4_16(value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = value; + store_u4_volatile(local_buffer, offset, reset); + } + } else { + Pack const* input = reinterpret_cast(params.input); + Pack reset = zero_pack(); + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + int offset = stage_offset + params.rank * params.rank_stride_packs + idx; + Pack value = input[idx]; + clear_pos_zero_pack(value); + store_pack_volatile(owner_buffer, offset, value); + } + + int start = params.rank * part; + int end = (params.rank == WorldSize - 1) ? params.num_packs : start + part; + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + for (int idx = start + tid; idx < end; idx += stride) { + Pack values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_pack_volatile(local_buffer, offset); + waiting |= has_pos_zero_pack(values[peer]); + } + if (!waiting) { + break; + } + } + Pack acc = reduce_loaded_packs(values); + reinterpret_cast(params.output)[idx] = acc; + Pack publish = acc; + clear_pos_zero_pack(publish); +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int offset = stage_offset + peer * params.rank_stride_packs + idx; + store_pack_volatile(local_buffer, offset, reset); + if (peer != params.rank) { + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(peer_buffer, final_offset, publish); + } + } + } + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int owner = rsag_owner_for_pack(idx, part); + if (owner == params.rank) { + continue; + } + int offset = stage_offset + owner * params.rank_stride_packs + idx; + Pack value; + while (true) { + value = load_pack_volatile(local_buffer, offset); + if (!has_pos_zero_pack(value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = value; + store_pack_volatile(local_buffer, offset, reset); + } + } + debug_commit_per_block_epoch(epoch_slot, epoch); +} + +template +__global__ __launch_bounds__(1024, 1) void ipc_topo_rsag8_push_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + pdl_grid_sync_const(); + + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + int epoch = advance_scratch_epoch(params.scratch_state, epoch_slot); + int stage_offset = epoch * params.epoch_stride_packs; + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int stride = gridDim.x * blockDim.x; + int part = params.num_packs / 4; + int base = params.rank < 4 ? 0 : 4; + int cross_base = base ^ 4; + int local_rank = params.rank - base; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + uint4 reset = {0u, 0u, 0u, 0u}; + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int chunk = rsag_owner_for_pack<4>(idx, part); + int owner = base + chunk; + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + int input_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + uint4 local_value = input[idx]; + uint4 publish_input = clear_pos_zero_u4_16(local_value); + store_u4_volatile(owner_buffer, input_offset, publish_input); + + if (local_rank == chunk) { + uint4 values[4]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer_local] = load_u4_volatile(local_buffer, offset); + waiting |= has_pos_zero_u4_16(values[peer_local]); + } + if (!waiting) { + break; + } + } + uint4 local_sum = values[0]; +#pragma unroll + for (int peer_local = 1; peer_local < 4; ++peer_local) { + local_sum = packed_add_u4(local_sum, values[peer_local]); + } +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + int offset = stage_offset + peer * params.rank_stride_packs + idx; + store_u4_volatile(local_buffer, offset, reset); + } + + int cross_owner = cross_base + chunk; + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + int cross_write = stage_offset + params.rank * params.rank_stride_packs + idx; + uint4 publish_sum = clear_pos_zero_u4_16(local_sum); + store_u4_volatile(cross_buffer, cross_write, publish_sum); + + int cross_read = stage_offset + cross_owner * params.rank_stride_packs + idx; + uint4 cross_sum; + while (true) { + cross_sum = load_u4_volatile(local_buffer, cross_read); + if (!has_pos_zero_u4_16(cross_sum)) { + break; + } + } + uint4 final_value = packed_add_u4(local_sum, cross_sum); + reinterpret_cast(params.output)[idx] = final_value; + store_u4_volatile(local_buffer, cross_read, reset); + + uint4 publish_final = clear_pos_zero_u4_16(final_value); +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + if (peer == params.rank) { + continue; + } + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(peer_buffer, final_offset, publish_final); + } + } else { + int final_offset = stage_offset + owner * params.rank_stride_packs + idx; + uint4 final_value; + while (true) { + final_value = load_u4_volatile(local_buffer, final_offset); + if (!has_pos_zero_u4_16(final_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = final_value; + store_u4_volatile(local_buffer, final_offset, reset); + } + } + } else { + Pack const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + Pack reset = zero_pack(); + + for (int idx = tid; idx < params.num_packs; idx += stride) { + int chunk = rsag_owner_for_pack<4>(idx, part); + int owner = base + chunk; + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + int input_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + Pack local_value = input[idx]; + Pack publish_input = local_value; + clear_pos_zero_pack(publish_input); + store_pack_volatile(owner_buffer, input_offset, publish_input); + + if (local_rank == chunk) { + Pack values[4]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + int offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer_local] = load_pack_volatile(local_buffer, offset); + waiting |= has_pos_zero_pack(values[peer_local]); + } + if (!waiting) { + break; + } + } + Pack local_sum = values[0]; +#pragma unroll + for (int peer_local = 1; peer_local < 4; ++peer_local) { + local_sum = add_pack(local_sum, values[peer_local]); + } +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + int offset = stage_offset + peer * params.rank_stride_packs + idx; + store_pack_volatile(local_buffer, offset, reset); + } + + int cross_owner = cross_base + chunk; + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + int cross_write = stage_offset + params.rank * params.rank_stride_packs + idx; + Pack publish_sum = local_sum; + clear_pos_zero_pack(publish_sum); + store_pack_volatile(cross_buffer, cross_write, publish_sum); + + int cross_read = stage_offset + cross_owner * params.rank_stride_packs + idx; + Pack cross_sum; + while (true) { + cross_sum = load_pack_volatile(local_buffer, cross_read); + if (!has_pos_zero_pack(cross_sum)) { + break; + } + } + Pack final_value = add_pack(local_sum, cross_sum); + reinterpret_cast(params.output)[idx] = final_value; + store_pack_volatile(local_buffer, cross_read, reset); + + Pack publish_final = final_value; + clear_pos_zero_pack(publish_final); +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + if (peer == params.rank) { + continue; + } + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(peer_buffer, final_offset, publish_final); + } + } else { + int final_offset = stage_offset + owner * params.rank_stride_packs + idx; + Pack final_value; + while (true) { + final_value = load_pack_volatile(local_buffer, final_offset); + if (!has_pos_zero_pack(final_value)) { + break; + } + } + reinterpret_cast(params.output)[idx] = final_value; + store_pack_volatile(local_buffer, final_offset, reset); + } + } + } + + debug_commit_per_block_epoch(epoch_slot, epoch); + pdl_grid_release_const(); +} + +// TP8 staged topology RS/AG push. +// +// ipc_topo_rsag8_block_param_kernel below partitions blocks by blockIdx.x & 3, +// so its four block groups push to four different island owners at the same +// instant: intra-island traffic is all-to-all, which is the expensive pattern +// on this fabric. +// +// This variant keeps the topology decomposition exactly as it is (island reduce +// -> owner-pair exchange across SYS -> island gather), because a topology-blind +// ring is slower at every block count. It only re-times the two +// intra-island phases: each is split into three passes, and in pass p a rank +// talks solely to island peer (local + 1 + p) % 4, which is a permutation, so +// each GPU has one outbound stream at a time. The cross-island exchange is +// already one-to-one and is left alone. +// +// Because chunks are now visited in time rather than assigned to block groups, +// blocks no longer have to be a multiple of four and a flat grid-stride loop +// covers each chunk. +// +// Costs six extra island barriers, so the policy only selects this above a +// payload threshold. +template +__global__ __launch_bounds__(1024, 1) void ipc_topo_rsag8_ring_push_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + static_assert(kSignalPhases >= 8, "topology ring push needs eight barrier phases"); + pdl_grid_sync_const(); + + int32_t* self_signal = reinterpret_cast(params.signal_ptrs[params.rank]); + // Unsigned arithmetic for the bump: signed overflow is UB, and this counter + // is meant to wrap. generation_reached() reads it back on the circle. + int flag = + static_cast(static_cast(load_acquire_i32( + self_signal + flag_offset(blockIdx.x, params.max_blocks, 8))) + + 1u); + + const int tid = blockIdx.x * blockDim.x + threadIdx.x; + const int stride = gridDim.x * blockDim.x; + const int part = params.num_packs >> 2; + const int base = params.rank < 4 ? 0 : 4; + const int local = params.rank & 3; + const uint32_t island_mask = params.rank < 4 ? 0x0fu : 0xf0u; + const int cross_owner = params.rank ^ 4; + // Call-level double buffer. The cross-island payload this rank publishes + // into its paired owner's slab has no read-complete edge coming back: phase + // 3 proves the write landed, but nothing stops the paired owner's *next* + // call from overwriting it while this one is still reading. Alternating + // halves supplies the missing distance, and two halves are exactly enough -- + // owner_pair_barrier is a two-sided rendezvous, so this rank cannot leave + // call k until its partner has entered call k, hence cannot reach call k+2 + // (the next use of this half) until the partner has left call k. Anyone + // making that barrier one-sided silently breaks this. + int32_t* scratch_state = params.scratch_state; + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + const int call_epoch = advance_scratch_epoch(scratch_state, epoch_slot); + const int stage_offset = call_epoch * params.epoch_stride_packs; + // This rank owns the chunk at its own position in the island. + const int my_start = local * part; + const int my_end = (local == 3) ? params.num_packs : my_start + part; + const int my_slot = stage_offset + params.rank * params.rank_stride_packs; + const int cross_slot = stage_offset + cross_owner * params.rank_stride_packs; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + + // Island reduce-scatter. The contribution to this rank's own chunk stays + // on this GPU, so it costs no fabric time and needs no pass of its own. + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_u4_volatile(local_buffer, my_slot + idx, input[idx]); + } + for (int p = 0; p < 3; ++p) { + const int t = (local + 1 + p) & 3; + const int owner_t = base + t; + const int t_start = t * part; + const int t_end = (t == 3) ? params.num_packs : t_start + part; + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner_t]); + for (int idx = t_start + tid; idx < t_end; idx += stride) { + store_u4_volatile(owner_buffer, my_slot + idx, input[idx]); + } + __threadfence_system(); + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, p, flag, + island_mask); + } + + // Island sum, then the one cross-SYS exchange with the paired owner. + for (int idx = my_start + tid; idx < my_end; idx += stride) { + uint4 v0 = load_u4_volatile(local_buffer, + stage_offset + (base + 0) * params.rank_stride_packs + idx); + uint4 v1 = load_u4_volatile(local_buffer, + stage_offset + (base + 1) * params.rank_stride_packs + idx); + uint4 v2 = load_u4_volatile(local_buffer, + stage_offset + (base + 2) * params.rank_stride_packs + idx); + uint4 v3 = load_u4_volatile(local_buffer, + stage_offset + (base + 3) * params.rank_stride_packs + idx); + uint4 local_sum = packed_add_u4(packed_add_u4(v0, v1), packed_add_u4(v2, v3)); + // Keep the island sum so the gather phase does not recompute it. + store_u4_volatile(local_buffer, my_slot + idx, local_sum); + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + store_u4_volatile(cross_buffer, my_slot + idx, local_sum); + } + __threadfence_system(); + owner_pair_barrier(params.signal_ptrs, params.rank, params.rank, cross_owner, params.max_blocks, + 3, flag); + debug_cross_read_stall(params.rank); + + // Final value for the owned chunk, written locally. + for (int idx = my_start + tid; idx < my_end; idx += stride) { + uint4 mine = load_u4_volatile(local_buffer, my_slot + idx); + uint4 theirs = load_u4_volatile(local_buffer, cross_slot + idx); + uint4 final_value = packed_add_u4(mine, theirs); + reinterpret_cast(params.output)[idx] = final_value; + store_u4_volatile(local_buffer, my_slot + idx, final_value); + } + __threadfence_system(); + + // Island all-gather, staged on the same permutation schedule. + for (int p = 0; p < 3; ++p) { + const int peer = base + ((local + 1 + p) & 3); + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_u4_volatile(peer_buffer, my_slot + idx, + load_u4_volatile(local_buffer, my_slot + idx)); + } + __threadfence_system(); + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, 4 + p, flag, + island_mask); + } + + // Collect the three chunks owned by the other island members. + for (int p = 0; p < 3; ++p) { + const int t = (local + 1 + p) & 3; + const int owner_t = base + t; + const int t_start = t * part; + const int t_end = (t == 3) ? params.num_packs : t_start + part; + const int owner_slot = stage_offset + owner_t * params.rank_stride_packs; + for (int idx = t_start + tid; idx < t_end; idx += stride) { + reinterpret_cast(params.output)[idx] = + load_u4_volatile(local_buffer, owner_slot + idx); + } + } + // Hold the island until everyone has finished reading before the next + // call's reduce-scatter starts writing the same slots. This covers the + // intra-island reuse only; the cross-island edge is what the epoch double + // buffer above supplies. + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, 7, flag, island_mask); + } else { + Pack const* input = reinterpret_cast(params.input); + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_pack_volatile(local_buffer, my_slot + idx, input[idx]); + } + for (int p = 0; p < 3; ++p) { + const int t = (local + 1 + p) & 3; + const int owner_t = base + t; + const int t_start = t * part; + const int t_end = (t == 3) ? params.num_packs : t_start + part; + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner_t]); + for (int idx = t_start + tid; idx < t_end; idx += stride) { + store_pack_volatile(owner_buffer, my_slot + idx, input[idx]); + } + __threadfence_system(); + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, p, flag, + island_mask); + } + + for (int idx = my_start + tid; idx < my_end; idx += stride) { + Pack values[4]; +#pragma unroll + for (int i = 0; i < 4; ++i) { + values[i] = load_pack_volatile( + local_buffer, stage_offset + (base + i) * params.rank_stride_packs + idx); + } + Pack local_sum = reduce_loaded_packs(values); + store_pack_volatile(local_buffer, my_slot + idx, local_sum); + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + store_pack_volatile(cross_buffer, my_slot + idx, local_sum); + } + __threadfence_system(); + owner_pair_barrier(params.signal_ptrs, params.rank, params.rank, cross_owner, params.max_blocks, + 3, flag); + debug_cross_read_stall(params.rank); + + for (int idx = my_start + tid; idx < my_end; idx += stride) { + Pack pair[2]; + pair[0] = load_pack_volatile(local_buffer, my_slot + idx); + pair[1] = load_pack_volatile(local_buffer, cross_slot + idx); + Pack final_value = reduce_loaded_packs(pair); + reinterpret_cast(params.output)[idx] = final_value; + store_pack_volatile(local_buffer, my_slot + idx, final_value); + } + __threadfence_system(); + + for (int p = 0; p < 3; ++p) { + const int peer = base + ((local + 1 + p) & 3); + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + for (int idx = my_start + tid; idx < my_end; idx += stride) { + store_pack_volatile(peer_buffer, my_slot + idx, + load_pack_volatile(local_buffer, my_slot + idx)); + } + __threadfence_system(); + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, 4 + p, flag, + island_mask); + } + + for (int p = 0; p < 3; ++p) { + const int t = (local + 1 + p) & 3; + const int owner_t = base + t; + const int t_start = t * part; + const int t_end = (t == 3) ? params.num_packs : t_start + part; + const int owner_slot = stage_offset + owner_t * params.rank_stride_packs; + for (int idx = t_start + tid; idx < t_end; idx += stride) { + reinterpret_cast(params.output)[idx] = + load_pack_volatile(local_buffer, owner_slot + idx); + } + } + block_barrier_mask(params.signal_ptrs, params.rank, 8, params.max_blocks, 7, flag, island_mask); + } + + if (threadIdx.x == 0) { + store_release_i32(self_signal + flag_offset(blockIdx.x, params.max_blocks, 8), flag); + } + // The barrier flag above must stay ahead of the release: a dependent kernel + // started by the trigger would otherwise observe this call's generation as + // not yet published. (The epoch is no longer a concern here -- it is + // committed at entry.) + debug_commit_per_block_epoch(epoch_slot, call_epoch); + pdl_grid_release_const(); +} + +template +__global__ __launch_bounds__(1024, 1) void ipc_topo_rsag8_block_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + pdl_grid_sync_const(); + + int32_t* self_signal = reinterpret_cast(params.signal_ptrs[params.rank]); + // Unsigned arithmetic for the bump: signed overflow is UB, and this counter + // is meant to wrap. generation_reached() reads it back on the circle. + int flag = + static_cast(static_cast(load_acquire_i32( + self_signal + flag_offset(blockIdx.x, params.max_blocks, 8))) + + 1u); + + // The caller's `blocks % 4 == 0` check is what keeps `blocks_per_chunk` + // non-zero; a zero stride below never advances the grid-stride loops. + int chunk = blockIdx.x & 3; + int chunk_block = blockIdx.x >> 2; + int blocks_per_chunk = gridDim.x >> 2; + int tid = chunk_block * blockDim.x + threadIdx.x; + int stride = blocks_per_chunk * blockDim.x; + int part = params.num_packs >> 2; + int start = chunk * part; + int end = (chunk == 3) ? params.num_packs : start + part; + int base = params.rank < 4 ? 0 : 4; + int owner = base + chunk; + int cross_owner = owner ^ 4; + // Call-level double buffer, for the same reason as the ring kernel: the + // phase 4 ack only covers this island, so nothing orders this rank's cross + // read against the paired owner's next-call cross write. Both TP8 kernels + // share this counter on purpose -- which one runs changes with the payload, + // and a counter advanced by only one of them would let a call land on + // a half two calls old. + int32_t* scratch_state = params.scratch_state; + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + const int call_epoch = advance_scratch_epoch(scratch_state, epoch_slot); + const int stage_offset = call_epoch * params.epoch_stride_packs; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + for (int idx = start + tid; idx < end; idx += stride) { + int offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(owner_buffer, offset, input[idx]); + } + __threadfence_system(); + island_owner_gather(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 1, flag); + + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + if (params.rank == owner) { + for (int idx = start + tid; idx < end; idx += stride) { + uint4 v0 = load_u4_volatile(local_buffer, + stage_offset + (base + 0) * params.rank_stride_packs + idx); + uint4 v1 = load_u4_volatile(local_buffer, + stage_offset + (base + 1) * params.rank_stride_packs + idx); + uint4 v2 = load_u4_volatile(local_buffer, + stage_offset + (base + 2) * params.rank_stride_packs + idx); + uint4 v3 = load_u4_volatile(local_buffer, + stage_offset + (base + 3) * params.rank_stride_packs + idx); + uint4 local_sum = packed_add_u4(packed_add_u4(v0, v1), packed_add_u4(v2, v3)); + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + int cross_write = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(cross_buffer, cross_write, local_sum); + } + } + if (params.rank == owner) { + __threadfence_system(); + } + owner_pair_barrier(params.signal_ptrs, params.rank, owner, cross_owner, params.max_blocks, 2, + flag); + debug_cross_read_stall(params.rank); + + if (params.rank == owner) { + for (int idx = start + tid; idx < end; idx += stride) { + uint4 v0 = load_u4_volatile(local_buffer, + stage_offset + (base + 0) * params.rank_stride_packs + idx); + uint4 v1 = load_u4_volatile(local_buffer, + stage_offset + (base + 1) * params.rank_stride_packs + idx); + uint4 v2 = load_u4_volatile(local_buffer, + stage_offset + (base + 2) * params.rank_stride_packs + idx); + uint4 v3 = load_u4_volatile(local_buffer, + stage_offset + (base + 3) * params.rank_stride_packs + idx); + uint4 local_sum = packed_add_u4(packed_add_u4(v0, v1), packed_add_u4(v2, v3)); + uint4 cross_sum = load_u4_volatile( + local_buffer, stage_offset + cross_owner * params.rank_stride_packs + idx); + uint4 final_value = packed_add_u4(local_sum, cross_sum); + reinterpret_cast(params.output)[idx] = final_value; +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + if (peer == params.rank) { + continue; + } + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(peer_buffer, final_offset, final_value); + } + } + } + if (params.rank == owner) { + __threadfence_system(); + } + island_owner_ready(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 3, flag); + + if (params.rank != owner) { + for (int idx = start + tid; idx < end; idx += stride) { + uint4 final_value = + load_u4_volatile(local_buffer, stage_offset + owner * params.rank_stride_packs + idx); + reinterpret_cast(params.output)[idx] = final_value; + } + } + // WRONG ORDER, and the reason enable_pdl is still refused at the binding: + // the trigger fires before island_owner_ack and before the barrier flag + // below, so a dependent kernel can start while this call's phase-4 ack and + // flag are still being written. Moving the release past both is the fix, + // but re-enabling PDL needs a per-kernel audit and an SM90 regression, so + // it is left visible rather than quietly reordered. + pdl_grid_release_const(); + island_owner_ack(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 4, flag); + } else { + Pack const* input = reinterpret_cast(params.input); + auto* owner_buffer = reinterpret_cast(params.tmp_ptrs[owner]); + for (int idx = start + tid; idx < end; idx += stride) { + int offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(owner_buffer, offset, input[idx]); + } + __threadfence_system(); + island_owner_gather(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 1, flag); + + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + if (params.rank == owner) { + for (int idx = start + tid; idx < end; idx += stride) { + Pack v0 = load_pack_volatile(local_buffer, + stage_offset + (base + 0) * params.rank_stride_packs + idx); + Pack v1 = load_pack_volatile(local_buffer, + stage_offset + (base + 1) * params.rank_stride_packs + idx); + Pack v2 = load_pack_volatile(local_buffer, + stage_offset + (base + 2) * params.rank_stride_packs + idx); + Pack v3 = load_pack_volatile(local_buffer, + stage_offset + (base + 3) * params.rank_stride_packs + idx); + Pack local_sum = add_pack(add_pack(v0, v1), add_pack(v2, v3)); + auto* cross_buffer = reinterpret_cast(params.tmp_ptrs[cross_owner]); + int cross_write = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(cross_buffer, cross_write, local_sum); + } + } + if (params.rank == owner) { + __threadfence_system(); + } + owner_pair_barrier(params.signal_ptrs, params.rank, owner, cross_owner, params.max_blocks, 2, + flag); + debug_cross_read_stall(params.rank); + + if (params.rank == owner) { + for (int idx = start + tid; idx < end; idx += stride) { + Pack v0 = load_pack_volatile(local_buffer, + stage_offset + (base + 0) * params.rank_stride_packs + idx); + Pack v1 = load_pack_volatile(local_buffer, + stage_offset + (base + 1) * params.rank_stride_packs + idx); + Pack v2 = load_pack_volatile(local_buffer, + stage_offset + (base + 2) * params.rank_stride_packs + idx); + Pack v3 = load_pack_volatile(local_buffer, + stage_offset + (base + 3) * params.rank_stride_packs + idx); + Pack local_sum = add_pack(add_pack(v0, v1), add_pack(v2, v3)); + Pack cross_sum = load_pack_volatile( + local_buffer, stage_offset + cross_owner * params.rank_stride_packs + idx); + Pack final_value = add_pack(local_sum, cross_sum); + reinterpret_cast(params.output)[idx] = final_value; +#pragma unroll + for (int peer_local = 0; peer_local < 4; ++peer_local) { + int peer = base + peer_local; + if (peer == params.rank) { + continue; + } + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int final_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(peer_buffer, final_offset, final_value); + } + } + } + if (params.rank == owner) { + __threadfence_system(); + } + island_owner_ready(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 3, flag); + + if (params.rank != owner) { + for (int idx = start + tid; idx < end; idx += stride) { + Pack final_value = load_pack_volatile( + local_buffer, stage_offset + owner * params.rank_stride_packs + idx); + reinterpret_cast(params.output)[idx] = final_value; + } + } + // WRONG ORDER, and the reason enable_pdl is still refused at the binding: + // the trigger fires before island_owner_ack and before the barrier flag + // below, so a dependent kernel can start while this call's phase-4 ack and + // flag are still being written. Moving the release past both is the fix, + // but re-enabling PDL needs a per-kernel audit and an SM90 regression, so + // it is left visible rather than quietly reordered. + pdl_grid_release_const(); + island_owner_ack(params.signal_ptrs, params.rank, base, owner, params.max_blocks, 4, flag); + } + + if (threadIdx.x == 0) { + store_release_i32(self_signal + flag_offset(blockIdx.x, params.max_blocks, 8), flag); + } + debug_commit_per_block_epoch(epoch_slot, call_epoch); +} + +template +__global__ __launch_bounds__(1024, 1) void push_oneshot_param_kernel( + const PushOneshotParamData __grid_constant__ params) { + using Pack = typename PackTraits::Pack; + pdl_grid_sync_const(); + + int32_t* epoch_slot = + params.epoch_slots + blockIdx.x; // used only by the PER_BLOCK_EPOCH debug build + int epoch = advance_scratch_epoch(params.scratch_state, epoch_slot); + int stage_offset = epoch * params.epoch_stride_packs; + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int stride = gridDim.x * blockDim.x; + + if constexpr (std::is_same_v || std::is_same_v) { + uint4 const* input = reinterpret_cast(params.input); + for (int idx = tid; idx < params.num_packs; idx += stride) { + uint4 value = clear_pos_zero_u4_16(input[idx]); +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + auto* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int peer_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_u4_volatile(peer_buffer, peer_offset, value); + } + } + + auto* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + uint4 reset = {0u, 0u, 0u, 0u}; + for (int idx = tid; idx < params.num_packs; idx += stride) { + uint4 values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int peer_offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_u4_volatile(local_buffer, peer_offset); + waiting |= has_pos_zero_u4_16(values[peer]); + } + if (!waiting) { + break; + } + } + + uint4 acc; + if constexpr (Fp32Reduce) { + acc = reduce_u4_fp32(values); + } else { + acc = values[0]; +#pragma unroll + for (int peer = 1; peer < WorldSize; ++peer) { + acc = packed_add_u4(acc, values[peer]); + } + } + reinterpret_cast(params.output)[idx] = acc; + +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int peer_offset = stage_offset + peer * params.rank_stride_packs + idx; + local_buffer[peer_offset] = reset; + } + } + } else { + Pack const* input = reinterpret_cast(params.input); + for (int idx = tid; idx < params.num_packs; idx += stride) { + Pack value = input[idx]; + clear_pos_zero_pack(value); +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + Pack* peer_buffer = reinterpret_cast(params.tmp_ptrs[peer]); + int peer_offset = stage_offset + params.rank * params.rank_stride_packs + idx; + store_pack_volatile(peer_buffer, peer_offset, value); + } + } + + Pack* local_buffer = reinterpret_cast(params.tmp_ptrs[params.rank]); + Pack reset = zero_pack(); + for (int idx = tid; idx < params.num_packs; idx += stride) { + Pack values[WorldSize]; + while (true) { + bool waiting = false; +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int peer_offset = stage_offset + peer * params.rank_stride_packs + idx; + values[peer] = load_pack_volatile(local_buffer, peer_offset); + waiting |= has_pos_zero_pack(values[peer]); + } + if (!waiting) { + break; + } + } + + Pack acc = reduce_loaded_packs(values); + reinterpret_cast(params.output)[idx] = acc; + +#pragma unroll + for (int peer = 0; peer < WorldSize; ++peer) { + int peer_offset = stage_offset + peer * params.rank_stride_packs + idx; + local_buffer[peer_offset] = reset; + } + } + } + + debug_commit_per_block_epoch(epoch_slot, epoch); + pdl_grid_release_const(); +} +// --------------------------------------------------------------------------- +// Host side +// --------------------------------------------------------------------------- + +// Byte layout of one rank's workspace slab. +// +// [ epoch slots | barrier phase slots | barrier flags +// | block-scratch epoch + arrival | pack scratch | block scratch ] +// +// Both scratch regions are sized for world_size ranks x a double-buffered +// epoch, so a rank may start collective N+1 before its peer has drained N. +// The epoch halves sit at fixed offsets derived from max_numel rather than +// from the current payload: if they moved with the payload, a rank that +// finished a large collective and flipped its epoch would start writing a +// small one inside the region a lagging peer is still draining. +struct WorkspaceLayout { + size_t signal_bytes; + size_t max_payload_bytes; + size_t scratch_bytes; // per scratch region + size_t total_bytes; +}; + +inline WorkspaceLayout compute_workspace_layout(int world_size, int64_t max_numel, int elem_size, + int max_blocks) { + const size_t epoch_slots = static_cast(max_blocks); + const size_t barrier_slots = static_cast(kSignalPhases) * + static_cast(max_blocks) * static_cast(world_size); + const size_t flag_slots = static_cast(max_blocks); + // {epoch, arrival} per scratch region, in ScratchRegion order. Appended at + // the tail so phase_offset() and flag_offset(), both anchored at the front, + // are unchanged. See scratch_state_offset(). + const size_t scratch_state_slots = 2 * 2; + const size_t signal_slots = epoch_slots + barrier_slots + flag_slots + scratch_state_slots; + auto align128 = [](size_t n) { return (n + 127u) & ~static_cast(127u); }; + WorkspaceLayout layout{}; + layout.signal_bytes = align128(sizeof(int32_t) * signal_slots); + layout.max_payload_bytes = + align128(static_cast(max_numel) * static_cast(elem_size)); + layout.scratch_bytes = align128(2 * static_cast(world_size) * layout.max_payload_bytes); + layout.total_bytes = layout.signal_bytes + 2 * layout.scratch_bytes; + return layout; +} + +// Bytes each rank must allocate and share over CUDA IPC. +inline int64_t workspace_size(int world_size, int64_t max_numel, int elem_size, int max_blocks) { + return static_cast( + compute_workspace_layout(world_size, max_numel, elem_size, max_blocks).total_bytes); +} + +// Per-region device pointers into every rank's slab, as seen by this process. +struct PeerViews { + uint64_t signal[kMaxWorldSize]; + uint64_t pack[kMaxWorldSize]; + uint64_t block[kMaxWorldSize]; + int32_t* self_signal; +}; + +// ipc_ptrs[i] must address rank i's slab; ipc_ptrs[rank] is this rank's own. +inline PeerViews make_peer_views(const int64_t* ipc_ptrs, int world_size, int rank, + const WorkspaceLayout& layout) { + PeerViews views{}; + for (int peer = 0; peer < world_size; ++peer) { + auto* base = reinterpret_cast(ipc_ptrs[peer]); + views.signal[peer] = reinterpret_cast(base); + auto* scratch = base + static_cast(layout.signal_bytes); + views.pack[peer] = reinterpret_cast(scratch); + views.block[peer] = reinterpret_cast(scratch + layout.scratch_bytes); + } + views.self_signal = reinterpret_cast(ipc_ptrs[rank]); + return views; +} + +template +inline cudaError_t launch(Kernel kernel, dim3 grid, dim3 block, cudaStream_t stream, bool use_pdl, + Args const&... args) { +#if CUDART_VERSION >= 12000 + if (use_pdl) { + cudaLaunchAttribute attr[1]; + attr[0].id = cudaLaunchAttributeProgrammaticStreamSerialization; + attr[0].val.programmaticStreamSerializationAllowed = 1; + cudaLaunchConfig_t config{}; + config.gridDim = grid; + config.blockDim = block; + config.dynamicSmemBytes = 0; + config.stream = stream; + config.attrs = attr; + config.numAttrs = 1; + return cudaLaunchKernelEx(&config, kernel, args...); + } +#else + if (use_pdl) return cudaErrorNotSupported; +#endif + kernel<<>>(args...); + return cudaGetLastError(); +} + +// Kernel selection. `variant` is picked by the caller rather than by a +// threshold here, because the crossovers depend on the fabric and are measured +// per machine. +// +// world variant kernel +// 2 kUnstaged ipc_tp2_remote_push_kernel (block scratch) +// 2 kStaged ipc_tp2_remote_push_kernel (block scratch) +// 4 kUnstaged push_oneshot_param_kernel (pack scratch) +// 4 kStaged ipc_rsag_push_param_kernel<4> (block scratch) +// 4 kStagedRing ipc_rsag_ring_push_param_kernel<4> (block scratch) +// 8 kUnstaged ipc_topo_rsag8_push_param_kernel (pack scratch) +// 8 kStaged ipc_topo_rsag8_block_param_kernel (block scratch) +// 8 kStagedRing ipc_topo_rsag8_ring_push_param_kernel (block scratch) +// 8 kFlatStaged ipc_rsag_push_param_kernel<8> (pack scratch) +// +// Preconditions the caller must have validated: world_size in {2,4,8}; the +// (world_size, variant) pair appears above; 0 < blocks <= max_blocks; +// 0 < threads <= 1024; numel and max_numel both divisible by the 16-byte pack +// width; numel * elem_size <= max_payload_bytes; and blocks % 4 == 0 for +// (8, kStaged), since that kernel derives its chunk from blockIdx.x & 3. +template +cudaError_t all_reduce(const T* input, T* output, int64_t numel, const PeerViews& views, int rank, + int world_size, int max_blocks, int64_t max_numel, int blocks, int threads, + Variant variant, bool use_pdl, cudaStream_t stream) { + using Traits = PackTraits; + const int num_packs = static_cast(numel / Traits::kPackElems); + const int rank_stride_packs = static_cast(max_numel / Traits::kPackElems); + const dim3 grid(static_cast(blocks)); + const dim3 cta(static_cast(threads)); + + if (world_size == 2) { + IpcTp2RemotePushData params{}; + params.tmp_ptrs[0] = views.block[0]; + params.tmp_ptrs[1] = views.block[1]; + params.input = input; + params.output = output; + params.epoch_slots = views.self_signal; + // TP2 stages through views.block under either variant, so it shares the + // block region's counter -- nominal here, since it is the only TP2 kernel, + // but the state must follow the region it actually writes. + params.scratch_state = + views.self_signal + scratch_state_offset(max_blocks, 2, ScratchRegion::kBlock); + params.num_packs = num_packs; + params.rank_stride_packs = rank_stride_packs; + params.rank = rank; + const bool staged = variant == Variant::kStaged; + if (use_pdl) { + return staged ? launch(ipc_tp2_remote_push_kernel, grid, cta, stream, true, + params) + : launch(ipc_tp2_remote_push_kernel, grid, cta, stream, true, + params); + } + return staged ? launch(ipc_tp2_remote_push_kernel, grid, cta, stream, false, + params) + : launch(ipc_tp2_remote_push_kernel, grid, cta, stream, false, + params); + } + + PushOneshotParamData params{}; + // Region and counter are chosen together: a kernel reading one region while + // advancing another's epoch would corrupt both. Which kernels may share a + // region is a protocol question, not a partitioning one -- see ScratchRegion. + const ScratchRegion region = (variant == Variant::kUnstaged || variant == Variant::kFlatStaged) + ? ScratchRegion::kPack + : ScratchRegion::kBlock; + const uint64_t* scratch = region == ScratchRegion::kPack ? views.pack : views.block; + for (int peer = 0; peer < world_size; ++peer) { + params.tmp_ptrs[peer] = scratch[peer]; + params.signal_ptrs[peer] = views.signal[peer]; + } + params.input = input; + params.output = output; + params.epoch_slots = views.self_signal; + params.scratch_state = views.self_signal + scratch_state_offset(max_blocks, world_size, region); + params.num_packs = num_packs; + params.rank_stride_packs = rank_stride_packs; + params.epoch_stride_packs = world_size * rank_stride_packs; + params.rank = rank; + params.max_blocks = max_blocks; + +#define FI_PCIE_IPC_LAUNCH(KERNEL_EXPR, PDL) launch(KERNEL_EXPR, grid, cta, stream, PDL, params) + +#define FI_PCIE_IPC_SELECT(PDL) \ + do { \ + if (world_size == 8) { \ + switch (variant) { \ + case Variant::kUnstaged: \ + return FI_PCIE_IPC_LAUNCH((ipc_topo_rsag8_push_param_kernel), PDL); \ + case Variant::kStaged: \ + return FI_PCIE_IPC_LAUNCH((ipc_topo_rsag8_block_param_kernel), PDL); \ + case Variant::kStagedRing: \ + return FI_PCIE_IPC_LAUNCH((ipc_topo_rsag8_ring_push_param_kernel), PDL); \ + case Variant::kFlatStaged: \ + return FI_PCIE_IPC_LAUNCH((ipc_rsag_push_param_kernel), PDL); \ + } \ + return cudaErrorInvalidValue; \ + } \ + switch (variant) { \ + case Variant::kUnstaged: \ + return FI_PCIE_IPC_LAUNCH((push_oneshot_param_kernel), PDL); \ + case Variant::kStaged: \ + return FI_PCIE_IPC_LAUNCH((ipc_rsag_push_param_kernel), PDL); \ + case Variant::kStagedRing: \ + return FI_PCIE_IPC_LAUNCH((ipc_rsag_ring_push_param_kernel), PDL); \ + default: \ + return cudaErrorInvalidValue; \ + } \ + } while (false) + + if (use_pdl) { + FI_PCIE_IPC_SELECT(true); + } + FI_PCIE_IPC_SELECT(false); + +#undef FI_PCIE_IPC_SELECT +#undef FI_PCIE_IPC_LAUNCH +} + +} // namespace pcie_ipc +} // namespace comm +} // namespace flashinfer + +#endif // FLASHINFER_COMM_PCIE_IPC_ALL_REDUCE_CUH_ diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/jit/comm.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/jit/comm.py new file mode 100644 index 0000000..17cdadb --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/fi/jit/comm.py @@ -0,0 +1,307 @@ +""" +Copyright (c) 2025 by FlashInfer team. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +from .core import JitSpec, gen_jit_spec, current_compilation_context +from .utils import write_if_different +from . import env as jit_env +import os +import pathlib +import jinja2 +from itertools import product +from typing import Dict, Tuple, List, Any + + +def gen_comm_alltoall_module() -> JitSpec: + return gen_jit_spec( + "comm", + [ + jit_env.FLASHINFER_CSRC_DIR / "trtllm_alltoall.cu", + jit_env.FLASHINFER_CSRC_DIR / "trtllm_alltoall_prepare.cu", + ], + ) + + +def gen_trtllm_mnnvl_comm_module() -> JitSpec: + # This MNNVL path is supported on Hopper/Blackwell datacenter targets only. + # Thor (SM11x) and RTX/consumer (SM12x) are intentionally excluded here. + nvcc_flags = current_compilation_context.get_nvcc_flags_list( + supported_major_versions=[9, 10] + ) + # MNNVL allreduce is fully deterministic across ranks. Oneshot specializes + # the common TP<=8 cases to keep the local value in registers and + # volatile-load only peers from the Lamport buffer. Larger world sizes use a + # compact deterministic fallback: the runtime benefit is thin there, while + # rank-specializing every case significantly increases JIT compile time. + return gen_jit_spec( + "trtllm_mnnvl_comm", + [ + jit_env.FLASHINFER_CSRC_DIR / "trtllm_mnnvl_allreduce.cu", + ], + extra_cuda_cflags=nvcc_flags, + ) + + +def gen_mixed_comm_module() -> JitSpec: + gen_directory = jit_env.FLASHINFER_GEN_SRC_DIR / "gen_mixed_comm" + os.makedirs(gen_directory, exist_ok=True) + source_paths = [jit_env.FLASHINFER_CSRC_DIR / "mixed_comm.cu"] + + def is_valid_op( + op: str, + use_local_tp: bool, + use_local_dp: bool, + use_inter_tp: bool, + use_inter_dp: bool, + ) -> bool: + # Should be aligned with is_valid_op in mixed_comm.cuh + use_tp = use_local_tp or use_inter_tp + use_dp = use_local_dp or use_inter_dp + use_mixed = use_tp and use_dp + if use_mixed: + return op in ("ALLREDUCE_ALLGATHER", "REDUCESCATTER_ALLREDUCE") + if use_tp: + return op == "ALLREDUCE" + return op in ("ALLGATHER", "REDUCESCATTER") + + def is_valid_mode(mode: str, use_local_tp: bool, use_inter_tp: bool) -> bool: + # Should be aligned with is_valid_mode in mixed_comm.cuh + if mode.startswith("OPT_WAITS_"): + return True + if mode.startswith("OPT_BYTES1_"): + return use_local_tp + if mode.startswith("OPT_BYTES2_"): + return use_inter_tp + return False + + def is_valid_block_y( + block_size_y: int, local_tp_size: int, local_dp_size: int, mode: str + ) -> bool: + # Should be aligned with is_valid_block_y in mixed_comm.cuh + if not mode.endswith("_MC"): + return block_size_y == 1 + if mode.startswith("OPT_WAITS_"): + return local_dp_size % block_size_y == 0 + return (local_tp_size * local_dp_size) % block_size_y == 0 + + with open(jit_env.FLASHINFER_CSRC_DIR / "mixed_comm_kernel_inst.jinja") as f: + kernel_inst_templ = jinja2.Template(f.read()) + kernel_dict = { + "allreduce_kernel": "ALLREDUCE", + "allgather_kernel": "ALLGATHER", + "reducescatter_kernel": "REDUCESCATTER", + "fused_allreduce_allgather_kernel": "ALLREDUCE_ALLGATHER", + "fused_reducescatter_allreduce_kernel": "REDUCESCATTER_ALLREDUCE", + } + mode_list = [ + "OPT_WAITS_MC", + "OPT_WAITS_UC", + "OPT_BYTES1_MC", + "OPT_BYTES1_UC", + "OPT_BYTES2_MC", + "OPT_BYTES2_UC", + ] + dtype_list = ["nv_half", "nv_bfloat16"] + local_size_list = [1, 2, 4, 8] + inter_flag_list = [False, True] + compile_dict: Dict[Tuple[str, str], List[Dict[str, Any]]] = { + (kernel_name, mode): [] + for kernel_name, mode in product(kernel_dict.keys(), mode_list) + } + for kernel_name, op in kernel_dict.items(): + for ( + block_size_y, + local_tp_size, + local_dp_size, + use_inter_tp, + use_inter_dp, + mode, + dtype, + ) in product( + local_size_list, + local_size_list, + local_size_list, + inter_flag_list, + inter_flag_list, + mode_list, + dtype_list, + ): + use_local_tp = local_tp_size > 1 + use_local_dp = local_dp_size > 1 + if not use_local_tp and not use_local_dp: + continue + if not is_valid_op( + op, use_local_tp, use_local_dp, use_inter_tp, use_inter_dp + ): + continue + if not is_valid_mode(mode, use_local_tp, use_inter_tp): + continue + if not is_valid_block_y(block_size_y, local_tp_size, local_dp_size, mode): + continue + compile_dict[(kernel_name, mode)].append( + { + "kernel_name": kernel_name, + "block_size_y": block_size_y, + "local_tp_size": local_tp_size, + "local_dp_size": local_dp_size, + "use_inter_tp": str(use_inter_tp).lower(), + "use_inter_dp": str(use_inter_dp).lower(), + "mode": mode, + "dtype": dtype, + } + ) + for (kernel_name, mode), instantiations in compile_dict.items(): + dest_path = gen_directory / f"mixed_comm_{kernel_name}_{mode}_inst.cu" + source_paths.append(dest_path) + source = kernel_inst_templ.render(instantiations=instantiations) + write_if_different(dest_path, source) + + import nvidia.nvshmem + + path_base = pathlib.Path(nvidia.nvshmem.__path__[0]) + + nvcc_flags = current_compilation_context.get_nvcc_flags_list( + supported_major_versions=[9, 10] + ) + ["-rdc=true"] + ldflags = [f"-L{str(path_base / 'lib')}"] + ["-lnvshmem_device"] + + return gen_jit_spec( + "mixed_comm", + source_paths, + extra_include_paths=[str(path_base / "include")], + extra_cuda_cflags=nvcc_flags, + extra_ldflags=ldflags, + needs_device_linking=True, + ) + + +def gen_trtllm_comm_module() -> JitSpec: + nvcc_flags = current_compilation_context.get_nvcc_flags_list( + supported_major_versions=[9, 10, 12] + ) + return gen_jit_spec( + "trtllm_comm", + [ + jit_env.FLASHINFER_CSRC_DIR / "trtllm_allreduce.cu", + jit_env.FLASHINFER_CSRC_DIR / "trtllm_allreduce_fusion.cu", + jit_env.FLASHINFER_CSRC_DIR / "trtllm_moe_allreduce_fusion.cu", + ], + extra_cuda_cflags=nvcc_flags, + ) + + +def gen_vllm_comm_module() -> JitSpec: + return gen_jit_spec( + "vllm_comm", + [ + jit_env.FLASHINFER_CSRC_DIR / "vllm_custom_all_reduce.cu", + ], + ) + + +def gen_ulysses_a2a_module() -> JitSpec: + return gen_jit_spec( + "ulysses_a2a", + [ + jit_env.FLASHINFER_CSRC_DIR / "ulysses_all_to_all.cu", + ], + ) + + +def gen_moe_alltoall_module() -> JitSpec: + return gen_jit_spec( + "mnnvl_moe_alltoall", + [ + jit_env.FLASHINFER_CSRC_DIR / "trtllm_moe_alltoall.cu", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "tensorrt_llm" + / "kernels" + / "communicationKernels" + / "moeAlltoAllKernels.cu", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "envUtils.cpp", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "tllmException.cpp", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "logger.cpp", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "stringUtils.cpp", + ], + extra_include_paths=[ + str(jit_env.FLASHINFER_CSRC_DIR / "nv_internal"), + str(jit_env.FLASHINFER_CSRC_DIR / "nv_internal" / "include"), + ], + extra_cuda_cflags=[ + "-DENABLE_BF16", + ], + ) + + +def gen_dcp_alltoall_module() -> JitSpec: + nvcc_flags = current_compilation_context.get_nvcc_flags_list( + supported_major_versions=[9, 10, 11, 12] + ) + return gen_jit_spec( + "dcp_alltoall", + [ + jit_env.FLASHINFER_CSRC_DIR / "trtllm_dcp_alltoall.cu", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "tensorrt_llm" + / "kernels" + / "helixAllToAll.cu", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "envUtils.cpp", + jit_env.FLASHINFER_CSRC_DIR + / "nv_internal" + / "cpp" + / "common" + / "tllmException.cpp", + ], + extra_include_paths=[ + str(jit_env.FLASHINFER_CSRC_DIR / "nv_internal"), + str(jit_env.FLASHINFER_CSRC_DIR / "nv_internal" / "include"), + ], + extra_cuda_cflags=nvcc_flags, + ) + + +def gen_pcie_ipc_comm_module() -> JitSpec: + # SSKJ-PIE (vendored from flashinfer main): no architecture restriction -- + # the kernels use only plain PTX loads/stores and CUDA IPC, both of which + # predate every architecture flashinfer builds for. The target is a PCIe + # machine without NVLink, which is orthogonal to the SM version. + return gen_jit_spec( + "pcie_ipc_comm", + [ + jit_env.FLASHINFER_CSRC_DIR / "pcie_ipc_all_reduce.cu", + ], + ) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/manifest.json b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/manifest.json new file mode 100644 index 0000000..0da881a --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/manifest.json @@ -0,0 +1,14 @@ +{ + "fi/comm/__init__.py": "83cac92703f8c49ef11a72861738f1e65d3fede45dfc4b099a9db34f3e458687", + "fi/comm/pcie_ipc_ar.py": "11c8826678d49a4f95074ebc60f069124f9f9618774255c54b03fb0e616520c7", + "fi/comm/pcie_ipc_policy.py": "d03c82bbca75a9d8df28f0aba5c863a8f9332c6faa22175a73d8b2f28678f3bb", + "fi/comm/pcie_ipc_topology.py": "26f473a1acdce258f1b8ca4240f6c4c028554b257d4b16a5c44902505546edb2", + "fi/comm/pcie_ipc_tuning.py": "d20d088fc287bb36a865d66e682dc696c2196d54be3e6dddad64ae4765133df8", + "fi/data/csrc/pcie_ipc_all_reduce.cu": "b2646bcef5e3e11dc9a5243cd7cd750cf75d54a8c199457f01fdcc56e28d817a", + "fi/data/include/flashinfer/comm/pcie_ipc_all_reduce.cuh": "f03a372c53c20c3176cddd0f90db6242222a01ebebcbc83af96b169c37bf57a1", + "fi/jit/comm.py": "b36c88d50ccfcfc529c45e93243b2123654405a43a40d0823363f14c6ff011c0", + "sg/base_runner.py": "d2b83190072f70aab725910368ffa69b9974e9c68af1c9a413d4038e32dec3fb", + "sg/environ.py": "45c5c4f586aa1275666ffef620681ae3a0a407c82e4e1a72bb3d13f1030aa969", + "sg/parallel_state.py": "9045173c2cf2b6ed836b70a9c13fbe78c3da1911fca651c445fdee487f4a8cc8", + "sg/pcie_ipc_ar.py": "953e0cc8b1f0383e6476d200835a4692991096a659a435f2ce5e89ecf5415613" +} \ No newline at end of file diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/base_runner.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/base_runner.py new file mode 100644 index 0000000..ef75ffe --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/base_runner.py @@ -0,0 +1,717 @@ +# Copyright 2023-2026 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Base class shared by EagerRunner and BaseCudaGraphRunner.""" + +from __future__ import annotations + +import inspect +import logging +from abc import ABC, abstractmethod +from types import SimpleNamespace +from typing import TYPE_CHECKING, Any, Optional, Tuple + +import torch + +from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin +from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config +from sglang.srt.environ import envs +from sglang.srt.layers import deep_gemm_wrapper +from sglang.srt.layers.dp_attention import ( + DpPaddingMode, + set_dp_buffer_len, + set_is_extend_in_batch, +) +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardBatch, + ForwardMode, + NgramEmbeddingInfo, + PPProxyTensors, + get_server_return_hidden_states_mode, +) +from sglang.srt.model_executor.forward_context import ForwardContext, forward_context +from sglang.srt.model_executor.runner.flashinfer_autotune import ( + maybe_flashinfer_autotune_extend, + run_flashinfer_autotune_forward, + should_run_flashinfer_autotune, +) +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_flags, + get_parallel, +) +from sglang.srt.speculative.spec_info import create_dummy_verify_input +from sglang.srt.utils import ( + empty_context, + log_info_on_rank0, + require_attn_tp_gather, + require_gathered_buffer, + require_mlp_tp_gather, +) + +if TYPE_CHECKING: + from sglang.srt.model_executor.model_runner import ModelRunner + +logger = logging.getLogger(__name__) + + +def _allocate_decode_buffers( + *, + device: torch.device, + max_bs: int, + max_num_token: int, + hidden_size: int, + vocab_size: int, + dtype: torch.dtype, + dp_size: int, + pp_size: int, + is_encoder_decoder: bool, + require_mlp_tp_gather: bool, + seq_len_fill_value: int, + encoder_len_fill_value: int, + num_tokens_per_req: int, + cache_loc_dtype: torch.dtype, + enable_mamba_track: bool, + ne_token_table: Optional[torch.Tensor] = None, + hc_hidden_size: Optional[int] = None, + pp_proxy_topk_size: Optional[int] = None, + pp_proxy_residual_num_blocks: Optional[int] = None, + allocate_logits_buffer: bool = True, +) -> SimpleNamespace: + """Allocate the FB-shared decode buffers.""" + with torch.device(device): + input_ids = torch.zeros((max_num_token,), dtype=torch.int64) + input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype) + req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64) + seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64) + out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype) + positions = torch.zeros((max_num_token,), dtype=torch.int64) + mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64) + num_token_non_padded = torch.zeros((1,), dtype=torch.int32) + custom_mask = torch.ones( + (max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_req, + dtype=torch.bool, + ) + # (max_num_token, vocab) fp32 is large (>10GB at 16k tokens); callers + # whose dummy runs never touch logits (run_lm_head=False autotune) opt out. + next_token_logits_buffer = ( + torch.zeros( + (max_num_token, vocab_size), + dtype=torch.float, + ) + if allocate_logits_buffer + else None + ) + mamba_track_indices = ( + torch.zeros((max_bs,), dtype=torch.int64) if enable_mamba_track else None + ) + mamba_track_mask = ( + torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None + ) + + if pp_size > 1: + # mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size). + is_mhc = hc_hidden_size is not None + hs = hc_hidden_size if is_mhc else hidden_size + pp_proxy_tensors = { + "hidden_states": torch.zeros((max_num_token, hs), dtype=dtype), + } + if not is_mhc: + # Only Kimi K3 supplies num_blocks: its PP bank is token-major + # [T, blocks, H]. Other models keep the legacy [max_bs, H]. + residual_shape = ( + (max_num_token, pp_proxy_residual_num_blocks, hidden_size) + if pp_proxy_residual_num_blocks is not None + else (max_num_token, hidden_size) + ) + pp_proxy_tensors["residual"] = torch.zeros(residual_shape, dtype=dtype) + if pp_proxy_topk_size is not None: + pp_proxy_tensors["topk_indices"] = torch.zeros( + (max_num_token, pp_proxy_topk_size), dtype=torch.int32 + ) + else: + pp_proxy_tensors = None + + if is_encoder_decoder: + encoder_lens = torch.full( + (max_bs,), encoder_len_fill_value, dtype=torch.int32 + ) + else: + encoder_lens = None + + if require_mlp_tp_gather: + global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32) + global_num_tokens_for_logprob_gpu = torch.zeros( + (dp_size,), dtype=torch.int32 + ) + else: + global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32) + global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32) + + ngram_embedding_info = ( + NgramEmbeddingInfo( + token_table=ne_token_table, + column_starts=torch.zeros([max_bs], dtype=torch.int32), + req_lens=torch.ones([max_bs], dtype=torch.int32), + out_column_starts=torch.zeros([max_bs], dtype=torch.int32), + out_req_lens=torch.ones([max_bs], dtype=torch.int32), + skip_token_table_update=torch.zeros([max_bs], dtype=torch.bool), + ) + if ne_token_table is not None + else None + ) + + if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get(): + rids_int = torch.zeros((max_bs,), dtype=torch.int64) + bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64) + else: + rids_int = None + bootstrap_room_ids_int = None + + seq_lens_cpu = torch.full( + (max_bs,), + seq_len_fill_value, + dtype=torch.int64, + device="cpu", + ) + + return SimpleNamespace( + input_ids=input_ids, + input_embeds=input_embeds, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + out_cache_loc=out_cache_loc, + positions=positions, + mrope_positions=mrope_positions, + num_token_non_padded=num_token_non_padded, + custom_mask=custom_mask, + next_token_logits_buffer=next_token_logits_buffer, + mamba_track_indices=mamba_track_indices, + mamba_track_mask=mamba_track_mask, + encoder_lens=encoder_lens, + global_num_tokens_gpu=global_num_tokens_gpu, + global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu, + pp_proxy_tensors=pp_proxy_tensors, + ngram_embedding_info=ngram_embedding_info, + rids_int=rids_int, + bootstrap_room_ids_int=bootstrap_room_ids_int, + ) + + +class BaseRunner(ABC): + def __init__(self, model_runner: ModelRunner) -> None: + self.model_runner = model_runner + self.device = model_runner.device + self.device_module = torch.get_device_module(self.device) + self.tp_size = get_parallel().tp_size + # elastic-EP scale-up rewrites dp_size on the published config + self.dp_size = get_parallel().dp_size + self.pp_size = get_parallel().pp_size + self.enable_pdmux = model_runner.server_args.enable_pdmux + self.return_hidden_states_mode = ( + CaptureHiddenMode.NULL + if model_runner.is_draft_worker + else get_server_return_hidden_states_mode() + ) + self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture() + self.attn_tp_size = get_parallel().attn_tp_size + self.attn_tp_rank = get_parallel().attn_tp_rank + self.tbo_plugin = TboCudaGraphRunnerPlugin() + + def warmup(self) -> None: + """Run kernel warmup + autotune once, gated by mr._kernel_warmed_up.""" + mr = self.model_runner + if getattr(mr, "_kernel_warmed_up", False): + return + mr._kernel_warmed_up = True + + if mr.device != "cuda": + return + + self._pre_initialize_flashinfer_allreduce_workspace() + self._pre_initialize_fi_a2a_workspace() + + # Model-owned communication resources may depend on the resolved + # request pool and must be compiled/allocated before graph capture. + prepare_model_resources = getattr( + mr.model, "prepare_before_cuda_graph_capture", None + ) + if prepare_model_resources is not None: + prepare_model_resources(mr) + + # SSKJ-PIE (sglang PR #34528 backport): build + tune + pre-resolve the + # PCIe-IPC AR workspace before the FlashInfer autotune context opens + # -- FlashInfer declines to profile a collective from inside a context + # it did not open, and an unresolved shape reaching CUDA graph capture + # fails there (its resolution is a collective with a host readback). + try: + from sglang.srt.distributed import get_tp_group + from sglang.srt.distributed.device_communicators.pcie_ipc_ar import ( + decode_width_from_args, + ) + + _pcie_ipc_comm = get_tp_group().pcie_ipc_comm + if _pcie_ipc_comm is not None: + _width = decode_width_from_args(getattr(mr, "server_args", None)) + _pcie_ipc_comm.prepare(mr.model_config.hidden_size, decode_width=_width) + except Exception as e: + logger.warning(f"SSKJ-PIE: pcie-ipc prepare failed: {e}") + + if should_run_flashinfer_autotune(self.model_runner): + buffers, batch_size = self._autotune_buffers() + assert ( + buffers is not None + ), "_autotune_buffers() must return a reusable buffer set for autotune" + self._flashinfer_autotune(buffers=buffers, batch_size=batch_size) + maybe_flashinfer_autotune_extend(self, decode_num_tokens=batch_size) + + if ( + envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get() + and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM + and mr.ps.pp_size > 1 + and not mr.spec_algorithm.is_speculative() + ): + from sglang.srt.layers.deep_gemm_wrapper.compile_utils import ( + pp_parallel_deep_gemm_warmup, + ) + + pp_parallel_deep_gemm_warmup(self) + + def _pre_initialize_flashinfer_allreduce_workspace(self): + """Allocate flashinfer allreduce workspaces; must run before CG capture + to keep broadcasts/barriers outside the capture context (else deadlock + with custom_all_reduce.register_graph_buffers). + """ + mr = self.model_runner + if get_exec().comm.flashinfer_allreduce_fusion_backend is None: + return + + from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE + from sglang.srt.layers.flashinfer_comm_fusion import pre_initialize_workspaces + + pre_initialize_workspaces( + max_token_num=FUSE_ALLREDUCE_MAX_BATCH_SIZE, + hidden_dim=mr.model_config.hidden_size, + dtype=mr.dtype, + ) + + def _pre_initialize_fi_a2a_workspace(self): + """Allocate the FlashInfer MNNVL all-to-all workspace for the fi_a2a DCP + comm backend; must run before CG capture (it syncs the stream + barriers + cross-rank, uncapturable) and raises early on non-MNNVL platforms. + """ + if ( + not get_parallel().dcp_enabled + or get_parallel().dcp_comm_backend != "fi_a2a" + ): + return + + from sglang.srt.layers.dcp import init_fi_a2a_workspace + + init_fi_a2a_workspace(get_parallel().dcp_group) + + def _flashinfer_autotune(self, *, buffers, batch_size): + """Run flashinfer autotune. + + buffers / batch_size: a prepared static decode-buffer set and its bs, + reused for the dummy forward instead of allocating a throwaway set. + Supplied by warmup() (the decode runner's captured buffers when a graph + runner exists; a freshly-allocated dummy set in the eager path). + """ + mr = self.model_runner + canary_run_ctx = ( + c.with_active_single_forward_manager(0) + if (c := mr.canary_manager) is not None + else empty_context() + ) + + def forward_fn(): + self._dummy_run( + batch_size=batch_size, + buffers=buffers, + run_ctx=canary_run_ctx, + ) + + run_flashinfer_autotune_forward( + self.model_runner, forward_fn, run_lm_head=False + ) + + def _alloc_dummy_decode_buffers( + self, + max_bs: int, + *, + num_tokens_per_req: int = 1, + allocate_logits_buffer: bool = True, + ): + """Allocate one static decode-buffer set for a dummy forward, sized to + (max_bs, max_bs * num_tokens_per_req). + + The PP-parallel DeepGEMM warmup sweeps batch sizes far larger than any + runner's max_bs (up to ~n_sms*block_m), so no pre-allocated runner buffer + set fits; it builds one here and hands it to _dummy_run (reused across the + sweep; _dummy_run slices it per shape). Eager FlashInfer autotune also + allocates decode-shaped scratch buffers here. Decode cuda-graph autotune + reuses the captured runner buffers instead. + """ + mr = self.model_runner + return _allocate_decode_buffers( + device=mr.device, + max_bs=max_bs, + max_num_token=max_bs * num_tokens_per_req, + hidden_size=mr.model_config.hidden_size, + vocab_size=mr.model_config.vocab_size, + dtype=mr.model_config.dtype, + dp_size=get_parallel().dp_size, + pp_size=get_parallel().pp_size, + is_encoder_decoder=mr.model_config.is_encoder_decoder, + require_mlp_tp_gather=require_mlp_tp_gather(), + seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(), + encoder_len_fill_value=( + getattr(mr.model_config.hf_config, "max_source_positions", 0) + if mr.model_config.is_encoder_decoder + else 0 + ), + num_tokens_per_req=num_tokens_per_req, + cache_loc_dtype=torch.int64, + enable_mamba_track=False, + ne_token_table=( + mr.ngram_embedding_manager.table + if mr.ngram_embedding_manager.enabled + else None + ), + hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None), + pp_proxy_topk_size=mr.get_pp_proxy_topk_size(), + pp_proxy_residual_num_blocks=mr.get_pp_proxy_residual_num_blocks(), + allocate_logits_buffer=allocate_logits_buffer, + ) + + def _dummy_run( + self, + batch_size: int, + run_ctx=None, + forward_mode_override: Optional[ForwardMode] = None, + *, + buffers, + extend_num_tokens_per_req: Optional[int] = None, + ): + """Run a dummy forward pass for warmup/profiling. + + forward_mode_override forces EXTEND/DECODE regardless of + is_generation (used by the PP-parallel DeepGEMM warmup). + + buffers: a prepared static buffer set (or lightweight adapter exposing + the same fields), sized >= this dummy shape, which _dummy_run slices to + (batch_size, num_tokens). The caller owns the shape and the allocation -- + the flashinfer autotune reuses an existing runner's buffers via + _autotune_buffers (the eager input registry, or the decode cuda-graph + runner's captured buffers); the PP-DeepGEMM warmup builds one via + _alloc_dummy_decode_buffers. _dummy_run never allocates and never re-pads + (autotune must run at the reused shape; the PP warmup pre-pads and sizes + its buffer to match). next_token_logits_buffer is optional -- a live + autotune forward returns logits fresh, so the eager-reuse path passes + None (only the PP warmup set still carries one). + """ + mr = self.model_runner + if forward_mode_override is not None: + capture_forward_mode = forward_mode_override + elif mr.is_generation: + capture_forward_mode = ForwardMode.DECODE + else: + capture_forward_mode = ForwardMode.EXTEND + capture_hidden_mode = ( + CaptureHiddenMode.NULL + if mr.is_draft_worker + else get_server_return_hidden_states_mode() + ) + num_tokens_per_req = 1 + # A PD prefill target worker's pool has no SpeculativeState, so a + # TARGET_VERIFY dummy forward would trip the linear-attn backend's + # pool-type assert. Warm up in plain DECODE instead. + _is_pd_prefill_target = ( + get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker + ) + if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target: + if mr.is_draft_worker: + assert ( + mr.spec_algorithm.supports_target_verify_for_draft() + ), "This should not happen" + capture_forward_mode = ForwardMode.TARGET_VERIFY + num_tokens_per_req = mr.decode_num_tokens_per_req() + if extend_num_tokens_per_req is not None: + assert capture_forward_mode == ForwardMode.EXTEND and ( + not mr.spec_algorithm.is_speculative() or _is_pd_prefill_target + ), ( + "extend_num_tokens_per_req requires an ordinary or PD-prefill " + "target EXTEND dummy" + ) + num_tokens_per_req = extend_num_tokens_per_req + + num_tokens = batch_size * num_tokens_per_req + + # Caller owns the shape: passes a static buffer >= the dummy shape; no + # allocation, no re-padding (would overflow the reused buffers). + assert ( + buffers is not None + and num_tokens <= buffers.input_ids.shape[0] + and batch_size <= buffers.seq_lens.shape[0] + ), ( + f"_dummy_run needs a static buffer >= (num_tokens={num_tokens}, " + f"batch_size={batch_size}); got " + + ( + "None" + if buffers is None + else f"(input_ids={buffers.input_ids.shape[0]}, " + f"seq_lens={buffers.seq_lens.shape[0]})" + ) + ) + + if get_flags().capture.enable_torch_compile: + set_torch_compile_config() + should_disable_torch_compile = not getattr( + mr.model, "_can_torch_compile", True + ) + if should_disable_torch_compile: + log_info_on_rank0( + logger, + "Transformers backend model reports it is not torch.compile " + "compatible (e.g. dynamic rope scaling). Disabling torch.compile.", + ) + get_flags().capture.enable_torch_compile = False + + # NOTE: aux hidden state capture (eagle3/dflash) is already + # configured by init_aux_hidden_state_capture() in initialize(). + + # Token-axis buffer views and counters. + input_ids = buffers.input_ids[:num_tokens] + positions = buffers.positions[:num_tokens] + out_cache_loc = buffers.out_cache_loc[:num_tokens] + mrope_positions = buffers.mrope_positions[:, :num_tokens] + buffers.num_token_non_padded[...] = num_tokens + + # Batch-axis buffer views. + req_pool_indices = buffers.req_pool_indices[:batch_size] + seq_lens = buffers.seq_lens[:batch_size] + seq_lens_cpu = buffers.seq_lens_cpu[:batch_size] + + # Optional buffer views. + # Eager-reuse drops the logits buffer; only buffer sets that carry one slice it. + next_token_logits_buffer = ( + buffers.next_token_logits_buffer[:num_tokens] + if buffers.next_token_logits_buffer is not None + else None + ) + encoder_lens = ( + buffers.encoder_lens[:batch_size] + if buffers.encoder_lens is not None + else None + ) + + # For extend mode + if capture_forward_mode == ForwardMode.EXTEND: + if extend_num_tokens_per_req is None: + per_req_extend_len = mr.attn_backend.get_cuda_graph_seq_len_fill_value() + else: + per_req_extend_len = extend_num_tokens_per_req + seq_lens.fill_(per_req_extend_len) + seq_lens_cpu.fill_(per_req_extend_len) + extend_prefix_lens_cpu = [0] * batch_size + extend_seq_lens_cpu = [per_req_extend_len] * batch_size + extend_num_tokens = num_tokens + extend_seq_lens = torch.full( + (batch_size,), per_req_extend_len, dtype=torch.int32, device=mr.device + ) + extend_prefix_lens = torch.zeros( + (batch_size,), dtype=torch.int32, device=mr.device + ) + extend_start_loc = torch.arange( + 0, num_tokens, num_tokens_per_req, dtype=torch.int32, device=mr.device + ) + else: + extend_prefix_lens_cpu = None + extend_seq_lens_cpu = None + extend_num_tokens = None + extend_seq_lens = None + extend_prefix_lens = None + extend_start_loc = None + + if get_parallel().pp_size > 1: + # PP0 already cp-split hidden_states before send. + pp_hidden_tokens = num_tokens + if ( + capture_forward_mode == ForwardMode.EXTEND + and mr.ps.pp_rank != 0 + and mr.ps.attn_cp_size > 1 + ): + pp_hidden_tokens = num_tokens // mr.ps.attn_cp_size + pp_proxy_tensors = PPProxyTensors( + {k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()} + ) + + # TP-gather requirements for global token metadata. + require_mlp_tp_gather_ = require_mlp_tp_gather() + require_attn_tp_gather_ = require_attn_tp_gather() + if require_gathered_buffer(): + assert require_mlp_tp_gather_ or require_attn_tp_gather_ + + if require_mlp_tp_gather_: + global_num_tokens_cpu = [num_tokens] * get_parallel().dp_size + elif require_attn_tp_gather_: + global_num_tokens_cpu = [num_tokens] + else: + global_num_tokens_cpu = None + + if global_num_tokens_cpu is not None: + global_dp_buffer_len = sum(global_num_tokens_cpu) + num_tokens_tensor = torch.tensor( + global_num_tokens_cpu, dtype=torch.int32, device=mr.device + ) + buffers.global_num_tokens_gpu.copy_(num_tokens_tensor) + buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor) + else: + global_dp_buffer_len = None + global_num_tokens_cpu = None + + # Speculative metadata and hidden-state capture mode. + spec_info = create_dummy_verify_input( + mr.spec_algorithm, + buffers.custom_mask, + num_tokens_per_req, + mr.is_draft_worker, + ) + if spec_info is not None and ( + mr.spec_algorithm.is_eagle() or mr.spec_algorithm.is_standalone() + ): + # MTP models (e.g. deepseek_nextn) read spec_info.hidden_states + # during forward; provide a dummy so warmup doesn't crash. + spec_info.hidden_states = torch.zeros( + (num_tokens, mr.model_config.hidden_size), + dtype=mr.dtype, + device=mr.device, + ) + if capture_hidden_mode != CaptureHiddenMode.FULL: + capture_hidden_mode = ( + spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL + ) + + # Optional LoRA metadata. + if mr.lora_manager is not None: + lora_ids = [None] * batch_size + else: + lora_ids = None + + forward_batch = ForwardBatch( + forward_mode=capture_forward_mode, + batch_size=batch_size, + input_ids=input_ids, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + next_token_logits_buffer=next_token_logits_buffer, + orig_seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + seq_lens_sum=seq_lens.sum().item(), + encoder_lens=encoder_lens, + return_logprob=False, + positions=positions, + extend_num_tokens=extend_num_tokens, + extend_seq_lens=extend_seq_lens, + extend_prefix_lens=extend_prefix_lens, + extend_start_loc=extend_start_loc, + extend_prefix_lens_cpu=extend_prefix_lens_cpu, + extend_seq_lens_cpu=extend_seq_lens_cpu, + global_num_tokens_gpu=buffers.global_num_tokens_gpu, + global_num_tokens_cpu=global_num_tokens_cpu, + global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu, + dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(), + global_dp_buffer_len=global_dp_buffer_len, + mrope_positions=mrope_positions, + spec_algorithm=mr.spec_algorithm, + spec_info=spec_info, + capture_hidden_mode=capture_hidden_mode, + num_token_non_padded=buffers.num_token_non_padded, + global_forward_mode=capture_forward_mode, + lora_ids=lora_ids, + ) + + if buffers.ngram_embedding_info is not None: + forward_batch.ngram_embedding_info = buffers.ngram_embedding_info.slice( + batch_size + ) + if lora_ids is not None: + mr.lora_manager.prepare_lora_batch(forward_batch) + + forward_batch = mr.prepare_dummy_forward_batch(forward_batch) + mr.attn_backend.init_forward_metadata(forward_batch) + + def run_once(): + # Reused dummy batches may carry DP-local lazy caches from a prior + # forward. Clear them, then refresh the process-wide DP buffer and + # MoE mode metadata read by model code during this standalone run. + forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None + set_dp_buffer_len( + global_dp_buffer_len, + num_tokens, + forward_batch.dp_padding_mode.is_max_len(), + global_num_tokens_cpu, + ) + set_is_extend_in_batch(False) + + kwargs = {} + if ( + get_parallel().pp_size > 1 + and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters + ): + kwargs["pp_proxy_tensors"] = PPProxyTensors( + {k: v.clone() for k, v in pp_proxy_tensors.tensors.items()} + ) + if not mr.is_generation: + kwargs["get_embedding"] = True + + logits_output_or_pp_proxy_tensors = mr.model.forward( + input_ids, + forward_batch.positions, + forward_batch, + **kwargs, + ) + return logits_output_or_pp_proxy_tensors + + torch.get_device_module(mr.device).synchronize() + mr.tp_group.barrier() + with forward_context(ForwardContext(attn_backend=mr.attn_backend)): + with run_ctx or empty_context(): + run_once() + + def _autotune_buffers(self) -> Tuple[Optional[Any], Optional[int]]: + """Return (buffers, bs) for the autotune dummy forward to reuse; the + EagerRunner and DecodeCudaGraphRunner override this.""" + return None, None + + @abstractmethod + def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ... + + @abstractmethod + def load_batch( + self, + forward_batch: ForwardBatch, + **kwargs, + ) -> Any: ... + + @abstractmethod + def execute( + self, + forward_batch: ForwardBatch, + **kwargs, + ) -> Any: ... diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/environ.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/environ.py new file mode 100644 index 0000000..9724e0d --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/environ.py @@ -0,0 +1,1848 @@ +import base64 +import functools +import json +import os +import warnings +from contextlib import contextmanager +from enum import IntEnum +from typing import Any, Callable, Dict, Optional + + +@functools.lru_cache(maxsize=1) +def _default_hip() -> bool: + """Lazy ROCm/HIP detection for platform-conditional env defaults. + + Avoids importing torch at environ import time (this module is intentionally + stdlib-only and loaded very early). Resolved on first EnvField.get() that uses + it as a default, by which point torch is already imported in any real run; + falls back to False if torch is unavailable. + """ + try: + import torch + + return torch.version.hip is not None + except Exception: + return False + + +_NON_UTF8_PREFIX = "base64:" + + +def _default_cache_subdir(name: str) -> str: + """A directory under SGLANG_CACHE_DIR, for env defaults that track it. + + Pass as a callable default: SGLANG_CACHE_DIR is declared further down the + Envs body, and resolving late also lets tests override it. + """ + return os.path.join(os.path.expanduser(envs.SGLANG_CACHE_DIR.get()), name) + + +def _default_tree_cache_sanity_check() -> bool: + """Enable the expensive tree-cache sanity check by default in CI.""" + return envs.SGLANG_IS_IN_CI.get() + + +class EnvField: + _allow_set_name = True + + def __init__(self, default: Any, secret: bool = False): + self.default = default + # NOTE: environ can only accept str values, so we need a flag to indicate + # whether the env var is explicitly set to None. + self._set_to_none = False + self.secret = secret + + def __set_name__(self, owner, name): + assert EnvField._allow_set_name, "Usage like `a = envs.A` is not allowed" + self.name = name + + def parse(self, value: str) -> Any: + raise NotImplementedError() + + def _resolve_default(self) -> Any: + # Support a callable default for lazily/platform-computed defaults + # (e.g. EnvBool(_default_hip)); evaluated only when the env is unset. + return self.default() if callable(self.default) else self.default + + def get(self) -> Any: + value = os.getenv(self.name) + + # Explicitly set to None + if self._set_to_none: + assert value == str(None) + return None + + # Not set, return default + if value is None: + return self._resolve_default() + + try: + return self.parse(value) + except ValueError as e: + default = self._resolve_default() + warnings.warn( + f'Invalid value for {self.name}: {e}, using default "{default}"' + ) + return default + + def is_set(self): + return self.name in os.environ + + def set(self, value: Any): + self._set_to_none = value is None + os.environ[self.name] = str(value) + + @contextmanager + def override(self, value: Any): + backup_present = self.name in os.environ + backup_value = os.environ.get(self.name) + backup_set_to_none = self._set_to_none + self.set(value) + yield + if backup_present: + os.environ[self.name] = backup_value + else: + os.environ.pop(self.name, None) + self._set_to_none = backup_set_to_none + + def clear(self): + os.environ.pop(self.name, None) + self._set_to_none = False + + def __bool__(self): + raise RuntimeError( + "Please use `envs.YOUR_FLAG.get()` instead of `envs.YOUR_FLAG`" + ) + + def __len__(self): + raise RuntimeError( + "Please use `envs.YOUR_FLAG.get()` instead of `envs.YOUR_FLAG`" + ) + + +class EnvTuple(EnvField): + def parse(self, value: str) -> tuple[str, ...]: + return tuple(s.strip() for s in value.split(",") if s.strip()) + + +class EnvStr(EnvField): + def parse(self, value: str) -> str: + return value + + +class EnvJSON(EnvField): + def parse(self, value: str | None) -> list | dict | None: + if not value: + return None + if os.path.exists(value): + with open(value) as f: + return json.load(f) + return json.loads(value) + + +class EnvBool(EnvField): + def parse(self, value: str) -> bool: + value = value.lower() + if value in ["true", "1", "yes", "y"]: + return True + if value in ["false", "0", "no", "n"]: + return False + raise ValueError(f'"{value}" is not a valid boolean value') + + +class EnvInt(EnvField): + def parse(self, value: str) -> int: + try: + return int(value) + except ValueError: + raise ValueError(f'"{value}" is not a valid integer value') + + +class _DeprecatedEnvFallback: + """Mixin for EnvField subclasses: if the canonical env var is not set, + check *deprecated_name* and emit DeprecationWarning before reading it. + + Usage: + SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias(True, deprecated_name="SGLANG_NSA_FUSE_TOPK") + """ + + def __init__(self, default: Any, deprecated_name: str, secret: bool = False): + super().__init__(default, secret=secret) + self.deprecated_name = deprecated_name + + def get(self) -> Any: + if os.getenv(self.name) is None: + fallback = os.getenv(self.deprecated_name) + if fallback is not None: + warnings.warn( + f"Environment variable '{self.deprecated_name}' is deprecated; " + f"use '{self.name}' instead. " + "The alias will be removed in a future release.", + DeprecationWarning, + stacklevel=2, + ) + os.environ[self.name] = fallback + return super().get() + + +class EnvBoolWithAlias(_DeprecatedEnvFallback, EnvBool): + pass + + +class EnvIntWithAlias(_DeprecatedEnvFallback, EnvInt): + pass + + +class EnvFloat(EnvField): + def parse(self, value: str) -> float: + try: + return float(value) + except ValueError: + raise ValueError(f'"{value}" is not a valid float value') + + +class GateGemvMode(IntEnum): + """Small-batch Inkling gate linear implementation. + + OFF: always the cublas GEMM + PAIR: PDL-chained GEMV and gate JIT kernels + FUSED: single-launch GEMV + gate epilogue (last-block ticket) + """ + + OFF = 0 + PAIR = 1 + FUSED = 2 + + +class ToolStrictLevel(IntEnum): + """ + Defines the strictness levels for tool call parsing and validation. + + OFF: No strict validation + FUNCTION: Enables structural tag constraints for all tools + PARAMETER: Enforces strict parameter validation for all tools + """ + + OFF = 0 + FUNCTION = 1 + PARAMETER = 2 + + +class InvariantCheckLevel(IntEnum): + """Signal level for value/index validity checks (see invariants.py). + + OFF: data layer only (sanitize/containment); no detection, no signal. + WARN: detect + throttled log/count; degrade, never crash (prod on-demand). + STRICT: detect + crash on GUARD/FATAL violations (CI default). + + The data layer is unconditional and independent of this level; only the + detection + signal layer is gated here. + """ + + OFF = 0 + WARN = 1 + STRICT = 2 + + +class DsparkFoldedSampling(IntEnum): + """Sampling support in the graph-folded DSpark draft proposal: OFF = + greedy-only folding, AUTO = on when its buffers fit in free GPU memory, + FORCE = always.""" + + OFF = 0 + AUTO = 1 + FORCE = 2 + + +class Envs: + # Organization principles for this registry: + # - Put every field in exactly one topical section. Prefer an existing + # section; add a new one only when no current section is a clear fit. + # - Group by the behavior and owning call sites, not by name similarity + # alone. Keep closely related lifecycle or feature knobs adjacent. + # - Keep each section focused and below 30 fields. Split growing sections + # by subsystem or lifecycle instead of creating catch-all groups. + # - Order broad runtime subsystems before shared storage and backends; keep + # platform- and model-specific integrations in dedicated later sections. + # - Use the same three-line section header everywhere; do not add ad hoc + # one-line headings or append unrelated fields at the end of a section. + # - Keep vendor-specific aliases with their owning integration, and keep + # test/debug knobs with the feature or test workflow they exercise. + # - Keep explanatory comments attached to their field when moving it. + # - For organization-only changes, AST-check that field names, descriptor + # types, and defaults are unchanged and that only field order moved. + + # =================================================================== + # Runtime configuration and process identity + # =================================================================== + # Per-role config-namespace bookkeeping: off / record / enforce (value is + # validated fail-loud in runtime_context, which resolves it once at import + # so the read stays dynamo-prunable). + SGLANG_ROLE_NAMESPACES = EnvStr("off") + # Record mode: append each newly observed (role, namespace) pair to this + # file so the audit survives signal-killed workers. + SGLANG_ROLE_NAMESPACES_OUT = EnvStr(None) + IS_H200 = EnvBool(False) + SGLANG_ENABLE_TORCH_INFERENCE_MODE = EnvBool(False) + + # =================================================================== + # Model configuration, discovery, and weight loading + # =================================================================== + SGLANG_USE_MODELSCOPE = EnvBool(False) + # Controls weight-file ordering for load-time I/O optimization. + # -1 : no sorting, no staggering; preserves original file order. + # 0 : sort files only; maximizes ordering but may reduce cross-rank I/O concurrency. + # k>0: sort files and stagger per-rank order with factor k. + # Files are processed in groups of (tp_size * k), and rank r starts each + # group at offset (r * k), improving multi-rank I/O concurrency while + # keeping access relatively ordered. + SGLANG_SORT_WEIGHT_FILES = EnvInt(0) + SGLANG_DISABLED_MODEL_ARCHS = EnvTuple(tuple()) + SGLANG_PREFETCH_BLOCK_SIZE_MB = EnvInt(16) + SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION = EnvBool(False) + SGLANG_ENABLE_WEIGHT_LOADER_V2 = EnvBool(False) + # Copy rank-local MoE slices into independent CPU storage before H2D when + # they reference a larger mmap-backed checkpoint storage. + SGLANG_MOE_COPY_WEIGHT_VIEWS_BEFORE_H2D = EnvBool(False) + SGLANG_LOAD_SNAPSHOT_USE_ZMQ = EnvBool(False) + SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN = EnvBool(False) + HF_HUB_DISABLE_XET = EnvBool(False) + # In seconds. If a warmup forward batch takes longer than this, the server will crash to prevent hanging. + # Recommend to increase warmup timeout to 1800 to accommodate some kernel JIT precache e.g. deep gemm + SGLANG_WARMUP_TIMEOUT = EnvFloat(-1) + SGLANG_EXTERNAL_MODEL_PACKAGE = EnvStr("") + SGLANG_EXTERNAL_MM_MODEL_ARCH = EnvStr("") + SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE = EnvStr("") + + # =================================================================== + # HTTP server and health + # =================================================================== + # Decompress request bodies tagged with `x-body-compressed`. + SGLANG_ENABLE_REQUEST_DECOMPRESSION = EnvBool(False) + # Override parsed request fields from headers. + SGLANG_ENABLE_REQUEST_HEADER_OVERRIDES = EnvBool(False) + DISABLE_OPENAPI_DOC = EnvBool(False) + SGLANG_TIMEOUT_KEEP_ALIVE = EnvInt(5) + # Uvicorn multiprocess supervisor pings each worker on this interval; default 5s is + # too short when many workers cold-start and load tokenizers in parallel. + SGLANG_UVICORN_WORKER_HEALTHCHECK_TIMEOUT = EnvInt(10) + SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION = EnvBool(True) + SGLANG_EXPOSE_OWN_ENV_VARS = EnvBool(False) + SGLANG_DIAG_BYPASS_HEALTH_GENERATE = EnvBool(False) + + # =================================================================== + # Logging + # =================================================================== + SGLANG_LOG_GC = EnvBool(False) + SGLANG_LOG_FORWARD_ITERS = EnvBool(False) + SGLANG_LOG_DECODE_GRAPH_KEY = EnvBool(False) + SGLANG_LOG_MS = EnvBool(False) + SGLANG_LOG_REQUEST_EXCEEDED_MS = EnvInt(-1) + SGLANG_LOG_REQUEST_HEADERS = EnvTuple(tuple()) + SGLANG_LOG_SCHEDULER_STATUS_TARGET = EnvStr("") + SGLANG_LOG_SCHEDULER_STATUS_INTERVAL = EnvFloat(60.0) + SGLANG_ENABLE_RANK_CONSENSUS_CHECKER = EnvBool(False) + + # =================================================================== + # IPC, broadcasters, and ports + # =================================================================== + SGLANG_USE_PICKLE_IPC = EnvBool(True) + # Log top-level PickleWrapper frames unwrapped on msgpack IPC decode. + SGLANG_LOG_PICKLE_IPC_OBJECTS = EnvBool(False) + SGLANG_USE_MESSAGE_QUEUE_BROADCASTER = EnvBool(True) + SGLANG_TCP_STORE_PORT = EnvInt(29600) + # Base port hint for ephemeral sockets (ZMQ, SHM broadcaster, etc.). + # When set, get_open_port() and shm_broadcast search upwards from this + # value instead of asking the OS for a random port. Useful to keep all + # SGLang ports in a predictable range behind a firewall. + SGLANG_PORT = EnvInt(None) + SGLANG_BACKUP_PORT_BASE = EnvInt(10000) + + # =================================================================== + # CI and test execution + # =================================================================== + SGLANG_IS_IN_CI = EnvBool(False) + SGLANG_IS_IN_CI_AMD = EnvBool(False) + # Set to true by the check-changes CI job when a PR touches nothing under + # rust/; default false so local and scheduled runs never skip the cargo tests. + SGLANG_SKIP_RUST_TESTS = EnvBool(False) + SGLANG_TEST_MAX_RETRY = EnvInt(None) + # Expand jit_kernel test grids to their full parameter ranges (nightly). + SGLANG_JIT_KERNEL_RUN_FULL_TESTS = EnvBool(False) + SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK = EnvBool(False) + + # =================================================================== + # Crash diagnostics and shutdown + # =================================================================== + SGLANG_CUDA_COREDUMP = EnvBool(False) + # None = unset, letting get_dump_dir() resolve the base (RUNNER_TEMP in CI, + # else /tmp); see debug_utils/cuda_coredump.py. + SGLANG_CUDA_COREDUMP_DIR = EnvStr(None) + SGLANG_FORCE_SHUTDOWN = EnvBool(False) + SGLANG_PYSPY_DUMP_BEFORE_CRASH = EnvBool(True) + SGLANG_CUDA_COREDUMP_BEFORE_CRASH = EnvBool(True) + SGLANG_CUDA_COREDUMP_BEFORE_CRASH_WAIT_SECS = EnvFloat(60.0) + + # =================================================================== + # Constrained decoding and grammar + # =================================================================== + SGLANG_GRAMMAR_POLL_INTERVAL = EnvFloat(0.005) + SGLANG_GRAMMAR_MAX_POLL_ITERATIONS = EnvInt(10000) + SGLANG_DISABLE_OUTLINES_DISK_CACHE = EnvBool(False) + + # =================================================================== + # Fault injection and regression tests + # =================================================================== + SGLANG_TEST_STUCK_DETOKENIZER = EnvFloat(0) + SGLANG_TEST_STUCK_DP_CONTROLLER = EnvFloat(0) + SGLANG_TEST_STUCK_SCHEDULER_INIT = EnvFloat(0) + SGLANG_TEST_STUCK_TOKENIZER = EnvFloat(0) + SGLANG_TEST_CRASH_AFTER_STREAM_OUTPUTS = EnvInt(0) + SGLANG_TEST_REQUEST_TIME_STATS = EnvBool(False) + SGLANG_TEST_DISAGG_FAILURE_PROB = EnvFloat(0.0) + SGLANG_TEST_RETRACT = EnvBool(False) + SGLANG_TEST_RETRACT_INTERVAL = EnvInt(3) + SGLANG_TEST_RETRACT_NO_PREFILL_BS = EnvInt(2**31) + # Scheduler: force lazy extra_buffer prealloc to fail at decode boundaries + SGLANG_TEST_MAMBA_LAZY_ALLOC_FAIL = EnvBool(False) + # KL tests: skip the cache-hit count assertion (e.g. when alloc failure reduces hits) + SGLANG_TEST_SKIP_CACHE_HIT_ASSERT = EnvBool(False) + + # =================================================================== + # CI reporting: per-model metrics jsonl for nightly XPU dashboard + # =================================================================== + # When set, XPU nightly tests append one JSON record per model to this file + # so xpu-ci-job-monitor.yml can render per-model ref/actual/status/duration + # tables. Unset (the default) is a full no-op — pre-existing CI unaffected. + SGLANG_TEST_METRICS_FILE = EnvStr(None) + + # =================================================================== + # PD and scripted-runtime tests + # =================================================================== + SGLANG_TEST_PD_DISAGG_BACKEND = EnvStr("mooncake") + SGLANG_TEST_PD_DISAGG_DEVICES = EnvStr(None) + SGLANG_TEST_FORCE_OPTIMISTIC_PREFILL_RETRY_PROB = EnvFloat(0.0) + SGLANG_TEST_SCRIPTED_RUNTIME = EnvBool(False) + SGLANG_TEST_SCRIPTED_RUNTIME_IPC_ADDR = EnvStr(None) + SGLANG_TEST_SCRIPTED_RUNTIME_OUT_OF_BAND_ERROR_PATH = EnvStr(None) + SGLANG_TEST_SCRIPTED_RUNTIME_SYS_PATH_ENTRY = EnvStr(None) + + # =================================================================== + # Profiling, tracing, and metrics + # =================================================================== + SGLANG_PROFILE_WITH_STACK = EnvBool(True) + SGLANG_PROFILE_RECORD_SHAPES = EnvBool(True) + SGLANG_PROFILE_V2 = EnvBool(False) + SGLANG_ENABLE_NVTX_SCHEDULER = EnvBoolWithAlias( + False, deprecated_name="SGLANG_ENABLE_NVTX" + ) + SGLANG_ENABLE_NVTX_OPERATIONS = EnvBoolWithAlias( + False, deprecated_name="SGLANG_OPERATIONS_ENABLE_PROFILE" + ) + SGLANG_RECORD_STEP_TIME = EnvBool(False) + SGLANG_ENABLE_CUDA_GRAPH_CAPTURE_TRACE = EnvBool(False) + # Opt-in: emit one CUDA-graph capture trace per captured batch size (per-bs). + # SGLANG_ENABLE_CUDA_GRAPH_CAPTURE_TRACE (single combined trace) takes + # precedence when both are set. + SGLANG_GRAPH_BATCH_CAPTURE = EnvBool(False) + SGLANG_TORCH_PROFILER_DIR = EnvStr("/tmp") + # Allocator-history buffer for /start_profile activities=["MEM"]; the + # default truncates long windows (each entry is one alloc/free event). + SGLANG_MEM_PROFILE_MAX_ENTRIES = EnvInt(100000) + SGLANG_OTLP_EXPORTER_SCHEDULE_DELAY_MILLIS = EnvInt(500) + SGLANG_OTLP_EXPORTER_MAX_EXPORT_BATCH_SIZE = EnvInt(64) + SGLANG_TRACE_ASYNC = EnvBool(False) + SGLANG_TRACE_ASYNC_FLUSH_THRESHOLD = EnvInt(100) + SGLANG_ENABLE_METRICS_DEVICE_TIMER = EnvBool(False) + SGLANG_ENABLE_METRICS_DP_ATTENTION = EnvBool(False) + SGLANG_TRACE_LOGITS_E2E = EnvBool(False) + SGLANG_TRACE_LOGITS_E2E_SYNC = EnvBool(False) + SGLANG_TRACE_SAMPLER_E2E = EnvBool(False) + SGLANG_TRACE_QWEN_MOE_DEEPEP_E2E = EnvBool(False) + SGLANG_DEEPEP_V2_TRACE_CONTIG = EnvBool(False) + SGLANG_DEEPEP_V2_TRACE_MASKED = EnvBool(False) + + # =================================================================== + # Debugging and invariant checks + # =================================================================== + SGLANG_DETECT_SLOW_RANK = EnvBool(False) + SGLANG_DEBUG_MEMORY_POOL = EnvBool(False) + SGLANG_VALIDATE_MAMBA_REPLAY_STATE_INDICES = EnvBool(False) + SGLANG_GDN_DECODE_FUSION_LOG_LAYER_HITS = EnvBool(False) + SGLANG_GDN_DECODE_FUSION_VERIFY_REAL_TENSORS = EnvBool(False) + # NaN-fill the unified memory pool at boot (debug repro switch). + SGLANG_DEBUG_POISON_POOL = EnvBool(False) + SGLANG_DEBUG_REVERT_PR = EnvInt(0) + SGLANG_PHASE_CHECKER_DEBUG = EnvBool(False) + SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK = EnvBool(True) + SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY = EnvInt(0) + SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE = EnvBool(True) + # The explicit environment variable still takes precedence over this CI + # default, so production remains opt-in and CI remains opt-out if needed. + SGLANG_ENABLE_TREE_CACHE_SANITY_CHECK = EnvBool(_default_tree_cache_sanity_check) + # Physical KV-page checks: committed<=allocated + no page alias. + SGLANG_CHECK_KV_PAGE_INVARIANTS = EnvBool(False) + SGLANG_TBO_DEBUG = EnvBool(False) + # Timing probe: run the swap-in fully but skip the host->device KV bytes, + # measuring the "IO is free" floor. GARBAGE OUTPUT -- benchmarking only. + SGLANG_DEBUG_HISPARSE_SKIP_IO = EnvBool(False) + # Master switch for all async-asserted invariant probes (NaN, Inf, OOB, + # page alignment). Off in prod; tests turn it on to fail-fast on + # numerical / index violations instead of getting silent NaN cascades. + SGLANG_ENABLE_ASYNC_ASSERT = EnvBool(False) + # Signal level for value/index validity checks (nan/inf/oob/...); see + # invariants.py. OFF (prod default) runs only the free data layer, WARN + # adds throttled logging, STRICT (CI default) crashes on violations. + # Supersedes SGLANG_ENABLE_ASYNC_ASSERT, which is bridged as STRICT until + # every callsite migrates. + SGLANG_INVARIANT_CHECK = EnvInt(InvariantCheckLevel.OFF) + + # =================================================================== + # Runtime simulations + # =================================================================== + SGLANG_SIMULATE_ACC_LEN = EnvFloat(-1) + SGLANG_SIMULATE_ACC_METHOD = EnvStr("match-expected") + SGLANG_SIMULATE_ACC_TOKEN_MODE = EnvStr("fixed") + SGLANG_SIMULATE_UNIFORM_EXPERTS = EnvBool(False) + SGLANG_SIMULATE_ROUND_ROBIN_EXPERTS = EnvBool(False) + + # =================================================================== + # DSpark speculative decoding + # =================================================================== + SGLANG_DSPARK_DEBUG_CONFIDENCE_PREFIX_SCHEDULER = EnvBool(False) + SGLANG_DSPARK_DEBUG_CONFIDENCE_METRICS = EnvBool(False) + SGLANG_DSPARK_DEBUG_DUMP = EnvTuple(tuple()) + SGLANG_DSPARK_LOG_SPS_PRED_INTERVAL = EnvInt(0) + SGLANG_DSPARK_STS_COLLECT_PATH = EnvStr("") + SGLANG_DSPARK_BLOCK_ACCEPT_ESTIMATE_PATH = EnvStr("") + SGLANG_DSPARK_BLOCK_ACCEPT_ONLINE_INTERVAL = EnvInt(0) + SGLANG_DSPARK_ENABLE_SPS_RECORD = EnvBool(False) + SGLANG_DSPARK_FAST_KERNEL = EnvBool(True) + SGLANG_DSPARK_FP32_LM_HEAD = EnvBool(False) + SGLANG_DSPARK_FAST_SAMPLING = EnvBool(True) + SGLANG_DSPARK_FOLDED_SAMPLING = EnvInt(DsparkFoldedSampling.AUTO) + SGLANG_DSPARK_FOLDED_PROPOSAL = EnvBool(True) + SGLANG_DSPARK_STACKED_CTX_KV = EnvBool(True) + SGLANG_DSPARK_EMBED_IN_GRAPH = EnvBool(True) + SGLANG_DSPARK_OPT_MARKOV_W2_BF16 = EnvBool(True) + SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD = EnvBool(True) + SGLANG_DSPARK_OPT_FUSED_GREEDY_MARKOV = EnvBool(False) + SGLANG_DSPARK_ENABLE_MULTI_STREAM = EnvBool(True) + SGLANG_DSPARK_CONFIDENCE_RELAY_LAG_STEPS = EnvInt(2) + + # =================================================================== + # Memory pools and KV-cache sizing + # =================================================================== + SGLANG_NATIVE_MOVE_KV_CACHE = EnvBool(False) + # Disable lazy compaction in the unified memory pool allocator and + # fall back to the per-free eager compaction. Used for production + # A/B and quick rollback. Default False (lazy compaction on). + SGLANG_DISABLE_LAZY_COMPACTION = EnvBool(False) + # Periodically log lazy-compaction stats per sub-pool (observability only). + SGLANG_LOG_LAZY_COMPACTION_STATS = EnvBool(False) + SGLANG_LOG_LAZY_COMPACTION_STATS_INTERVAL_SEC = EnvInt(30) + # HND KV layout folds (page, head) into one paged index for per-kv-head sparse + # page tables (DP attn); paged backends like trtllm_mha consume it directly. + SGLANG_USE_HND_KVCACHE = EnvBool(False) + + # Attention (aiter, ROCm): route NEXTN spec draft_extend (EAGLE-v2 KV + # catch-up) through aiter unified_attention (GQA-packed + split-KV) instead + # of the occupancy-starved mha_batch_prefill FMHA. Independent kill-switch + # for the new path; pairs with SGLANG_AITER_UNIFIED_VERIFY. Default on. + SGLANG_AITER_UNIFIED_DRAFT_EXTEND = EnvBool(True) + # size the KV pool after CUDA-graph capture + SGLANG_ENABLE_POST_CAPTURE_KV_SIZING = EnvBool(False) + + # =================================================================== + # Scheduler token budgeting and admission + # =================================================================== + SGLANG_INIT_NEW_TOKEN_RATIO = EnvFloat(0.7) + SGLANG_MIN_NEW_TOKEN_RATIO_FACTOR = EnvFloat(0.14) + SGLANG_NEW_TOKEN_RATIO_DECAY_STEPS = EnvInt(600) + SGLANG_RETRACT_DECODE_STEPS = EnvInt(20) + SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION = EnvInt(4096) + SGLANG_MAX_NEW_TOKENS_LIMIT = EnvInt(None) + SGLANG_DYNAMIC_CHUNKING_SMOOTH_FACTOR = EnvFloat(0.75) + # Window for the token-weighted recent cache-hit rate used to estimate + # waiting-queue prefill load. + SGLANG_CACHE_HIT_RATE_WINDOW_SECONDS = EnvFloat(15.0) + SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES = EnvInt(None) + SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK = EnvFloat(None) + SGLANG_DATA_PARALLEL_BUDGET_INTERVAL = EnvInt(1) + # Compact extend-attention scheduler tile-budget admission (AMD/HIP-only). + # Budget <= 0 disables; >0 sets the max prefix-extend tiles per batch. + SGLANG_PREFILL_TILE_BUDGET = EnvInt(0) + # Tile-budget mode: "compact" (default, counts actual per-request tiles) or + # "legacy" (rectangular grid, max_extend_len-shaped). + # Internal/testing only - users should not need to change this. + SGLANG_PREFILL_TILE_BUDGET_MODE = EnvStr("compact") + SGLANG_PREFILL_DELAYER_MAX_PREFILL_BS_WINDOW_SIZE = EnvInt(16) + + # =================================================================== + # Scheduler polling, timeouts, and output + # =================================================================== + SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT = EnvInt(1000) + SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DECODE = EnvInt(1) + SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_TARGET_VERIFY = EnvInt(1) + SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_NONE = EnvInt(1) + # in seconds. Set if you observe high memory accumulation over a long serving period. + SGLANG_EMPTY_CACHE_INTERVAL = EnvFloat(-1) + SGLANG_SCHEDULER_MAX_RECV_PER_POLL = EnvInt(-1) + SGLANG_SCHEDULER_SKIP_ALL_GATHER = EnvBool(False) + SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE = EnvBool(False) + SGLANG_KILLPG_ON_SCHEDULER_EXCEPTION = EnvBool(False) + SGLANG_REQ_WAITING_TIMEOUT = EnvFloat(-1) # in seconds + SGLANG_REQ_RUNNING_TIMEOUT = EnvFloat(-1) # in seconds + # For non-streaming requests, the scheduler still flushes intermediate + # output batches to the tokenizer manager every N decoded tokens so that + # `first_token_time`/TTFT can be recorded. Lower this (e.g. to 1) to get + # an accurate TTFT for benchmarking; the upstream default of 50 trades + # off some TTFT-metric accuracy for less IPC overhead. + SGLANG_FORCE_STREAM_INTERVAL = EnvInt(50) + + # =================================================================== + # Overlap scheduler and pipeline parallelism + # =================================================================== + SGLANG_DISABLE_CONSECUTIVE_PREFILL_OVERLAP = EnvBool(False) + # Force delay_sample_func for all overlap decode (not just grammar mode), + # allowing CPU result processing to overlap with subsequent forward computation + # and reducing the impact of sampling overhead on the critical path. + SGLANG_ENABLE_DELAY_SAMPLE = EnvBool(False) + # Force-enable the WAR (write-after-read) barrier for the overlap scheduler + # even when is_cuda() is False (e.g. AMD/ROCm). On CUDA the barrier is + # already enabled regardless of this flag (see start_event_loop). + SGLANG_ENABLE_WAR_BARRIER = EnvBool(False) + # Force the WAR barrier to wait for the whole forward instead of the + # read-done fastpath event. + SGLANG_FORCE_COARSE_WAR_BARRIER = EnvBool(False) + # Enable prefill read-done publication after compliant metadata initialization. + SGLANG_ENABLE_PREFILL_WAR_READ_DONE = EnvBool(False) + # PP: skip output send/recv when the entire batch consists of non-final chunked prefill requests, + # since process_batch_result_prefill discards next_token_ids for those anyway. + SGLANG_PP_SKIP_PURE_CHUNKED_OUTPUT_COMM = EnvBool(False) + SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH = EnvBool(False) + + # =================================================================== + # Radix and sparse KV caches + # =================================================================== + SGLANG_EXPERIMENTAL_CPP_RADIX_TREE = EnvBool(False) + SGLANG_RADIX_FORCE_MISS = EnvBool(False) + SGLANG_CHUNKED_PREFIX_CACHE_THRESHOLD = EnvInt(8192) + SGLANG_MAX_KV_CHUNK_CAPACITY = EnvInt(128 * 1024) + # Kill-switch for the shared-index (IndexShare) swap-in prefetch + # (auto-enabled for GLM-5.2-style DSA); set True to A/B synchronous swap-in. + SGLANG_DISABLE_HISPARSE_PREFETCH = EnvBool(False) + SGLANG_OPT_UNIFIED_CACHE_FREE_OUT_OF_WINDOW_SLOTS = EnvBool(True) + # Decode batches between SWA out-of-window evictions. + SGLANG_SWA_EVICTION_INTERVAL = EnvInt(128) + # Deprecated: the unified radix tree is the default tree cache now, so the + # registry no longer reads this. Kept because a few model/arch call sites + # still assert on it; do not use in new code. + SGLANG_ENABLE_UNIFIED_RADIX_TREE = EnvBool(False) + # Registered TreeCore backend serving the unified radix cache. + SGLANG_UNIFIED_RADIX_TREE_CORE_BACKEND = EnvStr("python") + # TODO(DSV4): @ispobock this has bug on main branch when retract + SGLANG_OPT_SWA_RADIX_CACHE_COMPACT = EnvBool(False) + SGLANG_OPT_SWA_SPLIT_LEAF_ON_INSERT = EnvBool(False) + SGLANG_OPT_SWA_RELEASE_LEAF_LOCK_AFTER_WINDOW = EnvBool(False) + + # =================================================================== + # PD disaggregation runtime + # =================================================================== + # NOTE: For SGLANG_DISAGGREGATION_THREAD_POOL_SIZE, the effective default is + # computed dynamically at runtime based on cpu_count; see disaggregation backends. + SGLANG_DISAGGREGATION_THREAD_POOL_SIZE = EnvInt(None) + SGLANG_DISAGGREGATION_QUEUE_SIZE = EnvInt(4) + SGLANG_DISAGGREGATION_BOOTSTRAP_TIMEOUT = EnvInt(300) + SGLANG_DISAGGREGATION_ZMQ_SEND_TIMEOUT = EnvInt(1) + SGLANG_DISAGGREGATION_HEARTBEAT_INTERVAL = EnvFloat(5.0) + SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2) + SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300) + SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX") + SGLANG_DISAGGREGATION_NIXL_BACKEND_PARAMS = EnvStr("{}") + SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX = EnvBool(True) + SGLANG_DISAGGREGATION_ZMQ_MAX_SOCKETS = EnvInt(16384) + SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False) + SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False) + SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0) + SGLANG_DISAGGREGATION_BOOTSTRAP_ENTRY_CLEANUP_INTERVAL = EnvInt(120) + # Deferred decode-side KV release: on abort, hold an in-flight request's KV + # pages/slot until the prefill acks the transfer drained, or the timeout + # below fires. Off by default (no behavior/perf impact when disabled). + SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE = EnvBool(False) + SGLANG_DISAGGREGATION_DEFERRED_DECODE_KV_RELEASE_TIMEOUT = EnvFloat(30.0) + + # =================================================================== + # Distributed and model-parallel runtime + # =================================================================== + SGLANG_ENABLE_CP_V2 = EnvBool(False) + SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS = EnvBool(False) + # Comma-separated bundle indices for Ray Custom PG mode (e.g., "0,1,2,7"). + SGLANG_RAY_BUNDLE_INDICES = EnvStr("") + # Override the distributed init method used by torch.distributed.init_process_group. + # Set to "env://" to use an externally-created TCPStore via MASTER_ADDR/MASTER_PORT. + SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE = EnvStr(None) + SGLANG_IS_FIRST_RANK_ON_NODE = EnvBool(True) + SGLANG_SYNC_TOKEN_IDS_ACROSS_TP = EnvBool(False) + SGLANG_ENABLE_COLOCATED_BATCH_GEN = EnvBool(False) + SGLANG_SHARED_EXPERT_TP1 = EnvBool(False) + # Replicate the input embedding across TP ranks instead of sharding it + # along the vocab dimension (saves an all-reduce/all-gather in the embed + # lookup at the cost of replicated embedding weights). Drives both the + # target and every draft that shares its embedding (see + # get_embedding_tp_kwargs); they must stay in lock-step. Currently only + # applies to the Deepseek-V2 family (Deepseek V3.1, Kimi K2.5) + drafts. + SGLANG_ENABLE_EMBED_REPLICATION = EnvBool(False) + + # =================================================================== + # Tool calling and native web search + # =================================================================== + SGLANG_FORWARD_UNKNOWN_TOOLS = EnvBool(False) + # Native web search (Exa). EXA_API_KEY is the vendor BYOK credential + # (kept as-is, not renamed to SGLANG_*); the SGLANG_EXA_* knobs tune the + # request defaults for the built-in GPT-OSS web_search tool. + EXA_API_KEY = EnvStr(None, secret=True) + SGLANG_EXA_NUM_RESULTS = EnvInt(10) + SGLANG_EXA_SEARCH_TYPE = EnvStr("auto") + SGLANG_EXA_INCLUDE_HIGHLIGHTS = EnvBool(True) + SGLANG_TOOL_STRICT_LEVEL = EnvInt(ToolStrictLevel.OFF) + + # =================================================================== + # HiCache storage backends and mmap allocation + # =================================================================== + # Per-call cudaHostRegister limit in GB. + SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB = EnvInt(256) + SGLANG_HICACHE_HF3FS_CONFIG_PATH = EnvStr(None) + SGLANG_HICACHE_DECODE_OFFLOAD_STRIDE = EnvInt(None) + SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR = EnvStr(None) + # File-backend LRU eviction (opt-in; sizes accept SI/IEC suffixes, "0" disables). + SGLANG_HICACHE_FILE_BACKEND_MAX_SIZE = EnvStr(None) + SGLANG_HICACHE_FILE_BACKEND_EVICTION_RATIO = EnvFloat(0.9) + SGLANG_HICACHE_FILE_BACKEND_MIN_FREE_SPACE = EnvStr("0") + # Enable client-side metadata caching to optimize filesystem checks (e.g. for Lustre/NFS/FUSE) + SGLANG_HICACHE_FILE_BACKEND_ENABLE_METADATA_CACHE = EnvBool(False) + # Positive cache TTL for filesystem metadata lookups (-1 disables positive expiration) + SGLANG_HICACHE_FILE_BACKEND_METADATA_TTL = EnvFloat(5.0) + # Buffer mode: pin a staged prefetch's device anchor from IO commit to + # consumption so eviction cannot waste the fetch; cap = fraction of pool. + SGLANG_ENABLE_HICACHE_BUFFER_ANCHOR_LOCK = EnvBool(False) + SGLANG_HICACHE_BUFFER_ANCHOR_LOCK_CAP = EnvFloat(0.5) + SGLANG_HICACHE_NIXL_BACKEND_STORAGE_DIR = EnvStr(None) + # Enable O_DIRECT when opening NIXL POSIX backend files (bypasses OS page cache). + # Disable with SGLANG_HICACHE_NIXL_USE_DIRECT_IO=0 or via the + # "use_direct_io": false key in --hicache-storage-backend-extra-config. + SGLANG_HICACHE_NIXL_USE_DIRECT_IO = EnvBool(True) + SGLANG_HUGEPAGE_SIZE = EnvStr("") + + # =================================================================== + # KV-transfer staging and Mooncake transport + # =================================================================== + # Staging buffer for heterogeneous TP KV transfer + SGLANG_DISAGG_STAGING_BUFFER = EnvBool(False) + SGLANG_DISAGG_STAGING_POOL_SIZE_MB = EnvInt(4096) + # TODO(yangminl): remove SGLANG_STAGING_USE_TORCH and the torch fallback in + # staging_buffer.py once Triton kernels are fully validated in production. + SGLANG_STAGING_USE_TORCH = EnvBool(False) + SGLANG_MOONCAKE_CUSTOM_MEM_POOL = EnvStr(None) + ENABLE_ASCEND_TRANSFER_WITH_MOONCAKE = EnvBool(False) + ASCEND_NPU_PHY_ID = EnvInt(-1) + SGLANG_MOONCAKE_SEND_AUX_TCP = EnvBool(False) + SGLANG_ENABLE_FAILED_SESSION_PROBE = EnvBool(False) + SGLANG_FAILED_SESSION_PROBE_INTERVAL_S = EnvFloat(30.0) + + # =================================================================== + # Mooncake store + # =================================================================== + SGLANG_HICACHE_MOONCAKE_CONFIG_PATH = EnvStr(None) + SGLANG_HICACHE_MOONCAKE_REUSE_TE = EnvBool(True) + MOONCAKE_MASTER = EnvStr(None) + MOONCAKE_CLIENT = EnvStr(None) + MOONCAKE_LOCAL_HOSTNAME = EnvStr("localhost") + MOONCAKE_TE_META_DATA_SERVER = EnvStr("P2PHANDSHAKE") + MOONCAKE_GLOBAL_SEGMENT_SIZE = EnvStr("4gb") + MOONCAKE_PROTOCOL = EnvStr("rdma") + MOONCAKE_DEVICE = EnvStr("") + MOONCAKE_MASTER_METRICS_PORT = EnvInt(9003) + MOONCAKE_CHECK_SERVER = EnvBool(False) + MOONCAKE_STANDALONE_STORAGE = EnvBool(False) + MOONCAKE_ENABLE_SSD_OFFLOAD = EnvBool(False) + MOONCAKE_OFFLOAD_FILE_STORAGE_PATH = EnvStr(None) + MOONCAKE_TENANT_ID = EnvStr("default") + + # =================================================================== + # MoRI transport and expert dispatch + # =================================================================== + SGLANG_DEEPEP_V2_FORCE_MAX_LEN = EnvBool(False) + # Send CPU-resident AUX data via RDMA instead of ZMQ TCP (default: TCP). + SGLANG_MORI_SEND_AUX_RDMA = EnvBool(False) + # Number of RDMA Queue Pairs (QPs) used per transfer operation. Higher + # values can increase parallelism and bandwidth utilization. + SGLANG_MORI_QP_PER_TRANSFER = EnvInt(4) + # Number of RDMA work requests posted in a single batch to each QP. Larger + # batch sizes reduce per-operation overhead and improve throughput at the + # cost of higher latency. -1 selects automatic sizing based on the number + # of merged work requests and available endpoints. + SGLANG_MORI_POST_BATCH_SIZE = EnvInt(-1) + # Number of worker threads in the RDMA executor thread pool. More workers + # can improve parallelism for large batch transfers across multiple QPs, + # but excessive threads may cause contention. + SGLANG_MORI_NUM_WORKERS = EnvInt(4) + # Number of sharded synchronous worker threads that drain KV transfers. + # Also the bound on outstanding (posted-but-not-completed) transfers, so it + # is the primary throttle keeping the RDMA send queue from overflowing. + SGLANG_MORI_TRANSFER_SHARDS = EnvInt(8) + # Poll cadence (ms) at which a transfer worker wakes to check the SLA while + # waiting for completion; real completion still wakes it immediately. + SGLANG_MORI_WAIT_POLL_MS = EnvInt(1000) + # Per-transfer SLA (ms) before a KV transfer is failed; 0 disables the SLA + # and relies on the RDMA retry-exceeded timeout only. + SGLANG_MORI_TRANSFER_TIMEOUT_MS = EnvInt(0) + SGLANG_MORI_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(4096) + + # =================================================================== + # AMD, ROCm, and AITER + # =================================================================== + SGLANG_USE_AITER = EnvBool(False) + SGLANG_USE_AITER_AG = EnvBool(True) + # Use reduce_scatter (instead of all_reduce + dp_scatter) for the equal-chunk + # MAX_LEN DP-MoE combine. Default ON for ROCm/HIP (uses the aiter custom + # symmetric-memory kernel), OFF elsewhere (would fall back to RCCL); override + # explicitly to force on/off on any platform. + SGLANG_DP_USE_REDUCE_SCATTER = EnvBool(_default_hip) + # Quantize the variable-length DP-MoE gather payload (SGLANG_DP_USE_GATHERV + # path, prefill/extend only) to fp8-e4m3 with per-token-group-128 scales: + # halves the gathered hidden-state bytes over NCCL; the combine + # (reduce_scatterv) leg stays bf16 (NCCL SUM cannot run on fp8). Lossy on + # the wire — same group quantization the MoE expert GEMMs apply to their + # input anyway, but router/shared-expert reads see rounded values, so this + # stays accuracy-gated and default OFF. + SGLANG_ENABLE_DP_GATHER_FP8 = EnvBool(False) + SGLANG_USE_AITER_UNIFIED_ATTN = EnvBool(False) + # Select the gate/up tile layout for AITER MoE: True -> interleave + # (matches FlyDSL `gate_mode="interleave"` kernels), False -> separated + # (matches `gate_mode="separated"`, the layout used by gptoss_fp4 tuned + # configs and by Mxfp4MoEMethod's post-fix weight shuffle). + SGLANG_USE_AITER_MOE_GU_ITLV = EnvBool(True) + # Fold `silu(gate) * up` into the triton MoE up-GEMM epilogue. W13 rows are + # permuted in place at load so gate/up land in adjacent columns of the same + # output tile, which removes intermediate_cache1 and the standalone + # activation launch per MoE layer. Opt-in because the in-place permute is + # not compatible with runtime weight updates or EPLB expert rearrangement, + # both of which assume the checkpoint's halves layout. + SGLANG_OPT_FUSE_SWIGLU_INTERLEAVED = EnvBool(False) + # Fuse the `residual_add + RMSNorm + zero-pad` triplet that appears + # before the MoE block for models whose MoE input hidden_size must be + # padded up to a stride (e.g. GPT-OSS MXFP4 needs pad to multiple of + # 256). When False (default) the pad runs as a separate + # torch.nn.functional.pad call inside the MoE method. When True, the + # aiter Triton kernel `fused_add_rmsnorm_pad` produces a padded + # post-attention layernorm output in one launch and the MoE method + # skips the explicit pad. Currently only takes effect on the + # post_attention_layernorm path with aiter backend and TP=1. + SGLANG_AITER_FUSE_RMSNORM_PAD = EnvBool(False) + # Physical layout for MHA KV cache. "nhd" (default) keeps the existing + # (size, head_num, head_dim) per-token storage that + # `aiter.mha.mha_batch_prefill_func`/`unified_attention` consume directly. + # "vectorized_5d" allocates K as (num_blocks, H_kv, head_dim/x, page_size, x) + # and V as (num_blocks, H_kv, page_size/x, head_dim, x) (x = 16 / dtype_size), + # matching the SHUFFLE layout that aiter's CK FmhaBatchPrefill kernel and + # `aiter.ops.triton.gluon.pa_decode_gluon` both consume natively. This is + # the SHUFFLE KV layout that enables pa_decode_gluon for full-attn + # decode without runtime permutes. + SGLANG_AITER_KV_CACHE_LAYOUT = EnvStr("nhd") + SGLANG_ROCM_FUSED_DECODE_MLA = EnvBool(False) + SGLANG_ROCM_DISABLE_LINEARQUANT = EnvBool(False) + USE_ROCM_AITER_ROPE_BACKEND = EnvStr("0") + # Enable dual-stream MoE (shared experts vs routed experts) on the + # ROCm/AITER path. Requires GPU_MAX_HW_QUEUES>=5 to avoid HW-queue serialization. + SGLANG_ROCM_USE_MULTI_STREAM = EnvBool(False) + SGLANG_HACK_FLASHMLA_BACKEND = EnvStr("tilelang") + SGLANG_USE_AITER_FP8_PER_TOKEN = EnvBool(False) + # Above 8192 tokens of context, aiter's non-static workspace is large enough + # that mem_fraction_static is scaled by 0.85 to leave room for it. Set this to + # honor an explicitly passed --mem-fraction-static instead. Off by default: + # the reserve is load-bearing, and skipping it OOMs long-context aiter serving + # that fits comfortably with it (67.32 GiB request against 47.40 GiB free on a + # 288 GB MI355 in nightly-4-gpu-mi35x-minimax-m3). Worth setting only when the + # scaled fraction is itself too small to hold the model weights. + SGLANG_AITER_HONOR_EXPLICIT_MEM_FRACTION = EnvBool(False) + # Route Kimi-K3-style h12 + fp8 MLA decode through aiter Triton Gluon when + # import and Triton cga_layout prerequisites hold. Set to 0 to force the + # zero-pad mla_decode_fwd fallback (benchmarking / emergency disable). + SGLANG_AITER_MLA_GLUON = EnvBool(True) + + # DSV4 Aiter flags + SGLANG_OPT_USE_AITER_SILU_MUL = EnvBool(False) + SGLANG_OPT_USE_FUSED_QK_NORM_ROPE = EnvBool(True) + # Unified KV wired the fused qk-norm-rope kernel to decode only, so MTP + # target-verify kept running the norm+RoPE as separate kernels. Set to 0 to + # go back to the unfused chain on the verify path. + SGLANG_OPT_FUSED_QK_NORM_ROPE_VERIFY = EnvBool(True) + SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False) + + # =================================================================== + # Apple Silicon and MLX + # =================================================================== + SGLANG_USE_MLX = EnvBool(False) + SGLANG_MLX_USE_CUSTOM_ROPE = EnvBool(False) + SGLANG_MLX_FUSE_SWIGLU = EnvBool(False) + # Number of decode steps between periodic mx.clear_cache() calls. + # Set to 0 to disable cache clearing entirely. + SGLANG_MLX_CLEAR_CACHE_STEPS = EnvInt(256) + # MLX buffer-cache cap in GB. + SGLANG_MLX_CACHE_LIMIT_GB = EnvFloat(None) + + # =================================================================== + # Ascend NPU + # =================================================================== + SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT = EnvBool(False) + SGLANG_NPU_USE_MULTI_STREAM = EnvBool(False) + SGLANG_NPU_USE_MLAPO = EnvBool(False) + # Forward native implementation for activation gelu tanh for model Skywork-Reward-Gemma-2-27B-v0.2 + SGLANG_NPU_FORWARD_NATIVE_GELUTANH = EnvBool(False) + # Forward native implementation for gemma rms norm for model Skywork-Reward-Gemma-2-27B-v0.2 + SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM = EnvBool(False) + # Delay all-gather after qlora for better performance for Deepseek v3.2 + SGLANG_USE_AG_AFTER_QLORA = EnvBool(False) + # Enable int4x2 weights loading + SGLANG_NPU_W4A4_NEW_PACKING = EnvBool(False) + # Use the graph-safe Triton-Ascend kernel for masked speculative KV commits. + SGLANG_NPU_USE_TRITON_PREFIX_KV_CACHE_STORE = EnvBoolWithAlias( + False, deprecated_name="SGLANG_NPU_USE_TRITON_KV_CACHE_STORE" + ) + # Quantize x to int8 in the dispatch operator (vendor alias consumed by the + # Ascend DeepEP library; the MTP draft-build scopes override it to False). + DEEP_NORMAL_MODE_USE_INT8_QUANT = EnvBool(False) + SGLANG_ZBAL_LOCAL_MEM_SIZE = EnvInt(0) + SGLANG_ZBAL_BOOTSTRAP_URL = EnvStr("") + + # =================================================================== + # MUSA + # =================================================================== + SGLANG_MUSA_FA3_FORCE_UPDATE_METADATA = EnvBool(False) + + # =================================================================== + # Quantization + # =================================================================== + SGLANG_INT4_WEIGHT = EnvBool(False) + SGLANG_CPU_QUANTIZATION = EnvBool(False) + SGLANG_USE_DYNAMIC_MXFP4_LINEAR = EnvBool(False) + SGLANG_FORCE_FP8_MARLIN = EnvBool(False) + SGLANG_MOE_NVFP4_DISPATCH = EnvBool(False) + SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN = EnvBool(False) + SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE = EnvBool(False) + SGLANG_QUANT_ALLOW_DOWNCASTING = EnvBool(False) + SGLANG_FP8_IGNORED_LAYERS = EnvStr("") + SGLANG_FP4_IGNORED_LAYERS = EnvStr("") + # On by default; set SGLANG_ENABLE_FP8_GEMM_CONFIG_TUNE=0 as a kill switch. + # Consults the tuned per-(N, K, M) Triton tile config table in + # apply_fp8_linear. When a tuned config exists for this GPU / weight shape / + # token count, run the Triton w8a8 FP8 GEMM with it; otherwise keep the + # default CUTLASS path. Only takes effect on a GPU with a matching + # dtype=fp8_w8a8_channelwise config JSON under + # kernels/ops/quantization/configs/ (currently L40S), so it is a no-op on + # any other GPU / untuned shape even when enabled. + SGLANG_ENABLE_FP8_GEMM_CONFIG_TUNE = EnvBool(True) + + # =================================================================== + # Humming quantization + # =================================================================== + SGLANG_HUMMING_ONLINE_QUANT_CONFIG = EnvJSON(None) + SGLANG_HUMMING_INPUT_QUANT_CONFIG = EnvJSON(None) + SGLANG_HUMMING_USE_F16_ACCUM = EnvBool(False) + SGLANG_HUMMING_MOE_GEMM_TYPE = EnvStr("") + + # =================================================================== + # FlashInfer, FlashMLA, and TRT-LLM + # =================================================================== + SGLANG_IS_FLASHINFER_AVAILABLE = EnvBool(True) + SGLANG_FLASHINFER_USE_PAGED = EnvBool(False) + # Default to the pick from flashinfer + SGLANG_FLASHINFER_WORKSPACE_SIZE = EnvInt(384 * 1024 * 1024) + # Per-rank dispatch capacity of the FlashInfer MoE A2A dispatcher. Unset + # means each call site keeps its own default. + SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(None) + # Enable per-token FP32 activation scaling for serialized ModelOpt FP4 with + # FlashInfer TRT-LLM or CuTe DSL v2 MoE. + SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False) + # Use BF16 activations with FlashInfer CuTe DSL NVFP4 dense and MoE weights. + SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16 = EnvBool(False) + # Launch the TRT-LLM MoE grouped GEMMs with PDL only at or below this + # token count. + SGLANG_TRTLLM_MOE_PDL_MAX_TOKENS = EnvInt(8192) + # Use FlashInfer's fused atomic CUTLASS/CuTe DSL MoE finalize. + SGLANG_FLASHINFER_MOE_FUSED_FINALIZE = EnvBool(True) + # Master switch for the experimental TRT-LLM LoRA fast path; when OFF (default) every + # fine-grained opt switch reads False, keeping non-experimental paths byte-identical. + SGLANG_EXPERIMENTAL_LORA_OPTI = EnvBool(False) + # SGLang needs to know FlashInfer NVFP4 4over6 config to compute the global scale factor. + FLASHINFER_NVFP4_4OVER6 = EnvBool(False) + FLASHINFER_NVFP4_4OVER6_E4M3_USE_256 = EnvBool(False) + # Skip-softmax threshold scale factor for TRT-LLM attention (prefill and decode separately). + # None = standard attention. See https://arxiv.org/abs/2512.12087 + SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR = EnvFloat(None) + SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR = EnvFloat(None) + # Split TRTLLM-GEN decode attention into sorted, equal-size request groups. + # One preserves the default single-call path; values above one are useful + # for batches whose KV sequence lengths have a large spread. + SGLANG_TRTLLM_MHA_DECODE_SEQ_LEN_SPLITS = EnvInt(1) + # SM120 FlashMLA decode backend: "flashinfer" (default), "triton", or "torch". + SGLANG_SM120_FLASHMLA_BACKEND = EnvStr("flashinfer") + SGLANG_FLASHINFER_PREFILL_SPLIT_TILE_SIZE = EnvInt(4096) + SGLANG_FLASHINFER_DECODE_SPLIT_TILE_SIZE = EnvInt(2048) + SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True) + # Also autotune one EXTEND-shaped dummy at max_prefill_tokens during + # warmup. Opt-in: the extra forward needs transient activation headroom + # that small-VRAM or tightly-packed configs may not have. + SGLANG_FLASHINFER_AUTOTUNE_EXTEND = EnvBool(False) + + # =================================================================== + # Triton and Torch compilation + # =================================================================== + SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS = EnvBool(False) + SGLANG_USE_CUSTOM_TRITON_KERNEL_CACHE = EnvBool(False) + # A-B kill-switch for Work-Centric (Lean) Attention. When True, forces the + # standard Triton decode kernel even if --enable-lean-attention or the auto-gate + # would select Lean. Used to isolate the Lean kernel in benchmarks. + SGLANG_DISABLE_LEAN_ATTENTION = EnvBool(False) + # Persistent-grid size multiplier for the Lean decode kernel: + # total_programs = round(device_CU_count * this). Default 1.0 (one CTA per CU), which + # maximizes KV work-tiles per CTA and minimizes the cross-CTA combine/atomic reduction. + # Kernel + E2E A/B sweeps found 1.0 beats 2.0 across uniform and ragged configs on both + # MI300X (gfx942) and MI355X (gfx950) — 2.0 oversubscribed the CUs and regressed high-batch + # decode. Exposed as a knob (e.g. set 2.0) for grid A/B tuning without a rebuild. + SGLANG_FORCE_LEAN_GRID_CU_MULT = EnvFloat(1.0) + + # Torch Compile + # Compact extend-attention query-tile grid: AMD/HIP-only optimization + # (parity with flash-attn's ragged-aware launch). The feature checks _is_hip + # explicitly in code; this env var allows override (0=force off, 1=force on). + SGLANG_TRITON_COMPACT_EXTEND_ATTENTION = EnvBool(True) + # Raise if Triton loads a kernel after the engine starts serving. This + # verifies that startup warmup covers every kernel specialization used at + # serving time. + SGLANG_CRASH_ON_TRITON_LOAD_AFTER_READY = EnvBool(False) + SGLANG_TRITON_SLOW_COMPILE_THRESHOLD_SECS = EnvFloat(1.0) + SGLANG_TRITON_LOAD_WARNING_THRESHOLD_GB = EnvFloat(1.0) + # gfx950 MLA decode stage-1: pick the launch geometry and split count per batch. + # Reorders the fp32 accumulation, so off by default. + SGLANG_MLA_DECODE_TUNE = EnvBool(False) + SGLANG_ENABLE_TORCH_COMPILE = EnvBool(False) + SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE = EnvInt(4096) + SGLANG_TRITON_DECODE_SPLIT_TILE_SIZE = EnvInt(256) + + # =================================================================== + # Expert parallel load balancing + # =================================================================== + SGLANG_EXPERT_LOCATION_UPDATER_LOG_INPUT = EnvBool(False) + SGLANG_EXPERT_LOCATION_UPDATER_CANARY = EnvBool(False) + SGLANG_EXPERT_LOCATION_UPDATER_LOG_METRICS = EnvBool(False) + SGLANG_LOG_EXPERT_LOCATION_METADATA = EnvBool(False) + SGLANG_EXPERT_DISTRIBUTION_RECORDER_DIR = EnvStr("/tmp") + SGLANG_EPLB_HEATMAP_COLLECTION_INTERVAL = EnvInt(0) + # Chunk size for the rebalance expert-weight P2P exchange; set + # >= num_physical_experts to submit a single batch_isend_irecv. + SGLANG_EPLB_P2P_BATCH_CHUNK_SIZE = EnvIntWithAlias( + 32, deprecated_name="SGLANG_EPLB_ROCM_P2P_BATCH_CHUNK_SIZE" + ) + + # =================================================================== + # DeepGEMM + # =================================================================== + SGLANG_ENABLE_JIT_DEEPGEMM = EnvBool(True) + # Enable the allowlisted low-M BF16 Split-K GEMM path on Blackwell. Shapes + # outside the measured allowlist continue to use CuTe DSL/cuBLAS. + SGLANG_ENABLE_BF16_SPLITK_GEMM = EnvBool(True) + SGLANG_DEEPGEMM_STANDARD_LAYOUT = EnvStr("auto") + SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION = EnvFloat(0.25) + # Cap the DeepGEMM masked grouped-GEMM per-expert padded capacity at + # round_up(max(masked_m), 256) instead of round_up(rank_tokens, 256): + # shrinks the [num_local_experts, m, *] MoE intermediates ~4x under + # load imbalance (they otherwise OOM saturated --moe-runner-backend + # deep_gemm serving). Costs one D2H sync per MoE layer. + SGLANG_OPT_DG_MASKED_M_CAP = EnvBool(False) + # Wide-DP eager prefill uses compact routing storage; masked storage scales + # with num_local_experts and can OOM on skewed batches. + SGLANG_OPT_DG_COMPACT_EAGER = EnvBool(False) + # Drop dp-attention MAX_LEN pad rows from MoE dispatch (StandardDispatcher + # post-translation topk_ids -> -1): pad rows otherwise run the router on + # stale hidden values and burn expert compute whose outputs are discarded; + # colliding pad top-ks also inflate the DeepGEMM masked-GEMM workspace to + # OOM at saturation. Capture-safe (reads only global_num_tokens_gpu). + SGLANG_OPT_MASK_DP_PAD_MOE = EnvBool(False) + SGLANG_JIT_DEEPGEMM_PRECOMPILE = EnvBool(True) + SGLANG_JIT_DEEPGEMM_FAST_WARMUP = EnvBool(False) + SGLANG_JIT_DEEPGEMM_COMPILE_WORKERS = EnvInt(4) + SGLANG_IN_DEEPGEMM_PRECOMPILE_STAGE = EnvBool(False) + # Resolved lazily so it tracks SGLANG_CACHE_DIR, which is defined below. + SGLANG_DG_CACHE_DIR = EnvStr(lambda: _default_cache_subdir("deep_gemm")) + SGLANG_DG_USE_NVRTC = EnvBool(False) + SGLANG_USE_DEEPGEMM_BMM = EnvBool(False) + SGLANG_DEEPGEMM_SANITY_CHECK = EnvBool(False) + SGLANG_DEEPGEMM_PDL = EnvBool(True) + SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP = EnvBool(False) + + # =================================================================== + # Cache directories + # =================================================================== + SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang")) + # JIT kernel build cache. None = unset, resolving to ~/.cache/sglang/jit; + # point it at a persistent mount to share builds across CI jobs. + SGLANG_JIT_CACHE_DIR = EnvStr(None) + # Log, at INFO, which dependency changed whenever a module is rebuilt. + SGLANG_JIT_CACHE_DEBUG = EnvBool(False) + # How many builds to keep per module variant. None = unset = keep all, which + # is what makes reverting an edit an instant hit instead of a rebuild; set + # it to trade that away for disk (1 keeps only the most recent build). + SGLANG_JIT_CACHE_KEEP = EnvInt(None) + + # =================================================================== + # Expert-parallel dispatch and MoE execution + # =================================================================== + # Deprecated in favor of '--deepep-dispatcher-output-dtype bf16' but still + # read by several call sites; do not use in new code. + SGLANG_DEEPEP_BF16_DISPATCH = EnvBool(False) + SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128) + # Per-rank buffer capacity, not a model token limit. + SGLANG_DEEPEP_V2_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128) + # 0 lets ElasticBuffer select its theoretical communication SM/QP counts. + SGLANG_DEEPEP_V2_NUM_SMS = EnvInt(0) + SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS = EnvInt(32) + SGLANG_BLACKWELL_OVERLAP_SHARED_EXPERTS_OUTSIDE_SBO = EnvBool(False) + SGLANG_ENABLE_QWEN_DEEPEP_SHARED_OVERLAP = EnvBool(True) + # Force dynamic Waterfill with runtime EP all-reduce instead of the default + # static local-batch path. + SGLANG_DISABLE_STATIC_WATERFILL = EnvBool(False) + SGLANG_NIXL_EP_BF16_DISPATCH = EnvBool(False) + SGLANG_NIXL_EP_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128) + SGLANG_PPLX_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(128) + SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(True) + # DeepSeek/GLM MoE (deepseek_v2.py): quantize the (dp-gathered) MoE input + # to per-token-group-128 fp8 ONCE and feed both the fused shared-expert + # GEMM (cutlass w8a8 linear) and the routed experts' triton fused runner, + # instead of quantizing the same [T, hidden] tensor twice with different + # scale layouts. Only engages on CUDA with fp8 block-128 weights, the + # standard dispatcher, and the triton MoE runner; falls back silently + # otherwise. + SGLANG_OPT_MOE_QUANT_ONCE = EnvBool(False) + + # =================================================================== + # DeepGEMM Mega MoE + # =================================================================== + SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK = EnvInt(8192) + # Blackwell MegaMoE uses a whole-grid software barrier. Keep a small + # residency margin so every cluster can launch beside other streams. + SGLANG_OPT_DEEPGEMM_MEGA_MOE_RESERVED_SMS = EnvInt(2) + + # =================================================================== + # Top-k kernels + # =================================================================== + SGLANG_OPT_USE_FUSED_HASH_TOPK = EnvBool(True) + # Opt-in: route DeepSeek-V3 grouped topk through the unified Triton router + # instead of the flashinfer/AOT grouped kernels. Off by default (flashinfer is + # the tuned production path); the Triton path is bit-exact on DeepSeek-V3.2 e2e + # and benchmarks at parity, so this is a consolidation escape hatch, not a perf flip. + SGLANG_OPT_USE_JIT_KERNEL_GROUPED_TOPK = EnvBool(False) + SGLANG_OPT_USE_TOPK_V2 = EnvBool(True) + + # =================================================================== + # Kernel selection and fused backends + # =================================================================== + # MiniCPM sparse attention developer switches + SGLANG_MINICPM_FUSE_TOPK = EnvBool(False) + SGLANG_MINICPM_DENSE_AS_SPARSE = EnvBool(False) + SGLANG_MINICPM_FORCE_DENSE = EnvBool(False) + + SGLANG_USE_SGL_FA3_KERNEL = EnvBool(True) + # Force every sglang.kernels BaseFusedOp onto one backend (a KernelBackend + # value, e.g. "torch" / "torch_compile" / "triton" / "aot"); unset = + # auto-select by priority. "torch" flips all fused ops to their pure-torch + # reference implementations for numerical-bug bisection. + SGLANG_FORCE_FUSED_OP_BACKEND = EnvStr(None) + USE_TRITON_W8A8_FP8_KERNEL = EnvBool(False) + SGLANG_MOE_PADDING = EnvBool(False) + + # =================================================================== + # Logits and log-probability processing + # =================================================================== + SGLANG_RETURN_ORIGINAL_LOGPROB = EnvBool(False) + # Sanitize NaN logits before sampling kernels and log a throttled warning + # (see sanitize_nan_logits). + SGLANG_SANITIZE_NAN_LOGITS = EnvBool(False) + SGLANG_ENABLE_LOGPROB_CHUNK = EnvBoolWithAlias( + True, deprecated_name="SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK" + ) + SGLANG_LOGPROB_CHUNK_SIZE = EnvIntWithAlias( + 2048, deprecated_name="SGLANG_LOGITS_PROCESSER_CHUNK_SIZE" + ) + # Compute input logprobs from logits via per-row logsumexp instead of + # materializing the full-vocab log-softmax. Escape hatch only; the two + # paths are mathematically identical. + SGLANG_ENABLE_FAST_INPUT_LOGPROBS = EnvBool(True) + + # =================================================================== + # Deterministic inference and all-reduce + # =================================================================== + SGLANG_ENABLE_DETERMINISTIC_INFERENCE = EnvBool(False) + # Use 1-stage all-reduce kernel on AMD (deterministic, fixed accumulation order) + # If not set: auto (enabled when --enable-deterministic-inference is on) + # Set to 1: force enable (even without --enable-deterministic-inference) + # Set to 0: force disable (use default Aiter AR even with --enable-deterministic-inference) + SGLANG_USE_1STAGE_ALLREDUCE = EnvBool(False) + # NCCL channel count pinned on CUDA so the all-reduce reduces a token the + # same way whatever else shares its batch. Raise it to buy back bandwidth + # on links that can drive more channels. + SGLANG_DETERMINISTIC_NCCL_NCHANNELS = EnvInt(8) + SGLANG_OPT_USE_CUSTOM_ALL_REDUCE_V2 = EnvBool(True) + # Default per-direction workspace cap for CustomAllReduceV2; explicit + # constructor sizes take precedence over this. + SGLANG_CUSTOM_ALL_REDUCE_V2_MAX_SIZE_KB = EnvInt(16 * 1024) + SGLANG_FORCE_CUSTOM_ALL_REDUCE_V2_PULL_SIZE_KB = EnvInt(None) + SGLANG_FORCE_CUSTOM_ALL_REDUCE_V2_PUSH_SIZE_KB = EnvInt(None) + # SSKJ-PIE (sglang PR #34528 backport): FlashInfer PCIe-IPC all-reduce, + # for switch-free intra-node hosts (no NVLink, no multicast) where the + # backends above do not apply. Which shapes the kernels take is + # FlashInfer's own decision; an unsupported shape keeps its NCCL path. + SGLANG_ENABLE_PCIE_IPC_ALLREDUCE = EnvBool(False) + # Elements its workspace is sized for. 0 sizes it for the widest decode + # (cuda_graph_config[decode].max_bs * hidden), leaving prefill chunks on + # NCCL -- measured faster than routing them through these kernels. + SGLANG_PCIE_IPC_MAX_NUMEL = EnvInt(0) + + # =================================================================== + # RoPE cache + # =================================================================== + SGLANG_SPEC_EXPANSION_SAFETY_FACTOR = EnvInt(2) + SGLANG_ROPE_CACHE_FP32 = EnvBool(False) + SGLANG_ROPE_CACHE_SAFETY_MARGIN = EnvInt(256) + SGLANG_ROPE_CACHE_ALIGN = EnvInt(128) + + # =================================================================== + # Speculative decoding + # =================================================================== + SGLANG_ENABLE_OVERLAP_PLAN_STREAM = EnvBool(False) + # Capture the per-replay attention-metadata prep (init_forward_metadata_out_graph) + # into a small CUDA graph, collapsing its host dispatch cost to one launch. + # Experimental; auto-falls back to eager if the backend's prep is not capturable. + SGLANG_ENABLE_METADATA_GLUE_GRAPH = EnvBool(False) + SGLANG_OPT_FUSED_KDA_VERIFY = EnvBool(False) + # A/B: keep the DFLASH draft greedy head eager (not folded in-graph). + SGLANG_DFLASH_EAGER_DRAFT_SAMPLER = EnvBool(False) + SGLANG_RAGGED_VERIFY_MODE = EnvStr("static") + SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE = EnvBool(False) + # Skip draft_extend while adaptive spec is at steps=0 (drafting disabled). + # Saves the per-step draft forward, but the draft KV goes stale: an upshift + # back to steps>0 starts from a cold draft state (low accept until it recovers). + SGLANG_SPEC_SKIP_ZERO_STEP_DRAFT_EXTEND = EnvBool(False) + # Which speculative decisions rank 0 broadcasts to its TP group; narrowing + # it under live traffic isolates where ranks actually diverge. Comma + # separated presets ("all", "rng", "init", "off"), or SpecTpSyncSite slugs + # and numbers, each negatable with a leading "-": "all,-dspark-plan,-6". + SGLANG_SPEC_TP_SYNC = EnvStr("all") + # Kill-switch for the draft-extend cuda graph. Draft extend then always runs + # eager. Escape hatch for setups where the capture's memory pool costs more + # than the graph saves (e.g. DeepEP MoE workspace captured at full dispatch + # capacity). + SGLANG_DISABLE_DRAFT_EXTEND_CUDA_GRAPH = EnvBool(False) + # Use the split-KV (flash-decode) kernel for EAGLE target-verify on the + # Triton backend (ROCm). Only active at speculative topk == 1; falls back to + # extend_attention_fwd for unsupported cases or when set false (e.g. for + # debugging). Correctness is unaffected; this only changes performance. + SGLANG_ENABLE_SPLITKV_VERIFY = EnvBool(True) + SGLANG_NGRAM_FORCE_GREEDY_VERIFY = EnvBool(False) + + # =================================================================== + # Multimodal processing + # =================================================================== + SGLANG_VLM_CACHE_SIZE_MB = EnvInt(100) + SGLANG_IMAGE_MAX_PIXELS = EnvInt(16384 * 28 * 28) + SGLANG_RESIZE_RESAMPLE = EnvStr("") + SGLANG_MM_BUFFER_SIZE_MB = EnvInt(0) + SGLANG_MM_PRECOMPUTE_HASH = EnvBool(False) + SGLANG_VIT_ENABLE_CUDA_GRAPH = EnvBool(False) + # Use the fully-vectorized ViT position-embedding interpolation (no per-image + # Python loop / CPU<->GPU sync). Bit-exact with the legacy implementation; + # set False to fall back to the per-image loop. + SGLANG_VIT_ENABLE_VECTORIZED_POS_EMBED = EnvBool(True) + SGLANG_MM_SKIP_COMPUTE_HASH = EnvBool(False) + # For pre-tokenized (list[int]) multimodal prompts, + # preserve the user's original tokens to avoid retokenization drift. + SGLANG_MM_AVOID_RETOKENIZE = EnvBool(True) + + # =================================================================== + # Multimodal CUDA IPC transport + # =================================================================== + SGLANG_USE_CUDA_IPC_TRANSPORT = EnvBool(False) + # Reuse the mapping for the already-allocated bounded CUDA IPC pool. This + # has no effect unless CUDA IPC feature transport is explicitly selected. + SGLANG_USE_IPC_POOL_HANDLE_CACHE = EnvBool(True) + SGLANG_MM_FEATURE_CACHE_MB = EnvInt(1 * 1024) + SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05) + + # =================================================================== + # Mamba state and cache + # =================================================================== + SGLANG_MAMBA_CONV_DTYPE = EnvStr("bfloat16") + SGLANG_MAMBA_SSM_DTYPE = EnvStr(None) + # Kill-switch for the fused per-slot conv clear/copy kernel (MambaPool); + # falls back to the per-conv-type Python loop. + SGLANG_DISABLE_FUSED_MAMBA_SLOT_OPS = EnvBool(False) + # Opt-in: on the unified radix tree, leave the matched-prefix mamba evictable + # during decode (it is already COW'd to the request's own slot) and shrink the + # mamba pool ratio accordingly. Frees one resident slot per running request, + # raising max_running_requests. Off = original locking + ratio (escape hatch). + SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK = EnvBool(False) + + # =================================================================== + # CUDA graphs and execution buffers + # =================================================================== + SGLANG_USE_BREAKABLE_CUDA_GRAPH = EnvBool(False) + # Guards CUDA graph executable dedup via cudaGraphExecUpdate. + SGLANG_ENABLE_CUDA_GRAPH_DEDUP = EnvBool(False) + SGLANG_MEMORY_SAVER_CUDA_GRAPH = EnvBool(False) + # Reuse wholly-free graph-pool segments for step-local eager allocations. + SGLANG_ENABLE_GRAPH_POOL_BORROW = EnvBool(False) + # Eager forward wraps the ForwardBatch's own tensors instead of copying them + # into the CUDA graph buffer registry (no per-iter device-to-device copy). + SGLANG_EAGER_INPUT_NO_COPY = EnvBool(False) + + # =================================================================== + # Tokenizer, request state, embeddings, and reasoning controls + # =================================================================== + SGLANG_EMBEDDINGS_SPARSE_HEAD = EnvStr(None) + # Think tokens budget: negative means unlimited, >= 0 caps thinking tokens + SGLANG_MAX_THINK_TOKENS = EnvInt(-1) + SGLANG_PATCH_TOKENIZER = EnvBool(True) + SGLANG_REQUEST_STATE_WAIT_TIMEOUT = EnvInt(4) + SGLANG_DEFAULT_THINKING = EnvBool(False) + + # =================================================================== + # Encoder pipeline and disaggregation + # =================================================================== + SGLANG_ENCODER_GRPC_TIMEOUT_SECS = EnvInt(60) + # Encoder receiver selection: http|grpc (used by EPD paths). + SGLANG_ENCODER_MM_RECEIVER_MODE = EnvStr("http") + SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0) + SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0) + SGLANG_ENCODER_HTTP_TIMEOUT = EnvFloat(1800.0) + SGLANG_ENCODER_REQ_TIMEOUT = EnvFloat(180.0) + SGLANG_ENCODER_DISPATCH_MIN_ITEMS = EnvInt(2) + SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU = EnvBool(False) + SGLANG_ENCODER_MAX_BATCH_SIZE = EnvInt(8) + SGLANG_ENCODER_PREPROC_WORKERS = EnvInt(8) + SGLANG_ENCODER_MM_LOAD_WORKERS = EnvInt(4) + # EncoderBootstrapServer health-check tuning. Interval == 0 disables it. + SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_INTERVAL = EnvFloat(10.0) + SGLANG_ENCODER_BOOTSTRAP_HEALTH_CHECK_TIMEOUT = EnvFloat(2.0) + # Seconds before permanently dropping an unhealthy encoder (0 = keep probing). + SGLANG_ENCODER_BOOTSTRAP_EVICTED_TTL = EnvFloat(600.0) + # Persistent receiver-side GPU embedding pool size for mooncake EPD transport. + # 0 disables (per-request register/deregister). 4096 = 4GB default per TP + SGLANG_EMBEDDING_POOL_SIZE_MB = EnvInt(4096) + SGLANG_ENCODER_DP_WORKER_MAX_INFLIGHT = EnvInt(64) + + # =================================================================== + # Native gRPC server + # =================================================================== + # Native gRPC server. SGLANG_GRPC_PORT is the env fallback for the + # --grpc-port CLI flag; setting either enables the native server alongside + # HTTP. The worker-threads knob stays env-only (internal tuning, no CLI + # surface). + SGLANG_GRPC_PORT = EnvInt(None) + SGLANG_GRPC_WORKER_THREADS = EnvInt(4) + + # =================================================================== + # NUMA and CPU affinity + # =================================================================== + SGLANG_SET_CPU_AFFINITY = EnvBool(False) + SGLANG_NUMA_BIND_V2 = EnvBool(True) + SGLANG_AUTO_NUMA_BIND = EnvBool(True) + SGLANG_CRASH_ON_NUMA_BIND_FAILURE = EnvBool(False) + + # =================================================================== + # DeepSeek V4 + # =================================================================== + + # Model and Quantization + # Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint. + SGLANG_DSV4_FP4_EXPERTS = EnvBool(True) + # Set True to dequantize the FP4 experts to FP8 at runtime + SGLANG_DSV4_FP4_DEQUANT = EnvBool(False) + # Flash-0731 also accepts "low"; the active profile is checkpoint-resolved. + SGLANG_DSV4_REASONING_EFFORT = EnvStr("") + # Quantize the SWA fp8 KV cache from bf16-rounded values (matches + # trainer-side QAT and the DSA-CP path) instead of fp32 registers. + SGLANG_DSV4_USE_BF16_KV_QUANT_SOURCE = EnvBool(False) + + # Kernels and indexer + SGLANG_OPT_DEEPGEMM_HC_PRENORM = EnvBool(True) + SGLANG_OPT_USE_TILELANG_MHC_PRE = EnvBool(True) + SGLANG_OPT_USE_TILELANG_MHC_POST = EnvBool(True) + SGLANG_OPT_USE_FLASHINFER_MHC = EnvBool(False) + SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(True) + SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False) + SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(True) + # Per-rank local query rows (after DP-attention sharding when enabled), + # not request ISL. + SGLANG_OPT_DSV4_NONPAGED_INDEXER_MIN_QUERY_TOKENS = EnvInt(8192) + SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True) + SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False) + SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False) + SGLANG_DSV4_COMPRESS_STATE_DTYPE = EnvStr("float32") + SGLANG_FP8_PAGED_MQA_LOGITS_TORCH = EnvBool(False) + SGLANG_OPT_FLASHMLA_SPARSE_PREFILL = EnvBool(True) + + # cache, GEMM, and distributed + SGLANG_OPT_FP8_WO_A_GEMM = EnvBool(True) + # Route the decode wo_a bf16 batched matmul off rocBLAS/Tensile onto aiter's + # tuned batched_gemm_bf16 (gfx95). Off by default; see deepseek_v4.py + # _apply_wo_a_bf16_matmul. + SGLANG_OPT_USE_AITER_BATCHED_GEMM = EnvBool(False) + SGLANG_OPT_BF16_FP32_GEMM_ALGO = EnvStr("cublas") + SGLANG_OPT_FUSE_WQA_WKV = EnvBool(True) + SGLANG_OPT_USE_MULTI_STREAM_OVERLAP = EnvBool(True) + + # =================================================================== + # Inkling + # =================================================================== + SGLANG_OPT_USE_FUSED_GATE_TOPK = EnvBool(True) + # Inside the fused gate: use the CUDA JIT top-k+renorm kernel (v2) instead + # of the triton kernel when the production Inkling shape applies. + SGLANG_OPT_USE_GATE_TOPK_JIT = EnvBool(True) + # Inside the fused gate: replace the cublas gate linear with the + # expert-per-block GEMV JIT kernel at small token counts (GateGemvMode). + SGLANG_OPT_GATE_GEMV_MODE = EnvInt(GateGemvMode.PAIR) + # Capture all multi-layer EAGLE draft-extend steps and the in-graph chain + # rotation into ONE CUDA graph instead of one captured graph per step. + SGLANG_ENABLE_SINGLE_CG_DRAFT = EnvBool(True) + # Draft sampler uses the Gumbel-max trick (argmax(probs / Exp(1))) instead of + # torch.multinomial, whose device-side validity assert breaks draft-graph replay. + SGLANG_OPT_USE_GUMBEL_SAMPLE = EnvBool(True) + # Multi-layer chain-MTP boundary-KV fix: widen the draft-extend window to + # rewrite rejected-draft KV rows before reuse (acc_len repair; on by default). + SGLANG_ENABLE_MTP_BOUNDARY_KV_FIX = EnvBool(True) + SGLANG_OPT_USE_INKLING_MULTI_STREAM_OVERLAP = EnvBool(True) + SGLANG_OPT_USE_INKLING_SHEARED_BIAS = EnvBool(True) + # Use feature-stacked GEMMs for the no-LoRA BF16 shared sink. Eligible LoRA + # serving enables this layout independently of the flag. + SGLANG_OPT_LINEARIZED_SHARED_SINK = EnvBool(True) + # Use the autotuned JIT all-reduce, falling back to torch multimem for + # shapes where it wins. + SGLANG_OPT_USE_INKLING_CUSTOM_AR = EnvBool(True) + # Fuse small-batch decode all-reduce, MLP convolution, and attention norm. + # Requires the custom all-reduce; other shapes use the unfused path. + SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV_NORM = EnvBool(True) + # Fuse eligible extend all-reduce, convolution, and cache updates. + # Supports scattered or full-width state and requires the custom all-reduce. + SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV = EnvBool(True) + # Fuse eligible convolution, QK norm, window, and KV-store prologue work. + # Non-BF16 caches retain the backend KV store. + SGLANG_OPT_USE_INKLING_FUSED_ATTN_PROLOGUE = EnvBool(True) + # Override shared-expert selection: true uses grouped GEMM, false uses BMM. + # When unset, selection follows model, quantization, and LoRA requirements. + SGLANG_OPT_USE_INKLING_SHARED_FUSED_MOE = EnvBool(True) + # Fold the conditional long-context log-scaling tau into its producers + # instead of separate output-sized scale kernels: the fused attn + # prologue's q path (bit-exact, before MXFP8 quantization there) and the + # rel_logits projection's r OPERAND (the diagonal scale commutes through + # the einsum, shrinking the pass by rel_extent/d_rel = 64x; rounding moves + # before the GEMM). Flag-off keeps the standalone apply_log_scaling_tau + # on the outputs. + # Fold the MoE shared-expert partials into the custom AR kernels instead + # of a separate {routed + shared} torch.add per MoE layer; some buckets + # keep a pre-add during the AR stage-in. torch.add numerics + # (bit-identical). Requires SGLANG_OPT_USE_INKLING_CUSTOM_AR. + SGLANG_OPT_USE_INKLING_FUSED_AR_SHARED = EnvBool(True) + SGLANG_OPT_USE_INKLING_FUSED_LOG_TAU = EnvBool(True) + # Dispatch the rel_logits projection around einsum's hidden compaction + # copy of the strided r operand (a view into the packed qkvr output): + # zero-copy strided-batched matmul at small t, JIT row-compact + einsum + # above the band, single-launch tau-folded kernel in the small-t tau + # band. Bit-identical to the plain einsum; flag-off restores it. + SGLANG_OPT_USE_INKLING_REL_PROJ_DISPATCH = EnvBool(True) + # Quantize and store MXFP8 K/V data and scales in one fused kernel. + SGLANG_OPT_INKLING_MXFP8_FUSED_QUANT_STORE = EnvBool(True) + # Default reasoning effort in [0.0, 0.99] when omitted by a request. + # An empty string falls back to the protocol default (0.9); the effort + # directive is always emitted. + SGLANG_INKLING_DEFAULT_REASONING_EFFORT = EnvStr("0.9") + SGLANG_INKLING_RS_MM_PREPROCESS = EnvBool(True) + + # =================================================================== + # DSA backend (GLM 5 and DeepSeek V3.2) + # =================================================================== + SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias( + True, deprecated_name="SGLANG_NSA_FUSE_TOPK" + ) + SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC = EnvBool(False) + SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None) + SGLANG_DSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD = EnvIntWithAlias( + 2048, deprecated_name="SGLANG_NSA_PREFILL_DENSE_ATTN_KV_LEN_THRESHOLD" + ) + SGLANG_DSA_HIP_DISABLE_PRESHUFFLE = EnvBoolWithAlias( + False, deprecated_name="SGLANG_NSA_HIP_DISABLE_PRESHUFFLE" + ) + SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2) + SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False) + SGLANG_DSA_TOPK_BROADCAST = EnvBool(False) + SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False) + # Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: fuse the + # absorbed q bmm with the nope/rope concat + fp8 cast so q is written + # directly in fp8 ("born fp8") and the standalone concat-cast kernel + # disappears. Not bit-exact vs the default path (same rounding stages, + # different GEMM accumulation order), hence default OFF until accuracy- + # gated (oracle + full-set gsm8k). + SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q = EnvBool(False) + # Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: pass a + # per-row valid-topk count (derived from the trailing -1 pad run of the + # topk indices) so the kernel skips whole pad-only topk blocks instead of + # computing masked zero contributions. Bit-exact by construction: skipped + # blocks contain only -1 pads, and -1 entries inside the consumed range + # still take the in-kernel clamp+mask path. + SGLANG_ENABLE_DSA_Q8KV8_TOPK_LENGTH = EnvBool(False) + # Opt-in: run the born-fp8 q-prep (absorbed bmm + concat + fp8 cast, + # ~173us/layer-call) on alt_stream underneath the DSA indexer — the two + # chains fork independently from the q_a_layernorm output. Requires + # SGLANG_ENABLE_DSA_Q8KV8_BORN_FP8_Q; eager-prefill-only via the born + # predicate. Coarse per-layer join keeps the single-slot born-q buffer + # WAR-safe. + SGLANG_ENABLE_DSA_Q8KV8_QPREP_OVERLAP = EnvBool(False) + # Opt-in: fuse the Q8KV8 non-prefix KV prep — cast-concat k/k_rope + # directly into the persistent fp8 kv buffer and zero the pad band in one + # Triton kernel (replaces bf16 _cat + copy_ cast + zero_ tail). + SGLANG_ENABLE_DSA_Q8KV8_KV_CAT_FUSION = EnvBool(False) + # Q8KV8 born-fp8 q-prep codegen: "auto" = per-K Triton dispatch (default); + # "cuda" = the hand-written SM90 WGMMA kernel (bitwise identical to the + # Triton two_dot variant, 1.16-1.38x faster across GLM/DS shapes). + SGLANG_OPT_Q8KV8_QPREP_VARIANT = EnvStr("auto") + + # =================================================================== + # MiniMax M3 + # =================================================================== + SGLANG_OPT_USE_BF16_ROUTER_GEMM = EnvBool(True) + SGLANG_OPT_USE_MINIMAX_DENSE_SPARSE_DECODE = EnvBool(False) + SGLANG_DISABLE_MSA = EnvBool(False) + SGLANG_OPT_USE_MSA_DECODE_UNDER_GRAPH = EnvBool(False) + # Kill switch for the derived fp8 attention-GEMM mode (m3_fp8_attn_gemm_enabled): + # forces the pre-fp8 behavior (bf16 indexer + widening sparse path, bf16 q) + # even when kv_cache_dtype fp8_e4m3 + trtllm_mha + SM100 would activate it. + SGLANG_DISABLE_M3_FP8_ATTN_GEMM = EnvBool(False) + # MiniMax-M3 sparse decode indexer: single JIT radix-select kernel replaces the 2-stage split-K Triton topk. + SGLANG_OPT_USE_MINIMAX_DECODE_TOPK_RADIX = EnvBool(True) + # Fused JIT store (minimax_store_kv_index) of main+index K/V instead of separate + # set_*_buffer copies; falls back when main/index dtypes differ or non-CUDA. + SGLANG_OPT_USE_MINIMAX_FUSED_KV_INDEX_STORE = EnvBool(True) + # MiniMax-M3 MXFP8 MoE experimental fusion toggles (default off; A/B only). + SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False) + SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False) + # MiniMax M3 NPU prefill MAIN-attention: route the sparse main attention through + # the native Ascend FA op `torch.ops.npu.npu_fused_infer_attention_score` (FIA) + # with a per-query CUSTOM block_table + SGLANG_MINIMAX_NPU_PREFILL_FIA = EnvBool(True) + # MiniMax-M3 NPU sparse INDEXER (decode + verify topk block selection): route + # through the native AscendC packed indexer op instead of the Triton indexer. + SGLANG_MINIMAX_NPU_NATIVE_INDEXER = EnvBool(False) + # MiniMax-M3 NPU sparse MAIN-attention (decode-main + verify-main): route the + # sparse main attention through the native AscendC sparse-attention op with the + # cached block_table override. + SGLANG_MINIMAX_NPU_NATIVE_ATTN = EnvBool(False) + # MiniMax-M3 on ROCm force-disables custom all-reduce in its model override + # (arg_groups/overrides.py) when aiter all-reduce fusion is off. Set this to + # opt back in and keep custom/quick all-reduce enabled -- e.g. to run the + # INT4 quick-reduce path via ROCM_QUICK_REDUCE_QUANTIZATION={INT4,INT6,INT8}. + SGLANG_M3_ALLOW_CUSTOM_AR = EnvBool(False) + + # =================================================================== + # Kimi K3 + # =================================================================== + # MNNVL fused all-reduce (bf16, TP8): zero-copy 1shot multicast-push for + # small messages and in-place NVLS 2shot on symmetric-memory tensors for + # large ones, with an optional fused residual add. Covers the KDA o_proj + # output and the latent|shared MoE reduce; everything else falls back to + # the regular all-reduce path. Auto-enabled on SM100/SM103 when + # CustomAllReduceV2 with multicast is available; set 0/1 to override in + # either direction. See srt/layers/k3_ar_fusion.py. + SGLANG_K3_AR_FUSION = EnvBool(False) + # K3 SP-MoE fused residual + reduce-scatter and matching all-gather over + # CustomAllReduceV2's MNNVL push workspace. Auto-probed for the validated + # TP8 GB300 configuration; set 0/1 to override. See + # srt/layers/k3_sp_collective.py. + SGLANG_K3_SP_COLLECTIVE = EnvBool(False) + # Keep K3's post-MoE residual stream token-sharded between consecutive + # SP-MoE layers. The next attention-residual aggregation and snapshot + # bank write run on the local shard, then only the normalized attention + # input is all-gathered. Requires SGLANG_K3_SP_COLLECTIVE. + SGLANG_K3_SP_ATTN_RES = EnvBool(False) + # Fused o_proj GEMM + all-reduce (bf16, TP 2..8, SM100+): one + # kernel computes the TP-local o_proj partial and the cross-rank sum over + # a P2P comm region, replacing the GEMM + NCCL AR pair at M <= 512. + SGLANG_K3_GEMM_AR = EnvBool(False) + # Merge the router gate and routed_expert_down_proj weights so the K3 MoE + # front reads hidden_states once, and run the top-k plus the bf16 cast in one + # epilogue kernel. See kernels/ops/moe/moe_front.py. Default on. + SGLANG_K3_FUSED_FRONT = EnvBool(True) + # Use the ROCm radix-4 router for covered K3 top-k workloads. + SGLANG_K3_RADIX4_TOPK = EnvBool(False) + SGLANG_KIMI_K3_VIT_CUDA_GRAPH_CACHE_CAPACITY = EnvInt(2) + SGLANG_KIMI_K3_VIT_CUDA_GRAPH_MIN_HITS = EnvInt(2) + SGLANG_KIMI_K3_VIT_CUDA_GRAPH_MAX_SEQLEN = EnvInt(6144) + + # =================================================================== + # Symmetric memory + # =================================================================== + SGLANG_SYMM_MEM_PREALLOC_GB_SIZE = EnvInt(-1) + SGLANG_DEBUG_SYMM_MEM = EnvBool(False) + + # Qwen3.5 and GDN + SGLANG_ENABLE_GDN_DECODE_FUSED_PROJ_CONV = EnvBool(True) + SGLANG_TRACE_QWEN35_FINAL_NORM = EnvBool(False) + SGLANG_QWEN35_NATIVE_FINAL_NORM = EnvBool(False) + # One switch enables deferred MoE finalize and AR + residual + RMSNorm. + SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION = EnvBool(False) + # Distinct workspace configurations allowed in one process. Production + # uses one model/configuration per rank, so fail closed on accidental reuse. + SGLANG_FLASHINFER_MNNVL_CUTEDSL_AR_FUSION_MAX_INSTANCES = EnvInt(1) + + # =================================================================== + # Plugin system + # =================================================================== + SGLANG_PLATFORM = EnvStr("") + SGLANG_PLUGINS = EnvStr("") + + # =================================================================== + # KV-Canary and Token-Oracle (testing only) + # =================================================================== + SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024) + SGLANG_KV_CANARY_STATS_PRINT_EVERY_N_STEPS = EnvInt(100) + SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False) + SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS = EnvInt(50) + SGLANG_KV_CANARY_PERTURB_REAL_KV_USED_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_REAL_KV_UNUSED_CACHE_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_REAL_KV_POST_FORWARD_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_PERTURB_TARGET_GROUP = EnvStr(None) + SGLANG_KV_CANARY_PERTURB_NEXT_TOKEN_SWAP_PROB = EnvFloat(0.0) + SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False) + SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT = EnvBool(False) + SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL = EnvInt(0) + SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False) + + # =================================================================== + # Rust server + # =================================================================== + SGLANG_RUST_SERVER = EnvBool(False) + # Build a missing Rust extension from source (auto), require a bundled or + # cached extension (never), or rebuild the local cache entry (force). + SGLANG_RUST_BUILD_MODE = EnvStr("auto") + # Most batched requests one /generate HTTP call may expand into. + SGLANG_MAX_BATCH_REQS_PER_HTTP_REQ = EnvInt(4096) + + # =================================================================== + # Weight Cache Daemon + # =================================================================== + # Paths the daemon and the engine ranks it serves must agree on. Both are + # format templates and must keep the {device_uuid} placeholder: each daemon + # is keyed by the physical GPU it runs on, so a GPU-independent path would + # let one job's client discover another job's daemon. + SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE = EnvStr( + "/tmp/sglang_weight_cache_{device_uuid}.sock" + ) + SGLANG_WEIGHT_CACHE_READY_TEMPLATE = EnvStr( + "/tmp/sglang_weight_cache_{device_uuid}.ready" + ) + + +envs = Envs() +EnvField._allow_set_name = False + + +def exportable_env_vars() -> dict[str, str]: + return { + field.name: _exportable_value(os.environ[field.name]) + for field in sorted( + (value for value in vars(Envs).values() if isinstance(value, EnvField)), + key=lambda field: field.name, + ) + if not field.secret and field.name in os.environ + } + + +def _exportable_value(value: str) -> str: + try: + value.encode() + except UnicodeEncodeError: + return ( + _NON_UTF8_PREFIX + + base64.b64encode(value.encode(errors="surrogateescape")).decode() + ) + return value + + +class _DeprecatedEnv: + """One deprecated env var: warn if it is set, and optionally forward its + (possibly transformed) value to a replacement env var.""" + + def __init__( + self, + replacement: Optional[str] = None, + transform: Optional[Callable[[str], str]] = None, + note: Optional[str] = None, + ): + self.replacement = replacement + self.transform = transform + self.note = note + + def apply(self, old_name: str): + if old_name not in os.environ: + return + message = f"Environment variable {old_name} is deprecated." + if self.replacement is not None: + message += f" Please use {self.replacement} instead." + if self.note is not None: + message += f" {self.note}" + warnings.warn(message) + if self.replacement is not None: + value = os.environ[old_name] + if self.transform is not None: + value = self.transform(value) + os.environ[self.replacement] = value + + +def _ms_to_s(value: str) -> str: + return str(float(value) / 1000.0) + + +def _invert_bool(value: str) -> str: + return "0" if value.lower() in ("true", "1", "yes", "y") else "1" + + +# The single registry for deprecated environment variables, processed once at +# import by _handle_deprecated_envs(). Add new deprecations here instead of +# ad-hoc warnings. For a rename where the old name must keep working through a +# descriptor, use EnvBoolWithAlias / EnvIntWithAlias instead. +_DEPRECATED_ENVS: Dict[str, _DeprecatedEnv] = { + # Renamed: the value is forwarded to the replacement. + "SGLANG_GC_LOG": _DeprecatedEnv(replacement="SGLANG_LOG_GC"), + "SGLANG_CUTEDSL_MOE_NVFP4_DISPATCH": _DeprecatedEnv( + replacement="SGLANG_MOE_NVFP4_DISPATCH" + ), + "SGLANG_ENABLE_THINKING": _DeprecatedEnv(replacement="SGLANG_DEFAULT_THINKING"), + "SGLANG_REASONING_EFFORT": _DeprecatedEnv( + replacement="SGLANG_DSV4_REASONING_EFFORT" + ), + "SGLANG_USE_JIT_ALL_REDUCE": _DeprecatedEnv( + replacement="SGLANG_OPT_USE_CUSTOM_ALL_REDUCE_V2" + ), + # The legacy DISABLE flags have the opposite polarity of their replacement. + "SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK": _DeprecatedEnv( + replacement="SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK", transform=_invert_bool + ), + # Renamed with a unit change. + "SGLANG_QUEUED_TIMEOUT_MS": _DeprecatedEnv( + replacement="SGLANG_REQ_WAITING_TIMEOUT", + transform=_ms_to_s, + note="Note the unit change: milliseconds -> seconds.", + ), + "SGLANG_FORWARD_TIMEOUT_MS": _DeprecatedEnv( + replacement="SGLANG_REQ_RUNNING_TIMEOUT", + transform=_ms_to_s, + note="Note the unit change: milliseconds -> seconds.", + ), + # Removed without replacement. + "SGLANG_PER_TOKEN_GROUP_QUANT_8BIT_V2": _DeprecatedEnv(), + # Superseded by the unified JIT per_token_group_quant, the default CUDA path. + "SGLANG_OPT_USE_JIT_PER_TOKEN_GROUP_QUANT": _DeprecatedEnv(), + "SGLANG_MASKED_GEMM_FAST_ACT": _DeprecatedEnv(), + # The unified free list is kept unsorted between flushes by design; the + # sort-after-merge A/B knob never left its off default and is gone. + "SGLANG_SORT_FREE_LIST_AFTER_MERGE": _DeprecatedEnv(), + "SGLANG_OPT_SWA_EVICT_DROP_PAGE_MARGIN": _DeprecatedEnv(), + # sconv-family kernels always use the CUDA-JIT ports when supported; no toggle. + "SGLANG_OPT_USE_CUDA_SCONV": _DeprecatedEnv(), + # The direct dense BF16 GEMM source is vendored in-tree. + "SGLANG_FLASHINFER_PR4266_SOURCE": _DeprecatedEnv(), + # DSV4 compressor V2 is always used. + "SGLANG_OPT_USE_COMPRESSOR_V2": _DeprecatedEnv(), + # Replaced by CLI flags. + "SGLANG_ENABLE_GRPC": _DeprecatedEnv( + note="Please use '--grpc-port' to enable the native gRPC server." + ), + "SGLANG_SCHEDULER_DECREASE_PREFILL_IDLE": _DeprecatedEnv( + note="Please use '--enable-prefill-delayer' instead." + ), + "SGLANG_PREFILL_DELAYER_MAX_DELAY_PASSES": _DeprecatedEnv( + note="Please use '--prefill-delayer-max-delay-passes' instead." + ), + "SGLANG_PREFILL_DELAYER_TOKEN_USAGE_LOW_WATERMARK": _DeprecatedEnv( + note="Please use '--prefill-delayer-token-usage-low-watermark' instead." + ), + "SGLANG_CUTLASS_MOE": _DeprecatedEnv( + note="Please use '--moe-runner-backend=cutlass' and/or " + "'--speculative-moe-runner-backend=cutlass' instead." + ), + "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_FP4_ACTS": _DeprecatedEnv( + note="Please use '--enable-w4a4-mxfp4-megamoe' instead." + ), + "SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_MXF4_KIND": _DeprecatedEnv( + note="Please use '--enable-w4a4-mxfp4-megamoe' instead." + ), + "SGLANG_DFLASH_PREFILL_REFILL_TARGET": _DeprecatedEnv( + note="DFlash now auto-enables the min-free-slots delay; unset this env. " + "To override the threshold, use '--min-free-slots-delay'." + ), + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": _DeprecatedEnv( + note="The unified radix tree is the default tree cache now; unset this " + "env. The field is still defined for legacy call sites." + ), +} + + +def _handle_deprecated_envs(): + for old_name, deprecation in _DEPRECATED_ENVS.items(): + deprecation.apply(old_name) + + # Rewrite the legacy SGL_ prefix to SGLANG_ (names not covered above). + for key, value in list(os.environ.items()): + if key.startswith("SGL_") and key not in _DEPRECATED_ENVS: + new_key = key.replace("SGL_", "SGLANG_", 1) + warnings.warn( + f"Environment variable {key} is deprecated, please use {new_key}" + ) + os.environ[new_key] = value + + +def third_party_cache_defaults() -> Dict[str, str]: + base = os.path.expanduser(envs.SGLANG_CACHE_DIR.get()) + return { + "TRITON_CACHE_DIR": os.path.join(base, "triton"), + "TORCHINDUCTOR_CACHE_DIR": os.path.join(base, "inductor"), + "CUDA_CACHE_PATH": os.path.join(base, "nv"), + # FlashInfer appends ".cache/flashinfer" to this base itself, so this + # is the base dir rather than the final cache dir. + "FLASHINFER_WORKSPACE_BASE": base, + } + + +def redirect_third_party_caches(): + """Point third-party JIT caches at SGLANG_CACHE_DIR, so a run's compiled + kernels can be cleaned, warmed or volume-mounted as one directory. + + Must be called early. The redirect silently does nothing if either of + these has already happened: + + - FlashInfer was imported. It resolves its workspace at import time. + - Inductor made its first ``cache_dir()`` call. That call setdefaults + TORCHINDUCTOR_CACHE_DIR itself. + """ + for key, value in third_party_cache_defaults().items(): + os.environ.setdefault(key, value) + + +_handle_deprecated_envs() + +# Trigger auto-injection of CUDA coredump env vars when SGLANG_CUDA_COREDUMP=1. +# Best-effort; for strict guarantees, set CUDA_* env vars in the shell before +# launching Python. Imported conditionally to keep the default import of this +# module free of non-stdlib side effects. +if envs.SGLANG_CUDA_COREDUMP.get(): + import sglang.srt.debug_utils.cuda_coredump # noqa: F401, E402 # isort: skip diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/parallel_state.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/parallel_state.py new file mode 100644 index 0000000..df4144e --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/parallel_state.py @@ -0,0 +1,3222 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# Adapted from https://github.com/vllm-project/vllm/blob/v0.6.4.post1/vllm/distributed/parallel_state.py + +# Copyright 2023 The vLLM team. +# Adapted from +# https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/parallel_state.py +# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved. +"""Distributed state. +It takes over the control of the distributed environment from PyTorch. +The typical workflow is: + +- call `init_distributed_environment` to initialize the distributed environment. +- call `initialize_model_parallel` or `ensure_model_parallel_initialized` to + initialize the model parallel groups. + +- any code dealing with the distributed stuff + +- call `destroy_model_parallel` to destroy the model parallel groups. +- call `destroy_distributed_environment` to destroy the distributed environment. + +If you only need to use the distributed environment without model/pipeline + parallelism, you can skip the model parallel initialization and destruction + steps. +""" + +import contextlib +import gc +import logging +import os +import pickle +import weakref +from collections import namedtuple +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass +from datetime import timedelta +from multiprocessing import shared_memory +from typing import Any, Callable, Dict, List, Optional, Tuple, Union +from unittest.mock import patch + +import torch +import torch.distributed +from torch.distributed import Backend, ProcessGroup + +from sglang.srt import platforms +from sglang.srt.compilation.compilation_config import register_split_op +from sglang.srt.distributed.utils import set_global_tcp_store +from sglang.srt.environ import envs +from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( + is_in_tc_piecewise_cuda_graph, +) +from sglang.srt.platforms.device_mixin import _DEVICE_TO_DISTRIBUTED_BACKEND +from sglang.srt.runtime_context import ( + derive_parallel_widths, + get_global_dwdp_manager, + get_parallel, + set_global_dwdp_manager, +) +from sglang.srt.utils import ( + get_current_device_stream_fast, + get_int_env_var, + is_cpu, + is_cuda, + is_cuda_alike, + is_gfx95_supported, + is_hip, + is_musa, + is_npu, + is_shm_available, + is_xpu, +) +from sglang.srt.utils.custom_op import register_custom_op +from sglang.srt.utils.network import get_local_ip_auto +from sglang.srt.utils.stale_shm_cleanup import make_shm_name + +_is_npu = is_npu() +_is_cpu = is_cpu() +_is_xpu = is_xpu() +_is_musa = is_musa() + +TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"]) + +# use int value instead of ReduceOp.SUM to support torch compile +REDUCE_OP_SUM = int(torch.distributed.ReduceOp.SUM) + +# Reuse the user-provided distributed timeout for model-parallel subgroup +# creation so runtime collectives do not silently fall back to backend defaults. +_MODEL_PARALLEL_GROUP_TIMEOUT: Optional[timedelta] = None + + +def get_torch_distributed_pg_options(group_name=None): + if not _is_npu: + return None + + # Only create HCCL options for default group or MoE-related groups + if group_name is not None and "moe" not in group_name: + return None + + import torch_npu + + options = torch_npu._C._distributed_c10d.ProcessGroupHCCL.Options() + hccl_buffer_size = int( + os.environ.get("DEEPEP_HCCL_BUFFSIZE") or os.environ.get("HCCL_BUFFSIZE") or 200 + ) + options.hccl_config = {"hccl_buffer_size": hccl_buffer_size} + return options + + +@dataclass +class GraphCaptureContext: + stream: torch.get_device_module().Stream + + +@dataclass +class P2PWork: + work: Optional[torch.distributed.Work] + payload: Optional[torch.Tensor] + + +def _split_tensor_dict( + tensor_dict: Dict[str, Union[torch.Tensor, Any]], +) -> Tuple[List[Tuple[str, Any]], List[torch.Tensor]]: + """Split the tensor dictionary into two parts: + 1. A list of (key, value) pairs. If the value is a tensor, it is replaced + by its metadata. + 2. A list of tensors. + """ + metadata_list: List[Tuple[str, Any]] = [] + tensor_list: List[torch.Tensor] = [] + for key, value in tensor_dict.items(): + if isinstance(value, torch.Tensor): + # Note: we cannot use `value.device` here, + # because it contains not only the device type but also the device + # index (e.g. "cuda:0"). We only need the device type. + # receiving side will set the device index. + device = value.device.type + metadata_list.append( + (key, TensorMetadata(device, value.dtype, value.size())) + ) + tensor_list.append(value) + else: + metadata_list.append((key, value)) + return metadata_list, tensor_list + + +_group_name_counter: Dict[str, int] = {} + + +def _get_unique_name(name: str) -> str: + """Get a unique name for the group. + Example: + _get_unique_name("tp") -> "tp:0" + _get_unique_name("tp") -> "tp:1" + """ + if name not in _group_name_counter: + _group_name_counter[name] = 0 + newname = f"{name}:{_group_name_counter[name]}" + _group_name_counter[name] += 1 + return newname + + +_groups: Dict[str, Callable[[], Optional["GroupCoordinator"]]] = {} + + +def _register_group(group: "GroupCoordinator") -> None: + _groups[group.unique_name] = weakref.ref(group) + + +@register_custom_op(mutates_args=["tensor"]) +@register_split_op() +def inplace_all_reduce(tensor: torch.Tensor, group_name: str) -> None: + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + group._all_reduce_in_place(tensor) + + +@register_custom_op(out_shape="tensor") +def outplace_all_reduce( + tensor: torch.Tensor, group_name: str, outplace_all_reduce_method: str +) -> torch.Tensor: + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + return group._all_reduce_out_place(tensor, outplace_all_reduce_method) + + +@register_custom_op(out_shape="tensor") +def flashinfer_allreduce(tensor: torch.Tensor, group_name: str) -> torch.Tensor: + """FlashInfer kAllReduce over ``group_name``. + + Registered as a custom op so it stays opaque under Dynamo and can run inside + piecewise CUDA graphs. Applicability is decided by + ``GroupCoordinator._can_use_flashinfer_allreduce`` before the call -- this op + has no fallback of its own. + """ + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + return group._flashinfer_allreduce(tensor) + + +@register_custom_op(mutates_args=["output"]) +def reg_all_gather_into_tensor( + output: torch.Tensor, input: torch.Tensor, group_name: str +) -> None: + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + group._all_gather_into_tensor(output, input) + + +@register_custom_op(mutates_args=["output"]) +def reg_reduce_scatter_tensor( + output: torch.Tensor, input: torch.Tensor, group_name: str +) -> None: + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + group._reduce_scatter_tensor(output, input) + + +@register_custom_op(mutates_args=["output"]) +def reg_all_to_all_single( + output: torch.Tensor, input: torch.Tensor, group_name: str +) -> None: + assert group_name in _groups, f"Group {group_name} is not found." + group = _groups[group_name]() + if group is None: + raise ValueError(f"Group {group_name} is destroyed.") + group._all_to_all_single(output, input) + + +class GroupCoordinator: + """ + PyTorch ProcessGroup wrapper for a group of processes. + PyTorch ProcessGroup is bound to one specific communication backend, + e.g. NCCL, Gloo, MPI, etc. + GroupCoordinator takes charge of all the communication operations among + the processes in the group. It can route the communication to + a specific implementation (e.g. switch allreduce implementation + based on the tensor size and cuda graph mode). + """ + + # available attributes: + rank: int # global rank + ranks: List[int] # global ranks in the group + world_size: int # size of the group + # difference between `local_rank` and `rank_in_group`: + # if we have a group of size 4 across two nodes: + # Process | Node | Rank | Local Rank | Rank in Group + # 0 | 0 | 0 | 0 | 0 + # 1 | 0 | 1 | 1 | 1 + # 2 | 1 | 2 | 0 | 2 + # 3 | 1 | 3 | 1 | 3 + local_rank: int # local rank used to assign devices + rank_in_group: int # rank inside the group + cpu_group: ProcessGroup # group for CPU communication + device_group: ProcessGroup # group for device communication + use_pynccl: bool # a hint of whether to use PyNccl + use_pymscclpp: bool # a hint of whether to use PyMsccl + use_custom_allreduce: bool # a hint of whether to use CustomAllreduce + use_torch_symm_mem_all_reduce: ( + bool # a hint of whether to use TorchSymmMemAllReduce + ) + use_message_queue_broadcaster: ( + bool # a hint of whether to use message queue broadcaster + ) + # communicators are only created for world size > 1 + pynccl_comm: Optional[Any] # PyNccl communicator + ca_comm: Optional[Any] # Custom allreduce communicator + torch_symm_mem_comm: Optional[Any] # Torch symm mem communicator + mq_broadcaster: Optional[Any] # shared memory broadcaster + + def __init__( + self, + group_ranks: List[List[int]], + local_rank: int, + torch_distributed_backend: Union[str, Backend], + use_pynccl: bool, + use_pymscclpp: bool, + use_custom_allreduce: bool, + use_torch_symm_mem_all_reduce: bool, + use_hpu_communicator: bool, + use_xpu_communicator: bool, + use_npu_communicator: bool, + use_message_queue_broadcaster: bool = False, + group_name: Optional[str] = None, + gloo_timeout: timedelta = timedelta(seconds=120 * 60), + recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, + ): + # Set group info + group_name = group_name or "anonymous" + self.unique_name = _get_unique_name(group_name) + _register_group(self) + + # Set rank info + self.rank = torch.distributed.get_rank() + # Joiner group ranks are local; shift them into global rank space. + if rank_offset > 0: + group_ranks = [[r + rank_offset for r in ranks] for ranks in group_ranks] + self.local_rank = local_rank + self.device_group = None + self.cpu_group = None + # Which FlashInfer fusion workspace this group owns, or None when the + # group is not eligible for the allreduce-only kAllReduce path. Stamped + # by _tag_groups_for_flashinfer_allreduce_only() after group init. + self._fi_workspace_hint: Optional[str] = None + self.local_size = get_int_env_var("LOCAL_SIZE", 0) + + if is_cuda_alike(): + device_id = ( + 0 if envs.SGLANG_ONE_VISIBLE_DEVICE_PER_PROCESS.get() else local_rank + ) + self.device = torch.device(f"cuda:{device_id}") + elif _is_npu: + self.device = torch.device(f"npu:{local_rank}") + elif _is_xpu: + self.device = torch.device(f"xpu:{local_rank}") + elif _is_musa: + self.device = torch.device(f"musa:{local_rank}") + else: + self.device = torch.device("cpu") + self.device_module = torch.get_device_module(self.device) + + for ranks in group_ranks: + subgroup_timeout = _MODEL_PARALLEL_GROUP_TIMEOUT + if "mooncake" in torch_distributed_backend: + from mooncake.pg import MooncakeBackendOptions + + pg_active_size = len(ranks) + if not recovered_rank and max_world_size is not None: + assert max_world_size >= len(ranks), ( + f"max_world_size ({max_world_size}) must be >= " + f"group size ({len(ranks)})" + ) + pg_active_size = max_world_size + + pg_active_ranks = torch.zeros( + pg_active_size, dtype=torch.int32, device=self.device + ) + pg_active_ranks[: len(ranks)] = 1 + pg_active_ranks_cpu = torch.zeros(pg_active_size, dtype=torch.int32) + pg_active_ranks_cpu[: len(ranks)] = 1 + + if not recovered_rank and max_world_size is not None: + dev_opts = MooncakeBackendOptions( + pg_active_ranks, recovered_rank, max_world_size + ) + cpu_opts = MooncakeBackendOptions( + pg_active_ranks_cpu, recovered_rank, max_world_size + ) + else: + dev_opts = MooncakeBackendOptions(pg_active_ranks, recovered_rank) + cpu_opts = MooncakeBackendOptions( + pg_active_ranks_cpu, recovered_rank + ) + + active_ranks = pg_active_ranks[: len(ranks)] + active_ranks_cpu = pg_active_ranks_cpu[: len(ranks)] + device_group = torch.distributed.new_group( + ranks, + backend="mooncake", + pg_options=dev_opts, + timeout=subgroup_timeout, + group_desc=f"{group_name}:device", + ) + cpu_group = torch.distributed.new_group( + ranks, + backend="mooncake-cpu", + pg_options=cpu_opts, + timeout=subgroup_timeout, + group_desc=f"{group_name}:cpu", + ) + else: + active_ranks = torch.ones( + len(ranks), dtype=torch.int32, device=self.device + ) + active_ranks_cpu = torch.ones(len(ranks), dtype=torch.int32) + pg_options = get_torch_distributed_pg_options(group_name) + device_group = torch.distributed.new_group( + ranks, + backend=torch_distributed_backend, + pg_options=pg_options, + timeout=subgroup_timeout, + group_desc=f"{group_name}:device", + ) + # a group with `gloo` backend, to allow direct coordination + # between processes through the CPU. + cpu_group = torch.distributed.new_group( + ranks, + backend="gloo", + timeout=gloo_timeout, + group_desc=f"{group_name}:cpu", + ) + if self.rank in ranks: + self.ranks = ranks + self.world_size = len(ranks) + self.rank_in_group = ranks.index(self.rank) + self.device_group = device_group + self.cpu_group = cpu_group + self.active_ranks = active_ranks + self.active_ranks_cpu = active_ranks_cpu + + assert self.cpu_group is not None + assert self.device_group is not None + + # Import communicators + self.use_pynccl = use_pynccl + self.use_pymscclpp = use_pymscclpp + self.use_custom_allreduce = use_custom_allreduce + self.use_torch_symm_mem_all_reduce = use_torch_symm_mem_all_reduce + self.use_hpu_communicator = use_hpu_communicator + self.use_xpu_communicator = use_xpu_communicator + self.use_npu_communicator = use_npu_communicator + self.use_message_queue_broadcaster = use_message_queue_broadcaster + + # Lazy import to avoid documentation build error + from sglang.srt.distributed.device_communicators.custom_all_reduce import ( + dispatch_custom_allreduce, + ) + from sglang.srt.distributed.device_communicators.pymscclpp import ( + PyMscclppCommunicator, + ) + from sglang.srt.distributed.device_communicators.pynccl import ( + PyNcclCommunicator, + ) + from sglang.srt.distributed.device_communicators.pynccl_allocator import ( + debug_check_symmetric_mempool, + is_symmetric_memory_enabled, + use_symmetric_memory, + ) + from sglang.srt.distributed.device_communicators.torch_symm_mem import ( + TorchSymmMemCommunicator, + ) + from sglang.srt.layers.dp_attention import is_allocation_symmetric + + self.is_symmetric_memory_enabled = is_symmetric_memory_enabled + self.use_symmetric_memory = use_symmetric_memory + self.is_allocation_symmetric = is_allocation_symmetric + self.debug_check_symmetric_mempool = debug_check_symmetric_mempool + if is_hip(): + from sglang.srt.distributed.device_communicators.quick_all_reduce import ( + QuickAllReduce, + qr_rocm_arch_available, + ) + + self.pynccl_comm: Optional[PyNcclCommunicator] = None + if use_pynccl and self.world_size > 1: + self.pynccl_comm = PyNcclCommunicator( + group=self.cpu_group, + device=self.device, + is_symmetric_memory_enabled=self.is_symmetric_memory_enabled(), + ) + + self.pymscclpp_comm: Optional[PyMscclppCommunicator] = None + if use_pymscclpp and self.world_size > 1: + self.pymscclpp_comm = PyMscclppCommunicator( + group=self.cpu_group, + device=self.device, + ) + + self.ca_comm: Optional[Any] = None + self.qr_comm: Optional[QuickAllReduce] = None + + self.pcie_ipc_comm: Optional[Any] = None + # SSKJ-PIE (sglang PR #34528 backport): only the tensor-parallel + # group issues the per-layer reductions these kernels target; + # other groups would just pin IPC buffers. + if ( + envs.SGLANG_ENABLE_PCIE_IPC_ALLREDUCE.get() + and self.world_size > 1 + and "tp" in self.unique_name + ): + try: + from sglang.srt.distributed.device_communicators.pcie_ipc_ar import ( + PcieIpcCommunicator, + ) + + # The IPC handshake needs the CUDA (NCCL) group, not the CPU + # one. Autotuning is the other way round: it rendezvouses on + # the host. + self.pcie_ipc_comm = PcieIpcCommunicator( + group=self.device_group, + device=self.device, + cpu_group=self.cpu_group, + ) + except Exception as e: + logger.warning(f"Setup FlashInfer PCIe-IPC allreduce failed with {e}.") + if use_custom_allreduce and self.world_size > 1: + # Initialize a custom fast all-reduce implementation. + try: + CAClass = dispatch_custom_allreduce( + group=self.cpu_group, + device=self.device, + ) + self.ca_comm = CAClass( + group=self.cpu_group, + device=self.device, + ) + except Exception as e: + logger.warning( + f"Setup Custom allreduce failed with {e}. To silence this " + "warning, specify --disable-custom-all-reduce explicitly." + ) + + if is_hip(): + try: + # Initialize a custom quick all-reduce implementation for AMD + # when rocm >= gfx942. Quick reduce is designed as a + # complement to custom allreduce. + # Based on quickreduce (https://github.com/mk1-project/quickreduce). + if qr_rocm_arch_available(): + self.qr_comm = QuickAllReduce( + group=self.cpu_group, device=self.device + ) + except Exception as e: + logger.warning(f"Failed to initialize QuickAllReduce: {e}") + elif self.world_size > 1 and is_hip(): + logger.info("[AR] All-reduce call path: NCCL (custom AR disabled)") + + self.torch_symm_mem_comm: Optional[TorchSymmMemCommunicator] = None + if self.use_torch_symm_mem_all_reduce and self.world_size > 1: + self.torch_symm_mem_comm = TorchSymmMemCommunicator( + group=self.cpu_group, + device=self.device, + ) + + # Create communicator for other hardware backends + from sglang.srt.distributed.device_communicators.hpu_communicator import ( + HpuCommunicator, + ) + from sglang.srt.distributed.device_communicators.npu_communicator import ( + NpuCommunicator, + ) + from sglang.srt.distributed.device_communicators.xpu_communicator import ( + XpuCommunicator, + ) + + self.hpu_communicator: Optional[HpuCommunicator] = None + if use_hpu_communicator and self.world_size > 1: + self.hpu_communicator = HpuCommunicator(group=self.device_group) + + self.xpu_communicator: Optional[XpuCommunicator] = None + if use_xpu_communicator and self.world_size > 1: + self.xpu_communicator = XpuCommunicator(group=self.device_group) + + self.npu_communicator: Optional[NpuCommunicator] = None + if use_npu_communicator and self.world_size > 1: + self.npu_communicator = NpuCommunicator(group=self.device_group) + + # Create message queue + from sglang.srt.distributed.device_communicators.shm_broadcast import ( + MessageQueue, + ) + + self.mq_broadcaster: Optional[MessageQueue] = None + if use_message_queue_broadcaster and self.world_size > 1 and not recovered_rank: + # Recovered ranks create their mq_broadcaster in elastic_ep.py + self.mq_broadcaster = MessageQueue.create_from_process_group( + self.cpu_group, 1 << 22, 6 + ) + + def __repr__(self): + return ( + f"ranks={self.ranks} rank={self.rank} local_rank={self.local_rank} use_pynccl={self.use_pynccl} " + f"device_group={self.device_group} cpu_group={self.cpu_group} unique_name={self.unique_name} " + f"world_size={self.world_size} rank_in_group={self.rank_in_group}" + ) + + @property + def first_rank(self): + """Return the global rank of the first process in the group""" + return self.ranks[0] + + @property + def last_rank(self): + """Return the global rank of the last process in the group""" + return self.ranks[-1] + + @property + def is_first_rank(self): + """Return whether the caller is the first process in the group""" + return self.rank == self.first_rank + + @property + def is_last_rank(self): + """Return whether the caller is the last process in the group""" + return self.rank == self.last_rank + + @property + def next_rank(self): + """Return the global rank of the process that follows the caller""" + rank_in_group = self.rank_in_group + world_size = self.world_size + return self.ranks[(rank_in_group + 1) % world_size] + + @property + def prev_rank(self): + """Return the global rank of the process that precedes the caller""" + rank_in_group = self.rank_in_group + world_size = self.world_size + return self.ranks[(rank_in_group - 1) % world_size] + + @contextmanager + def graph_capture( + self, + graph_capture_context: Optional[GraphCaptureContext] = None, + stream=None, + ): + if graph_capture_context is None: + if stream is None: + stream = self.device_module.Stream() + graph_capture_context = GraphCaptureContext(stream) + else: + stream = graph_capture_context.stream + # We don't need the context of custom quick allreduce because the ipc access + # is already collected in init() and we can capture the quick allreduce directly. + ca_comm = self.ca_comm + maybe_ca_context = nullcontext() if ca_comm is None else ca_comm.capture() + + # ensure all initialization operations complete before attempting to + # capture the graph on another stream + curr_stream = get_current_device_stream_fast() + if curr_stream != stream: + stream.wait_stream(curr_stream) + + with self.device_module.stream(stream), maybe_ca_context: + # In graph mode, we have to be very careful about the collective + # operations. The current status is: + # allreduce \ Mode | Eager | Graph | + # -------------------------------------------- + # quick allreduce | enabled | enabled | + # custom allreduce | enabled | enabled | + # PyNccl | disabled| enabled | + # PyMscclpp | disabled| enabled | + # TorchSymmMem | disabled| enabled | + # torch.distributed | enabled | disabled| + # + # Note: When custom quick allreduce is enabled, a runtime check + # will be performed. If the tensor size is too small, it will + # automatically fall back to the next available option. + # Note that custom allreduce will have a runtime check, if the + # tensor size is too large, it will fallback to the next + # available option. + # Note that the PyMsccl needs to register the tensor in ahead, + # which will introduce large overhead in the eager case, + # therefore it is only supported in the graph case. + # In summary: We select the appropriate allreduce method for + # each mode based on the algorithm order in the table and + # their usage conditions. + pynccl_comm = self.pynccl_comm + maybe_pynccl_context: Any + if not pynccl_comm: + maybe_pynccl_context = nullcontext() + else: + maybe_pynccl_context = pynccl_comm.change_state(enable=True) + + pymscclpp_comm = self.pymscclpp_comm + maybe_pymscclpp_context: Any + if not pymscclpp_comm: + maybe_pymscclpp_context = nullcontext() + else: + maybe_pymscclpp_context = pymscclpp_comm.change_state(enable=True) + with maybe_pynccl_context, maybe_pymscclpp_context: + yield graph_capture_context + + def all_reduce(self, input_: torch.Tensor) -> torch.Tensor: + """ + User-facing all-reduce function before we actually call the + all-reduce operation. + + We need this because Dynamo does not support passing an arbitrary + object (`self` in this case) to a custom op. We need to pass the + group name as a string, and then look up the group coordinator from + the group name, dispatch the all-reduce operation to the group + coordinator. + + In addition, PyTorch custom ops do not support mutation or returning + a new tensor in the same op. So we need to figure out if the op is + in-place or out-of-place ahead of time — except under Dynamo tracing, + where the method selection would guard on the symbolic shape; there we + always emit the out-of-place op with method "auto" and resolve the + method at runtime inside the op. + """ + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return input_ + + if input_.is_cpu: + if is_shm_available(input_.dtype, self.world_size, self.local_size): + torch.ops.sgl_kernel.shm_allreduce(input_, REDUCE_OP_SUM) + else: + torch.distributed.all_reduce(input_, group=self.device_group) + return input_ + + if self.hpu_communicator is not None and not self.hpu_communicator.disabled: + return self.hpu_communicator.all_reduce(input_) + + if self.xpu_communicator is not None and not self.xpu_communicator.disabled: + # Route through inplace_all_reduce custom op so Dynamo treats this as + # an opaque call and does not decompose it into _c10d_functional primitives + # (which invoke sycl_event.wait() and break XPU graph capture). + # Keeps the operation in-place; the all-reduce is performed by + # _all_reduce_in_place, which for XPU falls through to + # torch.distributed.all_reduce on self.device_group (the same group + # used by xpu_communicator). + inplace_all_reduce(input_, group_name=self.unique_name) + return input_ + + if self.npu_communicator is not None and not self.npu_communicator.disabled: + return self.npu_communicator.all_reduce(input_) + + if torch.compiler.is_compiling(): + if self._can_use_flashinfer_allreduce(input_): + return flashinfer_allreduce(input_, group_name=self.unique_name) + + # Byte-size thresholds in method selection (e.g. `_pick_algo` or + # `should_mscclpp_allreduce`) would guard on the symbolic token dim + # and recompile per shape; defer the selection to runtime inside + # the opaque custom op. Groups without any accelerated + # communicator keep the inplace split op so their collective + # stays outside captured graphs. The symmetric-memory in-place + # path below is deliberately bypassed under compile: its raw + # pynccl call is untraceable (hard error with fullgraph, graph + # break otherwise) and its in-place contract does not fit the + # outplace custom op. + if ( + self.ca_comm is None + and self.qr_comm is None + and self.pymscclpp_comm is None + and self.torch_symm_mem_comm is None + and self.pynccl_comm is None + ): + inplace_all_reduce(input_, group_name=self.unique_name) + return input_ + return outplace_all_reduce( + input_, + group_name=self.unique_name, + outplace_all_reduce_method="auto", + ) + + should_use_pymscclpp_allreduce = ( + self.pymscclpp_comm is not None + and self.pymscclpp_comm.should_mscclpp_allreduce(input_) + ) + should_use_custom_allreduce = ( + self.ca_comm is not None + and not self.ca_comm.disabled + and self.ca_comm.should_custom_ar(input_) + ) + if ( + self.pynccl_comm is not None + and self.is_symmetric_memory_enabled() + and not should_use_pymscclpp_allreduce + and not should_use_custom_allreduce + ): + self.debug_check_symmetric_mempool(self, {"input": input_}, "all_reduce") + with self.pynccl_comm.change_state(enable=True): + self.pynccl_comm.all_reduce(input_) + return input_ + + if self._can_use_flashinfer_allreduce(input_): + return flashinfer_allreduce(input_, group_name=self.unique_name) + + outplace_all_reduce_method = self._resolve_outplace_all_reduce_method( + input_=input_, + should_use_pymscclpp_allreduce=should_use_pymscclpp_allreduce, + ) + if outplace_all_reduce_method is not None: + return outplace_all_reduce( + input_, + group_name=self.unique_name, + outplace_all_reduce_method=outplace_all_reduce_method, + ) + else: + inplace_all_reduce(input_, group_name=self.unique_name) + return input_ + + def quant_all_reduce(self, input_: torch.Tensor) -> torch.Tensor: + """ + User-facing quant-all-reduce function similar to all-reduce. (NPU support only) + """ + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return input_ + + if self.npu_communicator is not None and not self.npu_communicator.disabled: + return self.npu_communicator.quant_all_reduce(input_) + else: + inplace_all_reduce(input_, group_name=self.unique_name) + return input_ + + def fused_allreduce_rmsnorm( + self, + input_: torch.Tensor, + residual_inp_: torch.Tensor, + weight_: torch.Tensor, + eps: float, + ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Attempt fused all-reduce + RMSNorm via custom all-reduce communicator. ROCm/HIP Only""" + ca_comm = self.ca_comm + if ca_comm is None or getattr(ca_comm, "disabled", True): + return None + + # Prefer communicator-native fused API when provided. + if hasattr(ca_comm, "fused_allreduce_rmsnorm"): + try: + return ca_comm.fused_allreduce_rmsnorm( + input_, residual_inp_, weight_, eps + ) + except Exception: + # Fall back to custom_fused_ar_rms path below. + pass + + if not hasattr(ca_comm, "custom_fused_ar_rms"): + return None + + # 1-stage vs 2-stage selection for fused AR+RMSNorm: + # The 1-stage kernel launches one block per token and is capped at + # 80 tokens (kMaxBlocks). Guard with a byte threshold so large + # prefill batches fall through to the 2-stage kernel instead of + # hitting a runtime error. AITER's C++ dispatch already gates + # which hidden_dims have valid 1-stage support. + if envs.SGLANG_USE_1STAGE_ALLREDUCE.is_set(): + use_1stage_ar = envs.SGLANG_USE_1STAGE_ALLREDUCE.get() + else: + total_bytes = input_.numel() * input_.element_size() + use_1stage_ar = total_bytes <= 128 * 1024 + + if ( + getattr(ca_comm, "_IS_CAPTURING", False) + and not torch.cuda.is_current_stream_capturing() + and is_in_tc_piecewise_cuda_graph() + ): + if not hasattr(ca_comm, "fused_ar_rms"): + return None + return ca_comm.fused_ar_rms( + input_, + residual_inp_, + w=weight_, + eps=eps, + registered=False, + use_1stage=use_1stage_ar, + ) + fused_outputs = ca_comm.custom_fused_ar_rms( + input_, + residual_inp_, + weight_, + eps, + use_1stage_ar, + ) + return fused_outputs + + def fused_allreduce_rmsnorm_quant_per_group( + self, + input_: torch.Tensor, + residual_inp_: torch.Tensor, + weight_: torch.Tensor, + eps: float, + group_size: int = 128, + emit_bf16: bool = False, + transpose_scale: bool = False, + ) -> Optional[Tuple[torch.Tensor, ...]]: + """Attempt fused all-reduce + RMSNorm + per-group FP8 quant. + + ROCm/aiter/gfx95-only entry point. Returns ``None`` on any other + platform or when the aiter custom-all-reduce communicator cannot + service the request, letting the caller fall back to the existing + ``fused_allreduce_rmsnorm`` + separate per-group quant path. + + When ``emit_bf16=True`` the fused kernel also writes the + pre-quantization bf16/fp16 normed output and returns + ``(fp8, residual_out, scale, bf16)`` — used by GDN-style layers that + need both an FP8 projection and a bf16 gating projection without + launching a separate per-group quant kernel. + + When ``transpose_scale=True`` the kernel writes the per-group scale in + the column-major layout the gfx95 bpreshuffle GEMM consumes, so the + caller can skip the post-kernel scale transpose. + """ + if not (is_hip() and is_gfx95_supported()): + return None + + ca_comm = self.ca_comm + if ca_comm is None or getattr(ca_comm, "disabled", True): + return None + if not hasattr(ca_comm, "custom_fused_ar_rms_per_group_quant"): + return None + + # Shape / size eligibility mirrors aiter's internal gate so we fail + # fast without entering the HIP kernel dispatch. + K = input_.shape[-1] + if K % group_size != 0 or K > 16384: + return None + total_bytes = input_.numel() * input_.element_size() + if total_bytes == 0 or total_bytes > 8 * 1024 * 8192: + return None + if self.world_size == 6: + return None + + if envs.SGLANG_USE_1STAGE_ALLREDUCE.is_set(): + use_1stage_ar = envs.SGLANG_USE_1STAGE_ALLREDUCE.get() + else: + token_num = input_.numel() // K + use_1stage_ar = total_bytes <= 128 * 1024 + if ( + # Keep the default 128 KiB cutoff except for the measured TP=8 + # K=7168 graph-replay crossover. K=4096 remains on the default + # rule because token_num=8/16 still favored 1-stage there. + self.world_size == 8 + and 4096 < K <= 7168 + and token_num >= 8 + and use_1stage_ar + ): + use_1stage_ar = False + + try: + return ca_comm.custom_fused_ar_rms_per_group_quant( + input_, + residual_inp_, + weight_, + eps, + group_size, + use_1stage_ar, + emit_bf16=emit_bf16, + transpose_scale=transpose_scale, + ) + except Exception: + return None + + def _resolve_outplace_all_reduce_method( + self, + input_: torch.Tensor, + should_use_pymscclpp_allreduce: Optional[bool] = None, + ) -> Optional[str]: + if should_use_pymscclpp_allreduce is None: + should_use_pymscclpp_allreduce = ( + self.pymscclpp_comm is not None + and self.pymscclpp_comm.should_mscclpp_allreduce(input_) + ) + if ( + self.ca_comm is not None + and not self.ca_comm.disabled + and not should_use_pymscclpp_allreduce + and self.ca_comm.should_custom_ar(input_) + ): + return "ca" + # SSKJ-PIE: after ``ca`` -- the PCIe-IPC kernels are for hosts where + # no fabric-specific backend applies. They do not probe for NVLink, + # so on a host that has it this ordering keeps the faster backend + # in front of them. On this stack stock CustomAllreduce rejects + # every shape on TP4 PCIe (world_size==2-or-NVLink gate), so this + # branch is the first accelerated one actually reached. + if ( + self.pcie_ipc_comm is not None + and not self.pcie_ipc_comm.disabled + and not should_use_pymscclpp_allreduce + and self.pcie_ipc_comm.should_pcie_ipc_ar(input_) + ): + return "pcie_ipc" + if ( + self.qr_comm is not None + and not self.qr_comm.disabled + and self.qr_comm.should_quick_allreduce(input_) + ): + return "qr" + if self.pymscclpp_comm is not None and should_use_pymscclpp_allreduce: + return "pymscclpp" + if ( + self.torch_symm_mem_comm is not None + and not self.torch_symm_mem_comm.disabled + and self.torch_symm_mem_comm.should_torch_symm_mem_allreduce(input_) + ): + return "torch_symm_mem" + if is_in_tc_piecewise_cuda_graph() and self.pynccl_comm is not None: + # For piecewise cuda graph, we use pynccl outplace allreduce + return "pynccl" + return None + + def _can_use_flashinfer_allreduce(self, input_: torch.Tensor) -> bool: + if self._fi_workspace_hint is None: + return False + from sglang.srt.layers.flashinfer_comm_fusion import ( + can_use_flashinfer_allreduce, + ) + + return can_use_flashinfer_allreduce( + input_, + use_attn_tp_group=(self._fi_workspace_hint == "attn_tp"), + expected_world_size=self.world_size, + expected_group_key=(self.device_group, self.cpu_group), + ) + + def _flashinfer_allreduce(self, input_: torch.Tensor) -> torch.Tensor: + from sglang.srt.layers.flashinfer_comm_fusion import ( + flashinfer_allreduce as _flashinfer_allreduce_impl, + ) + + return _flashinfer_allreduce_impl( + input_, use_attn_tp_group=(self._fi_workspace_hint == "attn_tp") + ) + + def _all_reduce_out_place( + self, input_: torch.Tensor, outplace_all_reduce_method: str + ) -> torch.Tensor: + if outplace_all_reduce_method == "auto": + outplace_all_reduce_method = self._resolve_outplace_all_reduce_method( + input_ + ) + if outplace_all_reduce_method == "pymscclpp": + # pymscclpp reduces in place and returns its input; feed it a + # clone to honor the op's no-mutation contract. + input_ = input_.clone() + elif outplace_all_reduce_method is None: + # Force pynccl over the in-place fallback: it is graph-capture + # safe and NCCL is natively out-of-place, avoiding the clone + # the in-place fallback needs. + if self.pynccl_comm is not None: + outplace_all_reduce_method = "pynccl" + else: + out = input_.clone() + self._all_reduce_in_place(out) + return out + ca_comm = self.ca_comm + qr_comm = self.qr_comm + pymscclpp_comm = self.pymscclpp_comm + torch_symm_mem_comm = self.torch_symm_mem_comm + pynccl_comm = self.pynccl_comm + pcie_ipc_comm = self.pcie_ipc_comm + assert any([qr_comm, ca_comm, pymscclpp_comm, torch_symm_mem_comm, pynccl_comm, pcie_ipc_comm]) + if outplace_all_reduce_method == "ca": + assert not ca_comm.disabled + out = ca_comm.custom_all_reduce(input_) + elif outplace_all_reduce_method == "pcie_ipc": + assert not pcie_ipc_comm.disabled + out = pcie_ipc_comm.pcie_ipc_all_reduce(input_) + elif outplace_all_reduce_method == "qr": + assert not qr_comm.disabled + out = qr_comm.quick_all_reduce(input_) + elif outplace_all_reduce_method == "torch_symm_mem": + assert not torch_symm_mem_comm.disabled + out = torch_symm_mem_comm.all_reduce(input_) + elif outplace_all_reduce_method == "pymscclpp": + assert not pymscclpp_comm.disabled + out = pymscclpp_comm.all_reduce(input_) + elif outplace_all_reduce_method == "pynccl": + with pynccl_comm.change_state(enable=True): + out = pynccl_comm.outplace_all_reduce(input_) + assert out is not None + return out + + def _all_reduce_in_place(self, input_: torch.Tensor) -> None: + pynccl_comm = self.pynccl_comm + torch_symm_mem_comm = self.torch_symm_mem_comm + if pynccl_comm is not None and not pynccl_comm.disabled: + pynccl_comm.all_reduce(input_) + elif ( + torch_symm_mem_comm is not None + and not torch_symm_mem_comm.disabled + and torch_symm_mem_comm.should_torch_symm_mem_allreduce(input_) + ): + torch_symm_mem_comm.all_reduce(input_, out=input_) + else: + torch.distributed.all_reduce(input_, group=self.device_group) + + def reduce_scatter_along_dim( + self, input_: torch.Tensor, dim: int = -1 + ) -> torch.Tensor: + world_size = self.world_size + # Bypass the function if we are using only 1 GPU. + if world_size == 1: + return input_ + assert ( + -input_.dim() <= dim < input_.dim() + ), f"Invalid dim ({dim}) for input tensor with shape {input_.size()}" + + if dim < 0: + # Convert negative dim to positive. + dim += input_.dim() + + with self.use_symmetric_memory(self): + # TODO: make sure whether tensor layout affects nccl reduce_scatter + # Note: This will produce an incorrect answer if we don't make + # the input_tensor contiguous. Possible bug in reduce_scatter_tensor? + input_tensor = input_.movedim(dim, 0).contiguous() + + assert input_tensor.shape[0] % world_size == 0 + chunk_size = input_tensor.shape[0] // world_size + output_shape = (chunk_size,) + input_tensor.shape[1:] + + with self.use_symmetric_memory(self): + output_tensor = torch.empty( + output_shape, + dtype=input_tensor.dtype, + device=input_tensor.device, + ) + + self.reduce_scatter_tensor(output_tensor, input_tensor) + + # Reshape before returning + return output_tensor.movedim(0, dim) + + def _reduce_scatter_tensor( + self, + output: torch.Tensor, + input: torch.Tensor, + ) -> torch.Tensor: + pynccl_comm = self.pynccl_comm + if pynccl_comm is not None and ( + not pynccl_comm.disabled or self.is_symmetric_memory_enabled() + ): + self.debug_check_symmetric_mempool( + self, {"output": output, "input": input}, "reduce_scatter_tensor" + ) + with pynccl_comm.change_state(enable=True): + pynccl_comm.reduce_scatter(output, input) + else: + torch.distributed.reduce_scatter_tensor( + output, input, group=self.device_group + ) + return output + + def reduce_scatter_tensor(self, output: torch.Tensor, input: torch.Tensor): + if _is_npu or _is_cpu: + # TODO: add optimized reduce_scatter_tensor kernel for cpu + self._reduce_scatter_tensor(output, input) + elif self._maybe_aiter_reduce_scatter(output, input): + return + else: + reg_reduce_scatter_tensor(output, input, group_name=self.unique_name) + + def _has_aiter_custom_reduce_scatter(self) -> bool: + ca_comm = self.ca_comm + return ( + ca_comm is not None + and not getattr(ca_comm, "disabled", True) + and hasattr(ca_comm, "should_custom_ar") + and hasattr(ca_comm, "reduce_scatter") + ) + + def _maybe_aiter_reduce_scatter( + self, output: torch.Tensor, input: torch.Tensor + ) -> bool: + # Aiter custom reduce-scatter (ROCm). Mirrors `_all_gather_into_tensor`'s + # custom all-gather path: an equal-chunk (no variable sizes) reduce-scatter + # using the registered symmetric-memory buffers, which is faster than the + # generic RCCL kernel for the small, latency-bound decode collective. + # Gated by SGLANG_DP_USE_REDUCE_SCATTER. Falls back (returns False) + # for non-ROCm / unsupported shape/size/topology so the caller uses RCCL. + if not ( + is_hip() + and envs.SGLANG_DP_USE_REDUCE_SCATTER.get() + and self._has_aiter_custom_reduce_scatter() + and input.is_contiguous() + and output.is_contiguous() + and input.dtype in (torch.float32, torch.float16, torch.bfloat16) + ): + return False + ca_comm = self.ca_comm + # input is the full (pre-reduce) buffer; should_custom_ar bounds its size. + if not ca_comm.should_custom_ar(input): + return False + # Equal-chunk only: input rows must split evenly into world_size chunks + # matching the per-rank output rows. + if input.shape[0] != output.shape[0] * self.world_size: + return False + if getattr(ca_comm, "_IS_CAPTURING", False): + if torch.cuda.is_current_stream_capturing(): + if envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get(): + ca_comm.reduce_scatter(input, output, registered=False) + else: + ca_comm.reduce_scatter(input, output, registered=True) + elif is_in_tc_piecewise_cuda_graph(): + ca_comm.reduce_scatter(input, output, registered=False) + else: + # True CUDA graph warmup: avoid a different host collective. + output.zero_() + return True + ca_comm.reduce_scatter(input, output, registered=False) + return True + + def _all_to_all_single(self, output: torch.Tensor, input: torch.Tensor) -> None: + # pynccl path keeps the a2a exchange CUDA-graph-capturable (DCP a2a backend). + pynccl_comm = self.pynccl_comm + if pynccl_comm is not None and not pynccl_comm.disabled: + pynccl_comm.all_to_all_single(output, input) + else: + torch.distributed.all_to_all_single(output, input, group=self.device_group) + + def all_to_all_single(self, output: torch.Tensor, input: torch.Tensor): + if self.world_size == 1: + output.copy_(input) + return + reg_all_to_all_single(output, input, group_name=self.unique_name) + + def reduce_scatter( + self, + output: torch.Tensor, + input_list: List[torch.Tensor], + ) -> None: + # TODO(ch-wan): support other backends + torch.distributed.reduce_scatter(output, input_list, group=self.device_group) + return output + + def reduce_scatterv( + self, + input_: torch.Tensor, + output: Optional[torch.Tensor] = None, + sizes: Optional[List[int]] = None, + ) -> torch.Tensor: + world_size = self.world_size + pynccl_comm = self.pynccl_comm + + with pynccl_comm.change_state(enable=True): + assert ( + pynccl_comm is not None and not pynccl_comm.disabled + ), "pynccl is required for reduce_scatterv" + + if sizes is not None: + assert len(sizes) == world_size + assert input_.shape[0] == sum(sizes) + chunk_size = sizes[self.rank_in_group] + else: + assert input_.shape[0] % world_size == 0 + chunk_size = input_.shape[0] // world_size + output_shape = (chunk_size,) + input_.shape[1:] + + if output is None: + output = torch.empty( + output_shape, dtype=input_.dtype, device=input_.device + ) + else: + assert output.shape == output_shape + + pynccl_comm.reduce_scatter(output, input_, sizes=sizes) + return output + + def _all_gather_into_tensor(self, output: torch.Tensor, input: torch.Tensor): + # Aiter custom all-gather (ROCm). Set SGLANG_USE_AITER_AG=0 to disable. + # Aiter's should_custom_ag still owns shape/layout validation: + # 16B alignment, weak-contiguous, supported topology, and per-rank + # size <= max_size/(world*2). + # On a hit, writes directly into the caller's pre-allocated `output` via + # all_gather_reg during CUDA-graph capture, and all_gather_unreg + # under torch_memory_saver and other paths. + ca_comm = self.ca_comm + if ( + is_hip() + and envs.SGLANG_USE_AITER_AG.get() + and self._has_aiter_custom_all_gather() + and input.is_contiguous() + and output.is_contiguous() + and input.dtype in (torch.float32, torch.float16, torch.bfloat16) + and ca_comm.should_custom_ag(input) + ): + if getattr(ca_comm, "_IS_CAPTURING", False): + if torch.cuda.is_current_stream_capturing(): + if envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get(): + ca_comm.all_gather_unreg(input, out=output, dim=0) + else: + ca_comm.all_gather_reg(input, out=output, dim=0) + elif is_in_tc_piecewise_cuda_graph(): + ca_comm.all_gather_unreg(input, out=output, dim=0) + else: + # True CUDA graph warmup: avoid a different host collective. + output.zero_() + return + else: + ca_comm.all_gather_unreg(input, out=output, dim=0) + return + + pynccl_comm = self.pynccl_comm + if pynccl_comm is not None and ( + not pynccl_comm.disabled or self.is_symmetric_memory_enabled() + ): + self.debug_check_symmetric_mempool( + self, {"output": output}, "all_gather_into_tensor" + ) + with pynccl_comm.change_state(enable=True): + pynccl_comm.all_gather(output, input) + else: + torch.distributed.all_gather_into_tensor( + output, input, group=self.device_group + ) + + def _has_aiter_custom_all_gather(self) -> bool: + if self._deterministic_collectives_enabled(): + return False + ca_comm = self.ca_comm + return ( + ca_comm is not None + and not getattr(ca_comm, "disabled", True) + and hasattr(ca_comm, "should_custom_ag") + and hasattr(ca_comm, "all_gather_reg") + and hasattr(ca_comm, "all_gather_unreg") + ) + + @staticmethod + def _deterministic_collectives_enabled() -> bool: + if envs.SGLANG_USE_1STAGE_ALLREDUCE.is_set(): + return envs.SGLANG_USE_1STAGE_ALLREDUCE.get() + return envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get() + + def all_gather_into_tensor(self, output: torch.Tensor, input: torch.Tensor): + if _is_npu or _is_cpu: + # TODO: add optimized all_gather_into_tensor kernel for cpu + self._all_gather_into_tensor(output, input) + else: + # XPU and CUDA both go through reg_all_gather_into_tensor (custom_op) to + # stay opaque to Dynamo. Calling torch.distributed.all_gather_into_tensor + # directly causes Dynamo to rewrite it as _c10d_functional.all_gather_into_tensor + # + wait_tensor, which invokes sycl_event.wait() and breaks XPU graph capture. + reg_all_gather_into_tensor(output, input, group_name=self.unique_name) + + def all_gather( + self, + input_: torch.Tensor, + dim: int = -1, + output_tensor_list: Optional[List[torch.Tensor]] = None, + ) -> torch.Tensor: + world_size = self.world_size + # Bypass the function if we are using only 1 GPU. + if world_size == 1: + if output_tensor_list is not None: + logger.warning( + "Performing in-place all-gather with a group size of 1. " + "This may be unnecessary; consider bypassing it for better efficiency." + ) + output_tensor_list[0].copy_(input_) + return None + else: + return input_ + + if output_tensor_list is not None: + # TODO(ch-wan): support other backends + return torch.distributed.all_gather( + output_tensor_list, input_, group=self.device_group + ) + + assert ( + -input_.dim() <= dim < input_.dim() + ), f"Invalid dim ({dim}) for input tensor with shape {input_.size()}" + + # For HPUs, use HPU communicator. + hpu_comm = self.hpu_communicator + if hpu_comm is not None and not hpu_comm.disabled: + return hpu_comm.all_gather(input_, dim) + + # For NPUs, use NPU communicator. + npu_comm = self.npu_communicator + if npu_comm is not None and not npu_comm.disabled: + return npu_comm.all_gather(input_, dim) + + if dim < 0: + # Convert negative dim to positive. + dim += input_.dim() + input_size = input_.size() + # NOTE: we have to use concat-style all-gather here, + # stack-style all-gather has compatibility issues with + # torch.compile . see https://github.com/pytorch/pytorch/issues/138795 + output_size = (input_size[0] * world_size,) + input_size[1:] + # Allocate output tensor. + with self.use_symmetric_memory( + self, disabled=not self.is_allocation_symmetric() + ): + output_tensor = torch.empty( + output_size, dtype=input_.dtype, device=input_.device + ) + + # All-gather. + if input_.is_cpu: + if is_shm_available(input_.dtype, self.world_size, self.local_size): + return torch.ops.sgl_kernel.shm_allgather(input_, dim) + else: + torch.distributed.all_gather_into_tensor( + output_tensor, input_, group=self.device_group + ) + else: + self.all_gather_into_tensor(output_tensor, input_) + + # Reshape + output_tensor = output_tensor.reshape((world_size,) + input_size) + output_tensor = output_tensor.movedim(0, dim) + output_tensor = output_tensor.reshape( + input_size[:dim] + (world_size * input_size[dim],) + input_size[dim + 1 :] + ) + return output_tensor + + def all_gatherv( + self, + input_: Union[torch.Tensor, List[torch.Tensor]], + sizes: Optional[List[int]] = None, + output: Optional[torch.Tensor] = None, + ) -> Union[torch.Tensor, List[torch.Tensor]]: + """ + Supports varying sizes per rank and input tensor list. + `sizes`: a list of len(world_size) with the number of items per rank to gather. + `output`: optional pre-allocated destination buffer (single-tensor input only). + When given, NCCL writes the gathered result directly into it, avoiding an + extra output allocation + caller-side copy. + """ + world_size = self.world_size + pynccl_comm = self.pynccl_comm + + with pynccl_comm.change_state(enable=True): + assert ( + pynccl_comm is not None and not pynccl_comm.disabled + ), "pynccl is required for all_gatherv" + + def _all_gather_allocate_output( + input_: torch.Tensor, + sizes: Optional[List[int]] = None, + output: Optional[torch.Tensor] = None, + ): + input_size = input_.size() + if sizes is not None: + assert len(sizes) == world_size + assert input_.shape[0] == sizes[self.rank_in_group] + output_size = (sum(sizes),) + input_size[1:] + # 'sizes' is not needed if all inputs in the same group have the same shape + if all(s == sizes[0] for s in sizes): + sizes = None + else: + output_size = (input_size[0] * world_size,) + input_size[1:] + if output is not None: + assert tuple(output.shape) == tuple(output_size), ( + f"all_gatherv output buffer shape {tuple(output.shape)} " + f"!= expected {tuple(output_size)}" + ) + return output, sizes + # Allocate output tensor. + with self.use_symmetric_memory(self, disabled=sizes is not None): + output_tensor = torch.empty( + output_size, dtype=input_.dtype, device=input_.device + ) + return output_tensor, sizes + + single_input = isinstance(input_, torch.Tensor) + if single_input: + input_ = [input_] + elif output is not None: + raise ValueError("all_gatherv `output` requires a single-tensor input") + + output_list = [] + size_list = [] + for inp in input_: + output_tensor, s = _all_gather_allocate_output( + inp, sizes=sizes, output=output + ) + output_list.append(output_tensor) + size_list.append(s) + + pynccl_comm.group_start() + for i, inp in enumerate(input_): + pynccl_comm.all_gather(output_list[i], inp, sizes=size_list[i]) + pynccl_comm.group_end() + + return output_list + + def gather( + self, input_: torch.Tensor, dst: int = 0, dim: int = -1 + ) -> Optional[torch.Tensor]: + """ + NOTE: We assume that the input tensor is on the same device across + all the ranks. + NOTE: `dst` is the local rank of the destination rank. + """ + world_size = self.world_size + # Bypass the function if we are using only 1 GPU. + if world_size == 1: + return input_ + assert ( + -input_.dim() <= dim < input_.dim() + ), f"Invalid dim ({dim}) for input tensor with shape {input_.size()}" + if dim < 0: + # Convert negative dim to positive. + dim += input_.dim() + if self.xpu_communicator is not None and not self.xpu_communicator.disabled: + return self.xpu_communicator.gather(input_, self.rank_in_group, dst, dim) + # Allocate output tensor. + if self.rank_in_group == dst: + gather_list = [torch.empty_like(input_) for _ in range(world_size)] + else: + gather_list = None + # Gather. + torch.distributed.gather( + input_, gather_list, dst=self.ranks[dst], group=self.device_group + ) + if self.rank_in_group == dst: + output_tensor = torch.cat(gather_list, dim=dim) + else: + output_tensor = None + return output_tensor + + def broadcast(self, input_: torch.Tensor, src: int = 0): + """Broadcast the input tensor. + NOTE: `src` is the local rank of the source rank. + """ + assert src < self.world_size, f"Invalid src rank ({src})" + + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return input_ + # Broadcast. + torch.distributed.broadcast( + input_, src=self.ranks[src], group=self.device_group + ) + return input_ + + def broadcast_object(self, obj: Optional[Any] = None, src: int = 0): + """Broadcast the input object. + NOTE: `src` is the local rank of the source rank. + """ + assert src < self.world_size, f"Invalid src rank ({src})" + + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return obj + if self.mq_broadcaster is not None: + assert src == 0, "Message queue broadcaster only supports src=0" + return self.mq_broadcaster.broadcast_object(obj) + if self.rank_in_group == src: + torch.distributed.broadcast_object_list( + [obj], src=self.ranks[src], group=self.cpu_group + ) + return obj + else: + recv = [None] + torch.distributed.broadcast_object_list( + recv, src=self.ranks[src], group=self.cpu_group + ) + return recv[0] + + def broadcast_object_list( + self, obj_list: List[Any], src: int = 0, group: Optional[ProcessGroup] = None + ): + """Broadcast the input object list. + NOTE: `src` is the local rank of the source rank. + """ + assert src < self.world_size, f"Invalid src rank ({src})" + + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return obj_list + # Broadcast. + torch.distributed.broadcast_object_list( + obj_list, src=self.ranks[src], group=self.device_group + ) + return obj_list + + def all_gather_object(self, obj: Any) -> List[Any]: + objs = [None] * self.world_size + torch.distributed.all_gather_object(objs, obj, group=self.cpu_group) + return objs + + def send_object( + self, + obj: Any, + dst: int, + async_send: bool = False, + tag: int = 0, + ) -> List[P2PWork]: + """ + Send the input object list to the destination rank. + This function uses the CPU group for all communications. + + TODO: If you want to use GPU communication, please add a new argument (e.g., data_group, group), + use other functions (e.g., send), or implement a new function (e.g., send_object_device). + + NOTE: `dst` is the local rank of the destination rank. + """ + + assert dst < self.world_size, f"Invalid dst rank ({dst})" + assert dst != self.rank_in_group, ( + "Invalid destination rank. Destination rank is the same " + "as the current rank." + ) + send_func = torch.distributed.isend if async_send else torch.distributed.send + + # Serialize object to tensor and get the size as well + object_tensor = torch.frombuffer(pickle.dumps(obj), dtype=torch.uint8) + size_tensor = torch.tensor( + [object_tensor.numel()], dtype=torch.long, device="cpu" + ) + + # Send object size + p2p_work = [] + size_work = send_func( + size_tensor, + self.ranks[dst], + group=self.cpu_group, + tag=tag, + ) + if async_send: + p2p_work.append(P2PWork(size_work, size_tensor)) + + object_work = send_func( + object_tensor, + self.ranks[dst], + group=self.cpu_group, + tag=tag, + ) + if async_send: + p2p_work.append(P2PWork(object_work, object_tensor)) + + return p2p_work + + def recv_object( + self, + src: int, + tag: int = 0, + ) -> Any: + """Receive the input object list from the source rank.""" + """NOTE: `src` is the local rank of the source rank.""" + + assert src < self.world_size, f"Invalid src rank ({src})" + assert ( + src != self.rank_in_group + ), "Invalid source rank. Source rank is the same as the current rank." + + size_tensor = torch.empty(1, dtype=torch.long, device="cpu") + + # Receive object size + # We have to use irecv here to make it work for both isend and send. + work = torch.distributed.irecv( + size_tensor, src=self.ranks[src], group=self.cpu_group, tag=tag + ) + work.wait() + + # Tensor to receive serialized objects into. + object_tensor: Any = torch.empty( # type: ignore[call-overload] + size_tensor.item(), # type: ignore[arg-type] + dtype=torch.uint8, + device="cpu", + ) + + work = torch.distributed.irecv( + object_tensor, src=self.ranks[src], group=self.cpu_group, tag=tag + ) + work.wait() + + obj = pickle.loads(object_tensor.numpy()) + return obj + + def broadcast_tensor_dict( + self, + tensor_dict: Optional[Dict[str, Union[torch.Tensor, Any]]] = None, + src: int = 0, + group: Optional[ProcessGroup] = None, + metadata_group: Optional[ProcessGroup] = None, + ) -> Optional[Dict[str, Union[torch.Tensor, Any]]]: + """Broadcast the input tensor dictionary. + NOTE: `src` is the local rank of the source rank. + """ + # Bypass the function if we are using only 1 GPU. + if not torch.distributed.is_initialized() or self.world_size == 1: + return tensor_dict + + group = self.device_group + metadata_group = self.cpu_group + assert src < self.world_size, f"Invalid src rank ({src})" + + rank_in_group = self.rank_in_group + if rank_in_group == src: + metadata_list: List[Tuple[Any, Any]] = [] + assert isinstance( + tensor_dict, dict + ), f"Expecting a dictionary, got {type(tensor_dict)}" + metadata_list, tensor_list = _split_tensor_dict(tensor_dict) + # `metadata_list` lives in CPU memory. + # `broadcast_object_list` has serialization & deserialization, + # all happening on CPU. Therefore, we can use the CPU group. + self.broadcast_object(metadata_list, src=src) + async_handles = [] + for tensor in tensor_list: + if tensor.numel() == 0: + # Skip broadcasting empty tensors. + continue + if tensor.is_cpu: + # use metadata_group for CPU tensors + handle = torch.distributed.broadcast( + tensor, src=self.ranks[src], group=metadata_group, async_op=True + ) + else: + # use group for GPU tensors + handle = torch.distributed.broadcast( + tensor, src=self.ranks[src], group=group, async_op=True + ) + async_handles.append(handle) + for async_handle in async_handles: + async_handle.wait() + + else: + metadata_list = self.broadcast_object(None, src=src) + tensor_dict = {} + async_handles = [] + for key, value in metadata_list: + if isinstance(value, TensorMetadata): + tensor = torch.empty( + value.size, dtype=value.dtype, device=value.device + ) + if tensor.numel() == 0: + # Skip broadcasting empty tensors. + tensor_dict[key] = tensor + continue + if tensor.is_cpu: + # use metadata_group for CPU tensors + handle = torch.distributed.broadcast( + tensor, + src=self.ranks[src], + group=metadata_group, + async_op=True, + ) + else: + # use group for GPU tensors + handle = torch.distributed.broadcast( + tensor, src=self.ranks[src], group=group, async_op=True + ) + async_handles.append(handle) + tensor_dict[key] = tensor + else: + tensor_dict[key] = value + for async_handle in async_handles: + async_handle.wait() + return tensor_dict + + def send_tensor_dict( + self, + tensor_dict: Dict[str, Union[torch.Tensor, Any]], + dst: Optional[int] = None, + all_gather_group: Optional["GroupCoordinator"] = None, + async_send: bool = False, + ) -> Optional[List[P2PWork]]: + """Send the input tensor dictionary. + NOTE: `dst` is the local rank of the source rank. + """ + # Bypass the function if we are using only 1 GPU. + if self.world_size == 1: + return tensor_dict + + all_gather_size = 1 if all_gather_group is None else all_gather_group.world_size + all_gather_rank = ( + 0 if all_gather_group is None else all_gather_group.rank_in_group + ) + + group = self.device_group + metadata_group = self.cpu_group + + if dst is None: + dst = (self.rank_in_group + 1) % self.world_size + assert dst < self.world_size, f"Invalid dst rank ({dst})" + + assert isinstance( + tensor_dict, dict + ), f"Expecting a dictionary, got {type(tensor_dict)}" + metadata_list, tensor_list = _split_tensor_dict(tensor_dict) + # Note: While switching to Device-to-Device (D2D) would introduce an extra + # Device-to-Host (D2H) memory copy overhead for serialization, our benchmarks + # show better overall transmission performance with D2D due to: + # 1. Superior D2D transfer bandwidth + # 2. Ability to overlap send and recv operations + # Thus the net performance gain justifies this approach. + + send_func = torch.distributed.isend if async_send else torch.distributed.send + p2p_works = self.send_object(metadata_list, dst=dst, async_send=async_send) + + for tensor in tensor_list: + if tensor.numel() == 0: + # Skip sending empty tensors. + continue + + # send-allgather: send only a slice, then do allgather. + if all_gather_group is not None and tensor.numel() % all_gather_size == 0: + tensor = tensor.reshape(all_gather_size, -1)[all_gather_rank] + + comm_group = metadata_group if tensor.is_cpu else group + work = send_func(tensor, self.ranks[dst], group=comm_group) + if async_send: + p2p_works.append(P2PWork(work, tensor)) + return p2p_works + + def recv_tensor_dict( + self, + src: Optional[int] = None, + all_gather_group: Optional["GroupCoordinator"] = None, + ) -> Optional[Dict[str, Union[torch.Tensor, Any]]]: + """Recv the input tensor dictionary. + NOTE: `src` is the local rank of the source rank. + """ + # Bypass the function if we are using only 1 GPU. + if not torch.distributed.is_initialized() or self.world_size == 1: + return None + + all_gather_size = 1 if all_gather_group is None else all_gather_group.world_size + all_gather_rank = ( + 0 if all_gather_group is None else all_gather_group.rank_in_group + ) + + group = self.device_group + metadata_group = self.cpu_group + + if src is None: + src = (self.rank_in_group - 1) % self.world_size + assert src < self.world_size, f"Invalid src rank ({src})" + + recv_metadata_list = self.recv_object(src=src) + tensor_dict: Dict[str, Any] = {} + for key, value in recv_metadata_list: + if isinstance(value, TensorMetadata): + tensor = torch.empty(value.size, dtype=value.dtype, device=value.device) + if tensor.numel() == 0: + # Skip broadcasting empty tensors. + tensor_dict[key] = tensor + continue + + # send-allgather: send only a slice, then do allgather. + use_all_gather = ( + all_gather_group is not None + and tensor.numel() % all_gather_size == 0 + ) + + if use_all_gather: + orig_shape = tensor.shape + tensor = tensor.reshape(all_gather_size, -1)[all_gather_rank] + + # We have to use irecv here to make it work for both isend and send. + comm_group = metadata_group if tensor.is_cpu else group + work = torch.distributed.irecv( + tensor, src=self.ranks[src], group=comm_group + ) + work.wait() + + if use_all_gather: + tensor = all_gather_group.all_gather(tensor, dim=0) + tensor = tensor.reshape(orig_shape) + + tensor_dict[key] = tensor + else: + tensor_dict[key] = value + return tensor_dict + + def barrier(self): + """Barrier synchronization among the group. + NOTE: don't use `device_group` here! `barrier` in NCCL is + terrible because it is internally a broadcast operation with + secretly created GPU tensors. It is easy to mess up the current + device. Use the CPU group instead. + """ + torch.distributed.barrier(group=self.cpu_group) + + def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None: + """Sends a tensor to the destination rank in a non-blocking way""" + """NOTE: `dst` is the local rank of the destination rank.""" + if dst is None: + dst = (self.rank_in_group + 1) % self.world_size + + pynccl_comm = self.pynccl_comm + if pynccl_comm is not None and not pynccl_comm.disabled: + pynccl_comm.send(tensor, dst) + else: + torch.distributed.send(tensor, self.ranks[dst], self.device_group) + + def recv( + self, size: torch.Size, dtype: torch.dtype, src: Optional[int] = None + ) -> torch.Tensor: + """Receives a tensor from the source rank.""" + """NOTE: `src` is the local rank of the source rank.""" + if src is None: + src = (self.rank_in_group - 1) % self.world_size + + tensor = torch.empty(size, dtype=dtype, device=self.device) + pynccl_comm = self.pynccl_comm + if pynccl_comm is not None and not pynccl_comm.disabled: + pynccl_comm.recv(tensor, src) + else: + torch.distributed.recv(tensor, self.ranks[src], self.device_group) + return tensor + + def destroy(self): + if self.device_group is not None: + torch.distributed.destroy_process_group(self.device_group) + self.device_group = None + if self.cpu_group is not None: + torch.distributed.destroy_process_group(self.cpu_group) + self.cpu_group = None + if self.pynccl_comm is not None: + self.pynccl_comm = None + if self.pymscclpp_comm is not None: + self.pymscclpp_comm.destroy() + if self.ca_comm is not None: + self.ca_comm = None + if self.pcie_ipc_comm is not None: + try: + self.pcie_ipc_comm.destroy() + except Exception: + pass + self.pcie_ipc_comm = None + if self.mq_broadcaster is not None: + self.mq_broadcaster = None + + +_WORLD: Optional[GroupCoordinator] = None + + +def get_world_group() -> GroupCoordinator: + assert _WORLD is not None, "world group is not initialized" + return _WORLD + + +def init_world_group( + ranks: List[int], local_rank: int, backend: str, recovered_rank: bool = False +) -> GroupCoordinator: + return GroupCoordinator( + group_ranks=[ranks], + local_rank=local_rank, + torch_distributed_backend=backend, + use_pynccl=False, + use_pymscclpp=False, + use_custom_allreduce=False, + use_torch_symm_mem_all_reduce=False, + use_hpu_communicator=False, + use_xpu_communicator=False, + use_npu_communicator=False, + group_name="world", + recovered_rank=recovered_rank, + ) + + +def init_model_parallel_group( + group_ranks: List[List[int]], + local_rank: int, + backend: str, + use_pynccl: Optional[bool] = None, + use_custom_allreduce: Optional[bool] = None, + use_message_queue_broadcaster: bool = False, + group_name: Optional[str] = None, + use_mscclpp_allreduce: Optional[bool] = None, + use_torch_symm_mem_allreduce: Optional[bool] = None, + recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, +) -> GroupCoordinator: + if use_custom_allreduce is None: + use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE + if use_mscclpp_allreduce is None: + use_mscclpp_allreduce = _ENABLE_MSCCLPP_ALL_REDUCE + if use_torch_symm_mem_allreduce is None: + use_torch_symm_mem_allreduce = _ENABLE_TORCH_SYMM_MEM_ALL_REDUCE + return GroupCoordinator( + group_ranks=group_ranks, + local_rank=local_rank, + torch_distributed_backend=backend, + use_pynccl=( + not (_is_npu or _is_xpu or backend == "mooncake") + if use_pynccl is None + else use_pynccl + ), + use_pymscclpp=use_mscclpp_allreduce, + use_custom_allreduce=use_custom_allreduce, + use_torch_symm_mem_all_reduce=use_torch_symm_mem_allreduce, + use_hpu_communicator=True, + use_xpu_communicator=True, + use_npu_communicator=True, + use_message_queue_broadcaster=use_message_queue_broadcaster, + group_name=group_name, + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + +_TP: Optional[GroupCoordinator] = None +_ATTN_TP: Optional[GroupCoordinator] = None +_ATTN_CP: Optional[GroupCoordinator] = None +_ATTN_CP_OVERLAP: Optional[GroupCoordinator] = None +_DCP: Optional[GroupCoordinator] = None + +# duplicate GroupCoordinator for prefill in PD-Multiplexing +_PDMUX_PREFILL_TP_GROUP: Optional[GroupCoordinator] = None + +_ENABLE_PDMUX_P_TP: bool = False + + +def set_pdmux_status(enable_prefill_multiplexing: bool): + global _ENABLE_PDMUX_P_TP + _ENABLE_PDMUX_P_TP = enable_prefill_multiplexing + + +def get_tp_group() -> GroupCoordinator: + if _ENABLE_PDMUX_P_TP: + assert ( + _PDMUX_PREFILL_TP_GROUP is not None + ), "tensor model parallel group for PD-Multiplexing Prefill is not initialized" + return _PDMUX_PREFILL_TP_GROUP + assert _TP is not None, "tensor model parallel group is not initialized" + return _TP + + +def get_attn_tp_group() -> GroupCoordinator: + assert ( + _ATTN_TP is not None + ), "attention tensor model parallel group is not initialized" + return _ATTN_TP + + +def get_attn_cp_group() -> GroupCoordinator: + assert ( + _ATTN_CP is not None + ), "attention context model parallel group is not initialized" + return _ATTN_CP + + +def get_attn_cp_overlap_group() -> GroupCoordinator: + return _ATTN_CP_OVERLAP if _ATTN_CP_OVERLAP is not None else get_attn_cp_group() + + +def _init_attn_cp_overlap_group( + *, + world_size: int, + attn_cp_size: int, + attn_tp_size: int, + backend: Optional[str], + recovered_rank: bool, + rank_offset: int, + max_world_size: Optional[int], +) -> None: + """Second communicator over the attention CP ranks; RCCL deadlocks when one + communicator is driven from two streams at once.""" + global _ATTN_CP_OVERLAP + assert ( + _ATTN_CP_OVERLAP is None + ), "attention context parallel overlap group is already initialized" + if attn_cp_size <= 1: + return + + span = attn_tp_size * attn_cp_size + group_ranks = [ + list(range(base + i, base + i + span, attn_tp_size)) + for base in range(0, world_size, span) + for i in range(attn_tp_size) + ] + rank = torch.distributed.get_rank() + mine = next(ranks for ranks in group_ranks if rank in ranks) + assert mine == get_attn_cp_group().ranks, ( + f"attn_cp_overlap partition {mine} does not match attn_cp " + f"{get_attn_cp_group().ranks}; the two communicators must span the " + "same ranks or the overlapped collectives will not pair up" + ) + + _ATTN_CP_OVERLAP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_message_queue_broadcaster=False, + group_name="attn_cp_overlap", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + +def get_dcp_group_no_assert() -> Optional[GroupCoordinator]: + return _DCP + + +def get_dcp_group() -> GroupCoordinator: + assert _DCP is not None, "decode context parallel group is not initialized" + return _DCP + + +_MOE_DP: Optional[GroupCoordinator] = None +_MOE_EP: Optional[GroupCoordinator] = None +_MOE_TP: Optional[GroupCoordinator] = None + + +def get_moe_dp_group() -> GroupCoordinator: + assert _MOE_DP is not None, "moe data parallel group is not initialized" + return _MOE_DP + + +def get_moe_ep_group() -> GroupCoordinator: + assert _MOE_EP is not None, "expert model parallel group is not initialized" + return _MOE_EP + + +def get_moe_tp_group() -> GroupCoordinator: + assert _MOE_TP is not None, "expert model parallel group is not initialized" + return _MOE_TP + + +# kept for backward compatibility +get_tensor_model_parallel_group = get_tp_group + +_PP: Optional[GroupCoordinator] = None +_SELF_PP: Optional[GroupCoordinator] = None + + +def get_self_pp_group() -> GroupCoordinator: + assert _SELF_PP is not None, "self pipeline group is not initialized" + return _SELF_PP + + +def get_pp_group() -> GroupCoordinator: + assert _PP is not None, "pipeline model parallel group is not initialized" + return _PP + + +# kept for backward compatibility +get_pipeline_model_parallel_group = get_pp_group + + +def get_mooncake_transfer_engine(): + """ + Return the shared MooncakeTransferEngine if initialized in device_communicators, + else None. Used by disaggregation mooncake backend and mem_cache mooncake_store. + """ + from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( + get_mooncake_transfer_engine as _get_engine, + ) + + return _get_engine() + + +@contextmanager +def graph_capture(stream=None): + """ + `graph_capture` is a context manager which should surround the code that + is capturing the CUDA graph. Its main purpose is to ensure that the + some operations will be run after the graph is captured, before the graph + is replayed. It returns a `GraphCaptureContext` object which contains the + necessary data for the graph capture. Currently, it only contains the + stream that the graph capture is running on. This stream is set to the + current CUDA stream when the context manager is entered and reset to the + default stream when the context manager is exited. This is to ensure that + the graph capture is running on a separate stream from the default stream, + in order to explicitly distinguish the kernels to capture + from other kernels possibly launched on background in the default stream. + """ + with ( + get_tp_group().graph_capture(stream=stream) as context, + get_pp_group().graph_capture(context), + ): + with contextlib.ExitStack() as stack: + seen = {id(_TP), id(_PP)} + for group in (_DCP, _ATTN_TP, _MOE_EP, _MOE_TP): + if group is not None and id(group) not in seen: + seen.add(id(group)) + stack.enter_context(group.graph_capture(context)) + yield context + + +logger = logging.getLogger(__name__) + +_ENABLE_CUSTOM_ALL_REDUCE = True +_ENABLE_MSCCLPP_ALL_REDUCE = False +_ENABLE_TORCH_SYMM_MEM_ALL_REDUCE = False +_ENABLE_FLASHINFER_ALLREDUCE_ONLY = False + + +def set_custom_all_reduce(enable: bool): + global _ENABLE_CUSTOM_ALL_REDUCE + _ENABLE_CUSTOM_ALL_REDUCE = enable + + +def set_mscclpp_all_reduce(enable: bool): + global _ENABLE_MSCCLPP_ALL_REDUCE + _ENABLE_MSCCLPP_ALL_REDUCE = enable + + +def set_torch_symm_mem_all_reduce(enable: bool): + global _ENABLE_TORCH_SYMM_MEM_ALL_REDUCE + _ENABLE_TORCH_SYMM_MEM_ALL_REDUCE = enable + + +def set_flashinfer_allreduce_only(enable: bool): + global _ENABLE_FLASHINFER_ALLREDUCE_ONLY + _ENABLE_FLASHINFER_ALLREDUCE_ONLY = enable + + +def _tag_groups_for_flashinfer_allreduce_only(): + """Stamp _fi_workspace_hint on the group coordinators that own a FlashInfer + fusion workspace, so all_reduce() can dispatch to flashinfer_allreduce() + without touching the call sites. + + Only two workspaces exist (see ``_get_workspace_manager``): one for + attention TP and one for MoE. A group may only be tagged for the workspace + that was rendezvoused on its own peers -- reducing over a workspace built + for a different set of peers silently returns wrong data. + + - ``_TP`` is deliberately absent: it *is* ``_ATTN_TP`` when + ``attn_tp_size == tp_size``, and a strict superset of it otherwise (DP + attention), where the attention workspace addresses the wrong peers. + - The MoE workspace rendezvouses on the EP group when ``moe_ep_size > 1`` + and on the MoE-TP group otherwise, so exactly one of ``_MOE_EP`` / + ``_MOE_TP`` is eligible. Tagging both makes a MoE-TP allreduce reduce + across the EP peers under hybrid EP+TP (e.g. tp=4, ep=2). + """ + if not _ENABLE_FLASHINFER_ALLREDUCE_ONLY: + return + + moe_group = _MOE_EP if (_MOE_EP is not None and _MOE_EP.world_size > 1) else _MOE_TP + # Attention is tagged last on purpose: when a coordinator backs both roles + # (e.g. _ATTN_TP is _MOE_EP is _TP at tp=4, ep=4) either workspace spans the + # same peers and is correct, so we just pick one deterministically. + for group, hint in ((moe_group, "moe"), (_ATTN_TP, "attn_tp")): + if group is not None: + group._fi_workspace_hint = hint + + +# TODO: refactor in-tree platforms to get rid of this wrapper +def get_default_distributed_backend(device: str) -> str: + # We deliberately go through ``platforms.current_platform`` (rather than + # ``from ... import current_platform``) so each call resolves through the + # platforms package's lazy ``__getattr__`` and picks up runtime overrides + # of ``_current_platform`` (e.g. in tests). + if device == platforms.current_platform.device_type: + return platforms.current_platform.get_torch_distributed_backend_str() + return _DEVICE_TO_DISTRIBUTED_BACKEND.get(device, "gloo") + + +def _create_global_tcp_store( + rank: int, + world_size: int, + dist_init_addr: Optional[str] = None, + allow_dynamic_membership: bool = False, +) -> None: + """Create a global TCPStore for coordination across ranks. + + This function creates a TCPStore that all ranks can use for coordination + (e.g., for NIXL buffer setup). + """ + from torch.distributed import TCPStore + + base_store_port = envs.SGLANG_TCP_STORE_PORT.get() + + master_ip = os.environ.get("MASTER_ADDR") + if not master_ip and allow_dynamic_membership and dist_init_addr: + addr = dist_init_addr + if addr.startswith("tcp://"): + addr = addr[len("tcp://") :] + master_ip = addr.rsplit(":", 1)[0] + if not master_ip: + logger.warning( + "Could not determine master IP for global TCPStore. " + "Broadcasting from rank 0 to all ranks." + ) + if rank == 0: + master_ip = get_local_ip_auto() + ip_list = [master_ip] + else: + ip_list = [None] + torch.distributed.broadcast_object_list(ip_list, src=0) + master_ip = ip_list[0] + + try: + if allow_dynamic_membership: + is_master = rank == 0 + tcp_store = TCPStore( + host_name=master_ip, + port=base_store_port, + is_master=is_master, + wait_for_workers=False, + ) + else: + is_master = rank == 0 + tcp_store = TCPStore( + host_name=master_ip, + port=base_store_port, + world_size=world_size, + is_master=is_master, + ) + set_global_tcp_store(tcp_store) + logger.info( + "Created global TCPStore at %s:%d (rank=%d, is_master=%s)", + master_ip, + base_store_port, + rank, + is_master, + ) + except Exception as e: + logger.warning( + "Failed to create global TCPStore at %s:%d: %s. " + "Components requiring TCPStore (like NIXL) may not work.", + master_ip, + base_store_port, + e, + ) + + +def init_distributed_environment( + world_size: int = -1, + rank: int = -1, + distributed_init_method: str = "env://", + local_rank: int = -1, + backend: str = "nccl", + timeout: Optional[int] = None, + moe_a2a_backend: Optional[str] = None, + recovered_rank: bool = False, + max_world_size: Optional[int] = None, +): + logger.debug( + "world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s", + world_size, + rank, + local_rank, + distributed_init_method, + backend, + ) + if not torch.distributed.is_initialized(): + global _MODEL_PARALLEL_GROUP_TIMEOUT + assert distributed_init_method is not None, ( + "distributed_init_method must be provided when initializing " + "distributed environment" + ) + if timeout is not None: + assert isinstance(timeout, (int)), "timeout must be a number" + assert timeout > 0, "timeout must be positive" + timeout = timedelta(seconds=timeout) + + _MODEL_PARALLEL_GROUP_TIMEOUT = timeout + + if backend == "mooncake": + from mooncake.pg import MooncakeBackendOptions + + use_max_ws = max_world_size and max_world_size > world_size + ar_size = max_world_size if use_max_ws else world_size + active_ranks = torch.zeros(ar_size, dtype=torch.int32, device="cuda") + active_ranks[:world_size] = 1 + if use_max_ws: + pg_options = MooncakeBackendOptions( + active_ranks, recovered_rank, max_world_size + ) + else: + pg_options = MooncakeBackendOptions(active_ranks, recovered_rank) + else: + pg_options = get_torch_distributed_pg_options() + + # this backend is used for WORLD + torch.distributed.init_process_group( + backend=backend, + init_method=distributed_init_method, + world_size=world_size, + rank=rank, + timeout=timeout, + pg_options=pg_options, + ) + + # Create a global TCPStore for coordination (used by NIXL) + if moe_a2a_backend == "nixl": + _create_global_tcp_store( + rank, + world_size, + dist_init_addr=distributed_init_method, + allow_dynamic_membership=( + recovered_rank + or (max_world_size is not None and max_world_size > world_size) + ), + ) + + # set the local rank + # local_rank is not available in torch ProcessGroup, + # see https://github.com/pytorch/pytorch/issues/122816 + if local_rank == -1: + # local rank not set, this usually happens in single-node + # setting, where we can use rank as local rank + if distributed_init_method == "env://": + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + else: + local_rank = rank + global _WORLD + if _WORLD is None: + ranks = list(range(torch.distributed.get_world_size())) + _WORLD = init_world_group( + ranks, local_rank, backend, recovered_rank=recovered_rank + ) + else: + assert ( + _WORLD.world_size == torch.distributed.get_world_size() + ), "world group already initialized with a different world size" + + +def initialize_model_parallel( + tensor_model_parallel_size: int = 1, + expert_model_parallel_size: int = 1, + pipeline_model_parallel_size: int = 1, + attention_data_parallel_size: int = 1, + attention_context_model_parallel_size: int = 1, + moe_data_model_parallel_size: int = 1, + decode_context_parallel_size: int = 1, + backend: Optional[str] = None, + duplicate_tp_group: bool = False, + duplicate_attn_cp_group: bool = False, + enable_symm_mem: bool = False, + recovered_rank: bool = False, + rank_offset: int = 0, + max_world_size: Optional[int] = None, +) -> None: + """ + Initialize model parallel groups. + + Arguments: + tensor_model_parallel_size: number of GPUs used for tensor model + parallelism. + expert_model_parallel_size: number of GPUs used for expert model + parallelism. + pipeline_model_parallel_size: number of GPUs used for pipeline model + parallelism. + attention_data_parallel_size: number of GPUs used for attention data + parallelism. + attention_context_model_parallel_size: number of GPUs used for attention context + parallelism. + moe_data_model_parallel_size: number of GPUs used for moe data + parallelism. + decode_context_parallel_size: number of GPUs used for decode context + parallelism, which splits the KV cache across GPUs within each + tensor-parallel group during decoding. Must be a divisor of + tensor_model_parallel_size and is currently only supported on the + AMD HIP platform. + + Let's say we have a total of 8 GPUs denoted by g0 ... g7 and we + use 2 GPUs to parallelize the model tensor, and 4 GPUs to parallelize + the model pipeline. The present function will + create 4 tensor model-parallel groups and 2 pipeline model-parallel groups: + 4 tensor model-parallel groups: + [g0, g1], [g2, g3], [g4, g5], [g6, g7] + 2 pipeline model-parallel groups: + [g0, g2, g4, g6], [g1, g3, g5, g7] + + Let's say we use 2 GPUs for attention context parallelism (attn_cp_size=2) and 4 GPUs for + attention tensor parallelism (attn_tp_size=4). As for MoE part, we use 2 GPUs for moe data + parallelism (moe_dp_size=2) and 4 GPUs for moe expert parallelism (moe_ep_size=4). Note that + this implies tensor_model_parallel_size=8 (attn_tp_size = tp_size // attn_cp_size // + attn_dp_size), so all 8 GPUs form a single tensor model-parallel group. The present + function will create the following groups: + 1 tensor model-parallel group: + [g0, g1, g2, g3, g4, g5, g6, g7] + 4 attention context-parallel groups: + [g0, g4], [g1, g5], [g2, g6], [g3, g7] + 2 attention tensor-parallel groups: + [g0, g1, g2, g3], [g4, g5, g6, g7] + 2 moe expert-parallel groups: + [g0, g1, g2, g3], [g4, g5, g6, g7] + 4 moe data-parallel groups: + [g0, g4], [g1, g5], [g2, g6], [g3, g7] + + Note that for efficiency, the caller should make sure adjacent ranks + are on the same DGX box. For example if we are using 2 DGX-1 boxes + with a total of 16 GPUs, rank 0 to 7 belong to the first box and + ranks 8 to 15 belong to the second box. + """ + # Get world size and rank. Ensure some consistencies. + assert torch.distributed.is_initialized() + backend = backend or torch.distributed.get_backend(get_world_group().device_group) + + # Joiners construct their local TP/PP layout in global rank space. + world_size: int = ( + tensor_model_parallel_size * pipeline_model_parallel_size + if recovered_rank + else torch.distributed.get_world_size() + ) + + if world_size != tensor_model_parallel_size * pipeline_model_parallel_size: + raise RuntimeError( + f"world_size ({world_size}) is not equal to " + f"tensor_model_parallel_size ({tensor_model_parallel_size}) x " + f"pipeline_model_parallel_size ({pipeline_model_parallel_size})" + ) + if decode_context_parallel_size < 1: + raise RuntimeError( + f"decode_context_parallel_size ({decode_context_parallel_size}) must be >= 1" + ) + if decode_context_parallel_size > 1 and not (is_hip() or is_cuda()): + raise RuntimeError( + "Decode context parallel (decode_context_parallel_size > 1) is " + "currently only supported on the AMD HIP platform or CUDA platform, but got " + f"decode_context_parallel_size ({decode_context_parallel_size}) " + "on a non-HIP or non-CUDA platform." + ) + if tensor_model_parallel_size % decode_context_parallel_size != 0: + raise RuntimeError( + f"tensor_model_parallel_size ({tensor_model_parallel_size}) must be divisible by " + f"decode_context_parallel_size ({decode_context_parallel_size})" + ) + + # Build the tensor model-parallel groups. + num_tensor_model_parallel_groups: int = world_size // tensor_model_parallel_size + global _TP + assert _TP is None, "tensor model parallel group is already initialized" + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + ranks = list( + range( + tp_group_idx * tensor_model_parallel_size, + (tp_group_idx + 1) * tensor_model_parallel_size, + ) + ) + group_ranks.append(ranks) + + # message queue broadcaster is only used in tensor model parallel group + _TP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), + group_name="tp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + if duplicate_tp_group: + global _PDMUX_PREFILL_TP_GROUP + assert ( + _PDMUX_PREFILL_TP_GROUP is None + ), "tensor model parallel group for PD-Multiplexing Prefill is already initialized" + _PDMUX_PREFILL_TP_GROUP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), + group_name="pdmux_prefill_tp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + if _TP.pynccl_comm: + _TP.pynccl_comm.disabled = False + _PDMUX_PREFILL_TP_GROUP.pynccl_comm.disabled = False + + # Build decode context-parallel groups inside each TP group only when DCP is enabled. + global _DCP + assert _DCP is None, "decode context parallel group is already initialized" + if decode_context_parallel_size > 1: + dcp_group_ranks = [] + for tp_group in group_ranks: + for start in range(0, len(tp_group), decode_context_parallel_size): + dcp_group_ranks.append( + tp_group[start : start + decode_context_parallel_size] + ) + _DCP = init_model_parallel_group( + dcp_group_ranks, + get_world_group().local_rank, + backend, + use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), + group_name="dcp", + recovered_rank=recovered_rank, + ) + if get_tensor_model_parallel_rank() == 0: + logger.info( + f"DCP enabled, dcp_size={decode_context_parallel_size}, tp_size={tensor_model_parallel_size}" + ) + + attn_dp_size = attention_data_parallel_size + attn_cp_size = attention_context_model_parallel_size + # The groups below are built at these numbers, and the same dict is stamped + # once they exist. + derived_widths = derive_parallel_widths( + tp_size=tensor_model_parallel_size, + attn_cp_size=attn_cp_size, + attn_dp_size=attn_dp_size, + moe_ep_size=expert_model_parallel_size, + moe_dp_size=moe_data_model_parallel_size, + dcp_size=decode_context_parallel_size, + dcp_enabled=_DCP is not None, + ) + attn_tp_size = derived_widths["attn_tp_size"] + + global _ATTN_CP + assert ( + _ATTN_CP is None + ), "attention context model parallel group is already initialized" + if attn_cp_size == tensor_model_parallel_size: + _ATTN_CP = _TP + else: + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + for dp_idx in range(attn_dp_size): + for attn_tp_idx in range(attn_tp_size): + st = ( + tp_group_idx * tensor_model_parallel_size + + dp_idx * attn_tp_size * attn_cp_size + + attn_tp_idx + ) + en = ( + tp_group_idx * tensor_model_parallel_size + + (dp_idx + 1) * attn_tp_size * attn_cp_size + + attn_tp_idx + ) + ranks = list(range(st, en, attn_tp_size)) + group_ranks.append(ranks) + _ATTN_CP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), + group_name="attn_cp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + if duplicate_attn_cp_group and is_hip(): + _init_attn_cp_overlap_group( + world_size=world_size, + attn_cp_size=attn_cp_size, + attn_tp_size=attn_tp_size, + backend=backend, + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + from sglang.srt.layers.sampler import SYNC_TOKEN_IDS_ACROSS_TP + + global _ATTN_TP + assert ( + _ATTN_TP is None + ), "attention tensor model parallel group is already initialized" + if attn_tp_size == tensor_model_parallel_size: + _ATTN_TP = _TP + else: + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + for cp_dp_combined_idx in range(attn_cp_size * attn_dp_size): + st = ( + tp_group_idx * tensor_model_parallel_size + + cp_dp_combined_idx * attn_tp_size + ) + en = ( + tp_group_idx * tensor_model_parallel_size + + (cp_dp_combined_idx + 1) * attn_tp_size + ) + ranks = list(range(st, en)) + group_ranks.append(ranks) + + _ATTN_TP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_pynccl=SYNC_TOKEN_IDS_ACROSS_TP or enable_symm_mem, + use_custom_allreduce=False, + use_torch_symm_mem_allreduce=False, + use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), + group_name="attention_tp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + moe_ep_size = expert_model_parallel_size + moe_dp_size = moe_data_model_parallel_size + moe_tp_size = derived_widths["moe_tp_size"] + + global _MOE_DP + assert _MOE_DP is None, "moe data parallel group is already initialized" + if attn_cp_size > moe_dp_size: + # When moe_dp_size < attn_cp_size, CP ranks must share tokens before MoE. + # The MOE_DP group includes these CP partners, so the existing DP + # allgather/scatter handles the token sharing. + _MOE_DP = _ATTN_CP + elif moe_dp_size == tensor_model_parallel_size: + _MOE_DP = _TP + else: + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + for tp_ep_combined_idx in range(moe_tp_size * moe_ep_size): + st = tp_group_idx * tensor_model_parallel_size + tp_ep_combined_idx + en = ( + tp_group_idx + 1 + ) * tensor_model_parallel_size + tp_ep_combined_idx + ranks = list(range(st, en, moe_tp_size * moe_ep_size)) + group_ranks.append(ranks) + _MOE_DP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + group_name="moe_dp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + global _MOE_EP + assert _MOE_EP is None, "expert model parallel group is already initialized" + # NPU requires a standalone group for MOE expert parallelism + if moe_ep_size == tensor_model_parallel_size and not _is_npu: + _MOE_EP = _TP + else: + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + for moe_dp_idx in range(moe_dp_size): + for moe_tp_idx in range(moe_tp_size): + st = ( + tp_group_idx * tensor_model_parallel_size + + moe_dp_idx * moe_ep_size * moe_tp_size + + moe_tp_idx + ) + en = st + moe_ep_size * moe_tp_size + ranks = list(range(st, en, moe_tp_size)) + group_ranks.append(ranks) + _MOE_EP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_pynccl=False, + use_custom_allreduce=False, + group_name="moe_ep", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + global _MOE_TP + assert _MOE_TP is None, "expert model parallel group is already initialized" + if moe_tp_size == tensor_model_parallel_size: + _MOE_TP = _TP + else: + group_ranks = [] + for tp_group_idx in range(num_tensor_model_parallel_groups): + for ep_dp_combined_idx in range(moe_ep_size * moe_dp_size): + st = ( + tp_group_idx * tensor_model_parallel_size + + ep_dp_combined_idx * moe_tp_size + ) + en = ( + tp_group_idx * tensor_model_parallel_size + + (ep_dp_combined_idx + 1) * moe_tp_size + ) + ranks = list(range(st, en)) + group_ranks.append(ranks) + _MOE_TP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_pynccl=False, + use_custom_allreduce=False, + group_name="moe_tp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + # Build the pipeline model-parallel groups. + num_pipeline_model_parallel_groups: int = world_size // pipeline_model_parallel_size + global _PP + assert _PP is None, "pipeline model parallel group is already initialized" + group_ranks = [] + for pp_group_idx in range(num_pipeline_model_parallel_groups): + ranks = list( + range(pp_group_idx, world_size, num_pipeline_model_parallel_groups) + ) + group_ranks.append(ranks) + # pipeline parallel does not need custom allreduce + _PP = init_model_parallel_group( + group_ranks, + get_world_group().local_rank, + backend, + use_custom_allreduce=False, + group_name="pp", + recovered_rank=recovered_rank, + rank_offset=rank_offset, + max_world_size=max_world_size, + ) + + # The one-layer draft uses a singleton PP group; every rank creates all groups + # because new_group is collective. + global _SELF_PP + if _SELF_PP is None: + _SELF_PP = init_model_parallel_group( + [[r] for r in range(world_size)], + get_world_group().local_rank, + backend, + use_custom_allreduce=False, + group_name="self_pp", + ) + + get_parallel().stamp_derived_widths(**derived_widths) + + +def create_custom_parallel_group( + group_ranks: List[int], backend: str = "gloo" +) -> Optional[torch.distributed.ProcessGroup]: + """ + Create a custom parallel group based on the provided ranks. + + Args: + group_ranks: The list of ranks that the CURRENT process wants to join. + (e.g., Rank 0 passes [0...7], Rank 8 passes [8...15]) + backend: The communication backend (default: "gloo"). + + Returns: + The ProcessGroup if the current rank is in group_ranks, else None. + """ + assert torch.distributed.is_initialized() + + world_size = torch.distributed.get_world_size() + rank = torch.distributed.get_rank() + + local_config = sorted(list(set(group_ranks))) + gathered_configs = [None for _ in range(world_size)] + + torch.distributed.all_gather_object(gathered_configs, local_config) + + unique_groups = [] + seen_signatures = set() + + for config in gathered_configs: + config_tuple = tuple(config) + if config_tuple not in seen_signatures: + seen_signatures.add(config_tuple) + unique_groups.append(list(config_tuple)) + + unique_groups.sort(key=lambda x: x[0]) + + my_new_group = None + + for g_ranks in unique_groups: + group = torch.distributed.new_group(ranks=g_ranks, backend=backend) + + if set(g_ranks) == set(local_config): + my_new_group = group + logger.debug( + f"Rank {rank} successfully created/joined custom group: {g_ranks}" + ) + + return my_new_group + + +def ensure_model_parallel_initialized( + tensor_model_parallel_size: int, + expert_model_parallel_size: int, + pipeline_model_parallel_size: int, + decode_context_parallel_size: int = 1, + backend: Optional[str] = None, +) -> None: + """Helper to initialize model parallel groups if they are not initialized, + or ensure tensor-parallel and pipeline-parallel sizes are equal to expected + values if the model parallel groups are initialized. + """ + backend = backend or torch.distributed.get_backend(get_world_group().device_group) + if not model_parallel_is_initialized(): + initialize_model_parallel( + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + decode_context_parallel_size=decode_context_parallel_size, + backend=backend, + ) + return + + assert get_tensor_model_parallel_world_size() == tensor_model_parallel_size, ( + "tensor parallel group already initialized, but of unexpected size: " + f"{get_tensor_model_parallel_world_size()=} vs. " + f"{tensor_model_parallel_size=}" + ) + pp_world_size = get_pp_group().world_size + assert pp_world_size == pipeline_model_parallel_size, ( + "pipeline parallel group already initialized, but of unexpected size: " + f"{pp_world_size=} vs. " + f"{pipeline_model_parallel_size=}" + ) + if decode_context_parallel_size > 1: + dcp_world_size = get_dcp_group().world_size + assert ( + dcp_world_size == decode_context_parallel_size + ), f"decode context parallel group already initialized, but of unexpected size: {dcp_world_size=} {decode_context_parallel_size=}" + + +def model_parallel_is_initialized(): + """Check if tensor and pipeline parallel groups are initialized.""" + return _TP is not None and _PP is not None + + +_TP_STATE_PATCHED = False +_PP_STATE_PATCHED = False + + +@contextmanager +def patch_pipeline_parallel_group(pp_group: GroupCoordinator): + """Patch the pp group temporarily until this function ends. + + This method is for draft workers of speculative decoding, whose model does not + span pipeline stages and must not read the target's pp topology. + """ + global _PP_STATE_PATCHED + assert not _PP_STATE_PATCHED, "Should not call when it's already patched" + + _PP_STATE_PATCHED = True + old_pp_group = get_pp_group() + global _PP + _PP = pp_group + try: + yield + finally: + _PP_STATE_PATCHED = False + _PP = old_pp_group + + +@contextmanager +def patch_tensor_parallel_group(tp_group: GroupCoordinator): + """Run under a different tensor-parallel group until this scope ends. + + This is for draft workers of speculative decoding, which run the draft model + at the target's attention-TP width rather than its global TP width. + + The scope replaces both the module global that ``get_tp_group()`` reads and + the three members the runtime context answers with. + + Args: + tp_group (GroupCoordinator): the tp group coordinator + """ + + global _TP_STATE_PATCHED + assert not _TP_STATE_PATCHED, "Should not call when it's already patched" + + _TP_STATE_PATCHED = True + old_tp_group = get_tp_group() + global _TP + _TP = tp_group + try: + with get_parallel().override( + tp_size=tp_group.world_size, + tp_rank=tp_group.rank_in_group, + tp_group=tp_group, + ): + yield + finally: + _TP_STATE_PATCHED = False + _TP = old_tp_group + + +def get_world_size(): + """Return world size for the world group.""" + return get_world_group().world_size + + +def get_world_rank(): + """Return my rank for the world group.""" + return get_world_group().rank_in_group + + +def get_tensor_model_parallel_world_size(): + """Return world size for the tensor model parallel group.""" + return get_tp_group().world_size + + +def get_dcp_world_size(): + return get_dcp_group().world_size + + +def get_dcp_rank(): + return get_dcp_group().rank_in_group + + +def get_tensor_model_parallel_rank(): + """Return my rank for the tensor model parallel group.""" + return get_tp_group().rank_in_group + + +# ATTN_TP +def get_attn_tensor_model_parallel_world_size(): + """Return world size for the attention tensor model parallel group.""" + return get_attn_tp_group().world_size + + +def get_attn_tensor_model_parallel_rank(): + """Return my rank for the attention tensor model parallel group.""" + return get_attn_tp_group().rank_in_group + + +# ATTN_CP +def get_attn_context_model_parallel_world_size(): + """Return world size for the attention context model parallel group.""" + return get_attn_cp_group().world_size + + +def get_attn_context_model_parallel_rank(): + """Return my rank for the attention context model parallel group.""" + return get_attn_cp_group().rank_in_group + + +def get_pipeline_model_parallel_world_size(): + """Return world size for the pipeline model parallel group.""" + return get_pp_group().world_size + + +def get_pipeline_model_parallel_rank(): + """Return my rank for the pipeline model parallel group.""" + return get_pp_group().rank_in_group + + +# MOE_DP +def get_moe_data_parallel_world_size(): + """Return world size for the moe data parallel group.""" + return get_moe_dp_group().world_size + + +def get_moe_data_parallel_rank(): + """Return my rank for the moe data parallel group.""" + return get_moe_dp_group().rank_in_group + + +# MOE_EP +def get_moe_expert_parallel_world_size(): + """Return world size for the moe expert parallel group.""" + return get_moe_ep_group().world_size + + +def get_moe_expert_parallel_rank(): + """Return my rank for the moe expert parallel group.""" + return get_moe_ep_group().rank_in_group + + +# MOE_TP +def get_moe_tensor_parallel_world_size(): + """Return world size for the moe tensor parallel group.""" + return get_moe_tp_group().world_size + + +def get_moe_tensor_parallel_rank(): + """Return my rank for the moe tensor parallel group.""" + return get_moe_tp_group().rank_in_group + + +def destroy_model_parallel(): + """Set the groups to none and destroy them.""" + get_parallel().clear_derived_widths() + dwdp_mgr = get_global_dwdp_manager() + if dwdp_mgr is not None: + dwdp_mgr.cleanup() + set_global_dwdp_manager(None) + + global _TP + if _TP: + _TP.destroy() + _TP = None + + global _PP + if _PP: + _PP.destroy() + _PP = None + + global _DCP + if _DCP: + _DCP.destroy() + _DCP = None + + global _MOE_EP + if _MOE_EP: + _MOE_EP.destroy() + _MOE_EP = None + + global _MOE_TP + if _MOE_TP: + _MOE_TP.destroy() + _MOE_TP = None + + global _ATTN_CP + global _ATTN_CP_OVERLAP + global _MOE_DP + # Destroy _MOE_DP before _ATTN_CP since it may alias _ATTN_CP. + # Only destroy if not aliasing another group. + if _MOE_DP and _MOE_DP is not _ATTN_CP and _MOE_DP is not _TP: + _MOE_DP.destroy() + _MOE_DP = None + if _ATTN_CP_OVERLAP: + _ATTN_CP_OVERLAP.destroy() + _ATTN_CP_OVERLAP = None + if _ATTN_CP: + _ATTN_CP.destroy() + _ATTN_CP = None + + global _ATTN_TP + if _ATTN_TP: + _ATTN_TP.destroy() + _ATTN_TP = None + + global _PDMUX_PREFILL_TP_GROUP + if _PDMUX_PREFILL_TP_GROUP: # type: ignore[union-attr] + _PDMUX_PREFILL_TP_GROUP.destroy() + _PDMUX_PREFILL_TP_GROUP = None + + +def destroy_distributed_environment(): + global _WORLD, _MODEL_PARALLEL_GROUP_TIMEOUT + if _WORLD: + _WORLD.destroy() + _WORLD = None + _MODEL_PARALLEL_GROUP_TIMEOUT = None + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() + + +def cleanup_dist_env_and_memory(shutdown_ray: bool = False): + destroy_model_parallel() + destroy_distributed_environment() + with contextlib.suppress(AssertionError): + torch.distributed.destroy_process_group() + if shutdown_ray: + import ray # Lazy import Ray + + ray.shutdown() + gc.collect() + if not _is_cpu: + if hasattr(torch, "cuda") and torch.cuda.is_available(): + torch.cuda.empty_cache() + if hasattr(torch._C, "_host_emptyCache"): + torch._C._host_emptyCache() + else: + logger.warning( + "torch._C._host_emptyCache() only available in Pytorch >=2.5" + ) + elif hasattr(torch, "xpu") and torch.xpu.is_available(): + torch.xpu.empty_cache() + elif hasattr(torch, "npu") and torch.npu.is_available(): + torch.npu.empty_cache() + elif hasattr(torch, "musa") and torch.musa.is_available(): + torch.musa.empty_cache() + + +def in_the_same_node_as(pg: ProcessGroup, source_rank: int = 0) -> List[bool]: + """ + This is a collective operation that returns if each rank is in the same node + as the source rank. It tests if processes are attached to the same + memory system (shared access to shared memory). + """ + assert ( + torch.distributed.get_backend(pg) != torch.distributed.Backend.NCCL + ), "in_the_same_node_as should be tested with a non-NCCL group." + # local rank inside the group + rank = torch.distributed.get_rank(group=pg) + world_size = torch.distributed.get_world_size(group=pg) + + # local tensor in each process to store the result + is_in_the_same_node = torch.tensor( + [0] * world_size, dtype=torch.int32, device="cpu" + ) + + # global ranks of the processes in the group + ranks = torch.distributed.get_process_group_ranks(pg) + + magic_message = b"magic_message" + shm = None + + try: + with contextlib.suppress(OSError): + if rank == source_rank: + # create a shared memory segment + shm = shared_memory.SharedMemory( + create=True, size=128, name=make_shm_name("nodecheck") + ) + shm.buf[: len(magic_message)] = magic_message + torch.distributed.broadcast_object_list( + [shm.name], src=ranks[source_rank], group=pg + ) + is_in_the_same_node[rank] = 1 + else: + # try to open the shared memory segment + recv = [None] + torch.distributed.broadcast_object_list( + recv, src=ranks[source_rank], group=pg + ) + name = recv[0] + # fix to https://stackoverflow.com/q/62748654/9191338 + # Python incorrectly tracks shared memory even if it is not + # created by the process. The following patch is a workaround. + with patch( + "multiprocessing.resource_tracker.register", + lambda *args, **kwargs: None, + ): + shm = shared_memory.SharedMemory(name=name) + if shm.buf[: len(magic_message)] == magic_message: + is_in_the_same_node[rank] = 1 + except Exception as e: + logger.error("Error ignored in is_in_the_same_node: %s", e) + finally: + if shm: + shm.close() + + torch.distributed.barrier(group=pg) + + # clean up the shared memory segment + with contextlib.suppress(OSError): + if rank == source_rank and shm: + shm.unlink() + torch.distributed.all_reduce(is_in_the_same_node, group=pg) + + return [x == 1 for x in is_in_the_same_node.tolist()] + + +vllm_get_pp_group = None +vllm_get_tp_group = None +vllm_get_world_group = None + + +def monkey_patch_vllm_parallel_state(reverse: bool = False): + try: + import vllm.distributed.parallel_state as vllm_parallel_state + except ImportError: + return + + global vllm_get_pp_group, vllm_get_tp_group, vllm_get_world_group + if vllm_get_pp_group is None: + vllm_get_pp_group = vllm_parallel_state.get_pp_group + vllm_get_tp_group = vllm_parallel_state.get_tp_group + vllm_get_world_group = vllm_parallel_state.get_world_group + if reverse: + setattr(vllm_parallel_state, "get_pp_group", vllm_get_pp_group) + setattr(vllm_parallel_state, "get_tp_group", vllm_get_tp_group) + setattr(vllm_parallel_state, "get_world_group", vllm_get_world_group) + else: + setattr(vllm_parallel_state, "get_pp_group", get_pp_group) + setattr(vllm_parallel_state, "get_tp_group", get_tp_group) + setattr(vllm_parallel_state, "get_world_group", get_world_group) diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/pcie_ipc_ar.py b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/pcie_ipc_ar.py new file mode 100644 index 0000000..290a12b --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/patches/pcie_ipc_v1/sg/pcie_ipc_ar.py @@ -0,0 +1,378 @@ +# SPDX-License-Identifier: Apache-2.0 +"""FlashInfer PCIe-IPC all-reduce for switch-free intra-node machines. + +Backport of sglang PR #34528 ("[SM120] Add optional FlashInfer PCIe-IPC +all-reduce for switch-free hosts") onto v0.5.19. The kernels come from +flashinfer main (PR flashinfer#4393, merged 2026-08-20) and are vendored into +the installed flashinfer 0.6.18 because no release carries them yet. + +FlashInfer's ``pcie_ipc`` kernels target hosts where every peer transfer crosses +the CPU root complex -- no NVLink, no multicast. The existing custom all-reduce +backends assume one of those fabrics, so on such a host every per-layer +reduction falls back to NCCL, which leaves a large margin on the table at +decode sizes (measured on 8x SM120, hidden 6144 bf16: 12 KiB 5.3 us vs NCCL +205.1 us). + +Adaptations from the upstream PR, required by this stack: + +- ``prepare()`` pre-resolves EVERY admissible row count (1..decode_width) via + ``workspace.prepare()`` after tuning. v0.5.19 captures 61 dense decode-graph + batches (1,2,3,...,8,10,12,...,512 -- not powers of two), and flashinfer + resolves an unseen shape with a group-agreement collective whose verdict is + read back on the host, which is impossible inside a CUDA graph capture. + Upstream only tunes power-of-two buckets, which covers the capture lists of + the machines it was benchmarked on but not this one's. +- ``_decode_width()`` falls back through the v0.5.19 argument names when + ``cuda_graph_config.decode.max_bs`` is unresolved. + +Workspace sizing is unchanged: decode width * hidden, so prefill chunks stay on +NCCL (handing prefill to these kernels was measured 66% worse on TTFT and buys +nothing on TPOT). The workspace is created on the first eligible tensor because +the hidden size is not known when the group is built; ranks run the same +sequence of reductions, so they all reach that first call with the same shape. +``SGLANG_PCIE_IPC_MAX_NUMEL`` overrides the derivation. +""" + +import logging +from typing import Any, Optional, Sequence + +import torch +import torch.distributed as dist +from torch.distributed import ProcessGroup + +from sglang.srt.environ import envs + +logger = logging.getLogger(__name__) + +# The kernels build IPC channels for these world sizes only. +_SUPPORTED_WORLD_SIZES = (2, 4, 8) + +# Fallback decode width when the server args are absent (unit tests, embedded +# use). Covers the batch sizes we serve; anything larger stays on NCCL. +_FALLBACK_DECODE_WIDTH = 64 + + +def _decode_max_bs_from_config(config) -> Optional[int]: + """Decode-phase ``max_bs`` from a cuda_graph_config in either shape. + + The resolved ServerArgs carries the parsed ``CudaGraphConfig`` dataclass on + some paths and the plain nested dict (as the startup banner prints it) on + others; treat both as first-class so a shape change upstream cannot shrink + the workspace back to the fallback. + """ + if config is None: + return None + decode = getattr(config, "decode", None) + if isinstance(decode, dict): + return decode.get("max_bs") + max_bs = getattr(decode, "max_bs", None) + if max_bs is not None: + return max_bs + if isinstance(config, dict): + decode = config.get("decode") + if isinstance(decode, dict): + return decode.get("max_bs") + return None + + +def decode_width_from_args(args) -> Optional[int]: + """Rows in the widest decode reduction for a ServerArgs instance, or None. + + Decode runs inside a captured graph, so the largest captured batch bounds + the reduction. Speculative decoding verifies several tokens per sequence in + one forward, and the decode phase's ``max_bs`` is the batch the runner + captures for, which already accounts for that. + """ + width = _decode_max_bs_from_config(getattr(args, "cuda_graph_config", None)) + if width is None: + width = getattr(args, "cuda_graph_max_bs_decode", None) + if width is None: + width = getattr(args, "cuda_graph_max_bs", None) + if width is None: + width = getattr(args, "max_running_requests", None) + return width + + +def _decode_width() -> Optional[int]: + """decode_width_from_args against the process-global ServerArgs, if published.""" + try: + from sglang.srt.server_args import get_global_server_args + + return decode_width_from_args(get_global_server_args()) + except Exception: + return None + + +def _tune_cache_path(world_size: int) -> Optional[str]: + """Where FlashInfer should persist the measured tactics. + + Named explicitly rather than left to default, because the tuning run has to + write the file for the *next* server to reuse it; a workspace built without + a cache path measures every start. + """ + try: + from flashinfer.comm.pcie_ipc_tuning import default_cache_path + + return str(default_cache_path(world_size)) + except (AttributeError, ImportError): + return None + + +def _tune_batches(max_rows: int) -> tuple: + """Row counts to profile, bounded by what the workspace can actually take. + + The autotuner measures one tactic per shape, so profiling rows the workspace + would reject is wasted startup time. Powers of two cover the captured decode + batches; ``max_rows`` is appended because it is the widest reduction served + and is not always a power of two. + """ + rows = [r for r in (1, 2, 4, 8, 16, 32, 64, 128, 256, 512) if r <= max_rows] + if max_rows >= 1 and max_rows not in rows: + rows.append(max_rows) + return tuple(rows) + + +class PcieIpcCommunicator: + """Adapter between ``GroupCoordinator`` and FlashInfer's PCIe-IPC all-reduce.""" + + def __init__( + self, + group: ProcessGroup, + device, + cpu_group: Optional[ProcessGroup] = None, + ): + self.disabled = True + self.max_numel = 0 + self._requested_width: Optional[int] = None + self._workspace: Optional[Any] = None + self._bound_stream: Optional[torch.cuda.Stream] = None + self._cpu_group = cpu_group + + world_size = dist.get_world_size(group=group) + if world_size not in _SUPPORTED_WORLD_SIZES: + logger.debug( + "FlashInfer PCIe-IPC all-reduce skipped: unsupported world size %d", + world_size, + ) + return + + if isinstance(device, int): + device = torch.device(f"cuda:{device}") + if device.type != "cuda": + return + + try: + from flashinfer.comm import PcieIpcAllReduceWorkspace + except ImportError: + logger.warning( + "SGLANG_ENABLE_PCIE_IPC_ALLREDUCE is set but this FlashInfer build has " + "no pcie_ipc modules; falling back to NCCL all-reduce." + ) + return + + # The workspace is built on the first eligible tensor: sizing it needs the + # hidden size, which the group does not know yet. + self._group = group + self._device = device + self._world_size = world_size + self._workspace_cls = PcieIpcAllReduceWorkspace + self._build_failed = False + self.disabled = False + logger.info( + "FlashInfer PCIe-IPC all-reduce enabled (world=%d, workspace sized on " + "first reduction)", + world_size, + ) + + def _ensure_workspace(self, inp: torch.Tensor) -> bool: + """Build the workspace for the widest decode, using ``inp`` for the hidden size. + + Every rank issues the same reductions in the same order, so they all arrive + here with the same shape and agree on ``max_numel`` without extra exchange. + """ + if self._workspace is not None: + return True + if self._build_failed: + return False + + hidden = inp.shape[-1] + override = envs.SGLANG_PCIE_IPC_MAX_NUMEL.get() + if override: + max_numel = override + else: + max_numel = ( + self._requested_width or _decode_width() or _FALLBACK_DECODE_WIDTH + ) * hidden + + try: + self._workspace = self._workspace_cls( + group=self._group, + max_numel=max_numel, + dtype=torch.bfloat16, + tune_batches=_tune_batches(max_numel // hidden), + tune_cache=_tune_cache_path(self._world_size), + ) + except Exception as e: + logger.warning( + "FlashInfer PCIe-IPC workspace (max_numel=%d) failed: %s; " + "reductions stay on NCCL", + max_numel, + e, + ) + self._build_failed = True + self.disabled = True + return False + + self.max_numel = max_numel + logger.info( + "FlashInfer PCIe-IPC workspace built (world=%d, max_numel=%d, " + "hidden=%d)", + self._world_size, + max_numel, + hidden, + ) + self._tune(hidden) + return True + + def prepare(self, hidden: int, decode_width: Optional[int] = None) -> None: + """Build, tune, and pre-resolve ahead of the first reduction. + + Call this from warmup, before any other autotuning starts. Left to the + first reduction instead, the build lands inside SGLang's own FlashInfer + autotune pass, and FlashInfer refuses to profile a collective from + inside an autotune context it did not open -- so the tuning call would + return having measured nothing. + + ``decode_width`` overrides the width derived from the server args; the + warmup hook passes the runner's own ServerArgs width, which is more + reliable than the process-global accessor. + """ + if self.disabled or self._workspace is not None: + return + if decode_width is not None: + self._requested_width = int(decode_width) + probe = torch.empty((1, hidden), dtype=torch.bfloat16, device=self._device) + if not self._ensure_workspace(probe): + return + # _ensure_workspace already tuned this hidden; resolve every admissible + # row count so no agreement collective can land inside a graph capture. + self._prepare_all_rows(hidden) + + def _tune(self, hidden: int) -> None: + """Measure the launch tactics for this workspace's shapes. + + Without this the kernels run FlashInfer's seed policy, and ``supports`` + answers from that policy rather than from measurements on this host. + Results persist to FlashInfer's cache, so only the first server pays. + """ + from flashinfer.autotuner import AutoTuner + + if AutoTuner.get().is_tuning_mode: + logger.warning( + "FlashInfer PCIe-IPC all-reduce reached its first reduction inside " + "another autotune context; skipping autotune and keeping the seed " + "policy. Call PcieIpcCommunicator.prepare() from warmup to tune." + ) + return + if self._cpu_group is None: + logger.warning( + "FlashInfer PCIe-IPC all-reduce has no host group to autotune on; " + "running FlashInfer's seed policy instead of measured tactics." + ) + return + if torch.cuda.is_current_stream_capturing(): + logger.warning( + "FlashInfer PCIe-IPC all-reduce reached its first reduction inside " + "a graph capture; skipping autotune and keeping the seed policy." + ) + return + try: + tuned = self._workspace.tune( + [hidden], dtype=torch.bfloat16, tune_group=self._cpu_group + ) + except Exception as e: + logger.warning( + "FlashInfer PCIe-IPC autotune failed (%s); keeping the seed policy.", + e, + ) + return + if not tuned: + logger.warning( + "FlashInfer PCIe-IPC autotune covered no shapes for hidden=%d " + "(requested rows %s); the kernels keep the seed policy.", + hidden, + _tune_batches(self.max_numel // hidden), + ) + return + logger.info( + "FlashInfer PCIe-IPC all-reduce autotuned %d shape(s) (hidden=%d, rows=%s)", + len(tuned), + hidden, + _tune_batches(self.max_numel // hidden), + ) + + def _prepare_all_rows(self, hidden: int) -> None: + """Resolve every admissible (rows, hidden) shape into the hot cache. + + flashinfer resolves an unseen shape with a group agreement whose verdict + is read back on the host -- impossible inside a CUDA graph capture, and + v0.5.19 captures dense decode batches (1,2,3,...,512), not just the + power-of-two buckets ``tune()`` covers. Pre-resolving the full + admissible row range costs one small collective per shape at startup + and makes a capture-time resolution impossible by construction. + """ + rows_cap = self.max_numel // hidden + shapes = [(r, hidden) for r in range(1, rows_cap + 1)] + try: + self._workspace.prepare(shapes, dtype=torch.bfloat16) + logger.info( + "FlashInfer PCIe-IPC pre-resolved %d shape(s) (hidden=%d, rows 1..%d)", + len(shapes), + hidden, + rows_cap, + ) + except Exception as e: + logger.warning( + "FlashInfer PCIe-IPC shape pre-resolve failed for hidden=%d: %s; " + "an unresolved shape reaching graph capture will fail there", + hidden, + e, + ) + + def should_pcie_ipc_ar(self, inp: torch.Tensor) -> bool: + """Whether FlashInfer can run this exact shape. + + ``supports`` is a capability question (dtype, contiguity, capacity); + an unsupported shape is rejected here and the caller keeps its NCCL + path. Admission is checked before any config resolution, so a shape too + large for the workspace falls back without touching the tuning state. + """ + if self.disabled or not inp.is_contiguous() or inp.dim() < 2: + return False + if not self._ensure_workspace(inp): + return False + if inp.numel() > self.max_numel: + return False + return self._workspace.supports(inp) + + def pcie_ipc_all_reduce(self, inp: torch.Tensor) -> Optional[torch.Tensor]: + """All-reduce ``inp``, rebinding the workspace when the stream changes. + + One workspace serves one stream: its epoch and arrival counters assume + the calls sharing it are totally ordered. SGLang switches streams at + phase boundaries -- graph capture, then replay -- so the binding is + moved with the caller rather than kept on whichever stream got there + first. This is safe because those phases do not overlap; it would not + be safe if two streams issued reductions concurrently on the same group. + """ + stream = torch.cuda.current_stream() + if self._bound_stream is None or stream != self._bound_stream: + self._workspace.rebind_stream() + self._bound_stream = stream + return self._workspace.all_reduce(inp) + + def destroy(self) -> None: + if self._workspace is not None: + self._workspace.destroy() + self._workspace = None + self.disabled = True diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/results/results_raw.md b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/results/results_raw.md new file mode 100644 index 0000000..ecf904d --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/results/results_raw.md @@ -0,0 +1,51 @@ +# 原始压测数据 — DSV4 优化点迁移实验(60.1 / 6000D-1,2026-09-09) + +口径:i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐 tok/s(bench_corpus.py, +PG19 真实语料固定窗口 run-id 9311-9313 池C + 9314/9315 pool-override, +radix OFF 输入逐 token 相同)。每臂 3 轮取中位;arm5 回切确认为 1 轮。 +cc40/64 点 nreq=80/128(沿用方案 D 口径,表注见双场景报告)。 + +## 输出吞吐(out tok/s)逐轮 + 中位 + +| 臂 | cc8 (r1/r2/r3) | cc16 | cc32 | cc40 | cc64 | +|---|---|---|---|---|---| +| **0a** nightly-dev-20260828 基线 | 96.38/104.25/**102.22** | 148.08/148.98/**148.21** | 192.32/191.13/**191.13** | 202.00/198.41/**198.41** | 208.07/208.18/**208.18** | +| **0b** latest(0.5.19+fi0.6.18) 基线 | 87.85/105.79/**101.51** | 147.71/146.66/**147.64** | 190.64/191.07/**191.07** | 200.54/198.01/**198.01** | 207.16/207.74/**207.61** | +| **arm2** autotune+持久缓存 (latest) | 104.78/103.33/**104.78** | 152.53/150.89/**151.08** | 195.03/193.96/**195.03** | 200.63/200.27/**200.63** | 210.01/209.83/**210.01** | +| **arm3** PCIe-IPC 12文件包 (latest) | 96.89/98.57/**97.91** | 143.18/143.30/**143.18** | 187.00/187.08/**187.06** | 198.78/195.17/**197.30** | 204.52/204.34/**204.52** | +| **arm5** 回切确认 stock-latest(1轮) | 102.57 | 147.01 | 190.94 | 200.58 | 207.80 | + +注:0b cc8 r1=87.85 为该臂首轮冷 JIT(fi_jit_cache/sglang_cache 挂载此臂首次填充); +arm5 同为全新容器首轮但缓存已热 → 102.57,直接落在 0b r2/r3 暖机水平。 + +## 中位数对比(vs 0b latest 基线) + +| 臂 | cc8 | cc16 | cc32 | cc40 | cc64 | 判定 | +|---|---|---|---|---|---|---| +| 0a(nightly) | +0.7% | +0.4% | +0.0% | +0.2% | +0.3% | 持平 → 按用户规则转 latest | +| **arm2(autotune)** | **+3.2%** | **+2.3%** | **+2.1%** | **+1.3%** | **+1.2%** | **五点全胜,唯一胜者** | +| arm3(IPC) | −3.5% | −3.0% | −2.1% | −0.4% | −1.5% | 五点全降,判负 | +| arm5(回切确认) | +1.0% | −0.4% | −0.1% | +1.3% | +0.1% | 全落基线带,A/B/A 闭环 | + +分布不重叠性:arm2 vs 0b 在 cc16/32/64 轮间分布不重叠(arm2 min > 0b max); +arm3 vs 0b 在 cc8/16/32/64 分布不重叠(方向为劣)。 + +## TPOT 副指标(r1 轮 mean,秒/token) + +| 臂 | cc8 | cc16 | cc64 | +|---|---|---|---| +| 0b 基线 | 0.0592 | 0.0778 | 0.1430 | +| arm2 autotune | 0.0565 (−4.6%) | 0.0745 (−4.2%) | 0.1401 (−2.0%) | +| arm3 IPC | 0.0628 (+6.1%) | 0.0813 (+4.5%) | 0.1455 (+1.7%) | + +## 显存与 KV 池 + +全臂 KV 池恒定 **max_total_num_tokens=1,040,384**(无容量漂移)。 +latest 镜像下 available_gpu_mem:stock 15.64 GB → autotune 11.79 GB +(**autotune tactic 缓冲代价 ≈3.9 GB/卡**;池不缩)。 + +## 质量门(quality_gate_605.sh) + +全臂(0a/0b/arm2/arm3)一致 **6/7**:GSM8K×5 + 中文推理全过; +tool-call FAIL 为方案 D 无 --tool-call-parser 的配置缺口(生产上需补 +--tool-call-parser glm47 与 --reasoning-parser glm45),非本实验变量。 diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/deploy_glm53_exp.sh b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/deploy_glm53_exp.sh new file mode 100644 index 0000000..6428779 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/deploy_glm53_exp.sh @@ -0,0 +1,70 @@ +#!/bin/bash +# deploy_glm53_exp.sh — 方案D(TP2PP4) 配置克隆部署器,用于 DSV4 报告优化点迁移实验。 +# 用法(环境变量驱动): +# EXP_NAME=glm53-exp0a EXP_IMAGE=lmsysorg/sglang:nightly-dev-20260828-daf63171 bash deploy_glm53_exp.sh +# AUTOTUNE=1 -> 去掉 --disable-flashinfer-autotune +# EXTRA_ENV='-e A=1 -e B=2' -> 追加 docker 环境变量 +# INJECT_DIR=/path -> 含 inject_manifest.txt(每行 relpath:container_path)的注入目录 +# 缓存挂载(/data/glm53_exp/ 下 fi_jit_cache/sglang_cache/triton_cache)全臂恒开, +# 保证除研究变量外环境完全一致(DSV4 教训:先排除环境噪声再定罪)。 +# 原生产容器 glm53-pp4 全程不动,本部署器只创建/销毁 glm53-exp* 专用名。 +set -uo pipefail +NAME=${EXP_NAME:?EXP_NAME required} +IMAGE=${EXP_IMAGE:?EXP_IMAGE required} +AUTOTUNE=${AUTOTUNE:-0} +EXTRA_ENV=${EXTRA_ENV:-} +INJECT_DIR=${INJECT_DIR:-} +PORT=30000 +mkdir -p /data/glm53_exp/fi_jit_cache /data/glm53_exp/sglang_cache /data/glm53_exp/triton_cache /data/glm53_exp/logs + +docker rm -f "$NAME" >/dev/null 2>&1 || true +for i in $(seq 1 15); do docker ps -a --format '{{.Names}}' | grep -q "^${NAME}$" || break; sleep 2; done +for i in $(seq 1 90); do + m=$(nvidia-smi --query-gpu=memory.used --format=csv,noheader,nounits | sort -n | tail -1) + echo "drain check ${i}: max_gpu_mem=${m}MiB $(date +%T)" + [ "$m" -lt 2000 ] && break + sleep 10 +done +if [ "${m:-99999}" -ge 2000 ]; then echo "DRAIN_TIMEOUT max=${m}MiB"; exit 1; fi + +AUTOTUNE_FLAG="--disable-flashinfer-autotune" +if [ "$AUTOTUNE" = "1" ]; then AUTOTUNE_FLAG=""; fi + +docker create --gpus all --shm-size 64g --ipc=host -p ${PORT}:${PORT} \ + -v /data/hf_models:/data/hf_models \ + -v /data/glm53_exp/fi_jit_cache:/root/.cache/flashinfer \ + -v /data/glm53_exp/sglang_cache:/root/.cache/sglang \ + -v /data/glm53_exp/triton_cache:/root/.triton \ + -e SGLANG_CACHE_DIR=/root/.cache/sglang \ + $EXTRA_ENV \ + --name "$NAME" "$IMAGE" \ + python3 -m sglang.launch_server --model-path /data/hf_models/GLM-5.3-NVFP4 \ + --tp-size 2 --pp-size 4 --mem-fraction-static 0.85 --max-running-requests 48 \ + --disable-radix-cache --disable-shared-experts-fusion \ + --moe-runner-backend flashinfer_cutlass $AUTOTUNE_FLAG \ + --disable-custom-all-reduce --chunked-prefill-size 16384 \ + --host 0.0.0.0 --port ${PORT} \ + --json-model-override-args '{"index_topk_freq": 4}' >/dev/null || { echo CREATE_FAILED; exit 1; } + +CID=$(docker ps -aq --filter "name=${NAME}") +[ -n "$CID" ] || { echo CREATE_FAILED_NO_CID; exit 1; } + +if [ -n "$INJECT_DIR" ] && [ -f "$INJECT_DIR/inject_manifest.txt" ]; then + while IFS=':' read -r rel target; do + case "$rel" in ''|\#*) continue;; esac + docker cp "$INJECT_DIR/$rel" "$CID:$target" || { echo "CP_FAILED $rel"; exit 1; } + done < "$INJECT_DIR/inject_manifest.txt" +fi + +docker start "$NAME" || { echo START_FAILED; exit 1; } +READY=0 +for i in $(seq 1 90); do + sleep 20 + curl -sf localhost:${PORT}/health >/dev/null 2>&1 && { READY=1; break; } + docker ps --format '{{.Names}}' | grep -q "^${NAME}$" || { echo "CONTAINER_DIED at check $i"; docker logs "$NAME" 2>&1 | tail -30; exit 1; } +done +if [ "$READY" != "1" ]; then echo "START_TIMEOUT"; docker logs "$NAME" 2>&1 | tail -40; exit 1; fi +echo "READY at $(date +%H:%M:%S)" +docker logs "$NAME" 2>&1 | grep -m1 "max_total_num_tokens" || true +docker ps --filter "name=${NAME}" --format '{{.Names}} {{.Status}}' +echo "DEPLOY_DONE $NAME" diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_arm_chain.sh b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_arm_chain.sh new file mode 100644 index 0000000..d340da8 --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_arm_chain.sh @@ -0,0 +1,16 @@ +#!/bin/bash +# run_arm_chain.sh — 单臂完整链:质量门 + 五点×3 轮。 +# 用法: CONTAINER=glm53-exp0a TAGPREFIX=arm0a bash run_arm_chain.sh +set -uo pipefail +C=${CONTAINER:?CONTAINER required} +P=${TAGPREFIX:?TAGPREFIX required} +L=/data/glm53_exp/logs + +bash /root/quality_gate_605.sh > $L/gate_${P}.log 2>&1 +echo "GATE_DONE rc=$? $(date +%T)" >> $L/chain_${P}.log + +for r in r1 r2 r3; do + CONTAINER=$C TAG=${P}_${r} bash /root/run_exp_s2.sh >> $L/sweep_${P}.log 2>&1 + echo "ROUND ${P}_${r} rc=$? $(date +%T)" >> $L/chain_${P}.log +done +echo "CHAIN_DONE ${P} $(date +%T)" >> $L/chain_${P}.log diff --git a/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_exp_s2.sh b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_exp_s2.sh new file mode 100644 index 0000000..89cec5b --- /dev/null +++ b/experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/run_exp_s2.sh @@ -0,0 +1,31 @@ +#!/bin/bash +# run_exp_s2.sh — 16k/512 五点扫描(cc 8/16/32/40/64),单臂单轮。 +# 语料窗口与 run_mono_s2.sh 完全一致(9311-9313 池C + 9314/9315 pool-override), +# 所有臂用同一窗口 -> 输入逐 token 相同,A/B 最干净(radix OFF 无缓存污染)。 +# 用法: CONTAINER=glm53-exp0a TAG=arm0a_r1 bash run_exp_s2.sh +set -uo pipefail +C=${CONTAINER:?CONTAINER required} +TAG=${TAG:?TAG required} +mkdir -p /root/bench_logs +curl -s -m 5 -o /dev/null http://127.0.0.1:30000/health || { echo "ABORT: server unhealthy"; exit 1; } +docker ps --format '{{.Names}}' | grep -qx "$C" || { echo "ABORT: $C not running"; exit 1; } + +run_point() { + CC=$1; NR=$2; RID=$3; PO=$4 + EXTRA="" + [ -n "$PO" ] && EXTRA="--pool-override $PO" + echo "=== [$TAG] cc=$CC nreq=$NR run=$RID pool=$PO start $(date +%T) ===" >> /root/bench_logs/exp_runner.log + timeout 1800 python3 /root/bench_corpus.py --corpus /root/corpus_ids.json \ + --input-len 16384 --concurrency "$CC" --num-requests "$NR" --run-id "$RID" \ + --shared-frac 0 --output-len 512 --container "$C" $EXTRA \ + > "/root/bench_logs/exp_${TAG}_run${RID}.log" 2>&1 + echo "=== [$TAG] run=$RID done rc=$? $(date +%T) ===" >> /root/bench_logs/exp_runner.log +} + +run_point 8 16 9311 "" +run_point 16 32 9312 "" +run_point 32 32 9313 "" +run_point 40 80 9314 4000000 +run_point 64 128 9315 6000000 +echo "[$TAG] ALL_DONE $(date +%T)" >> /root/bench_logs/exp_runner.log +echo "SWEEP_DONE $TAG"