glm53 dsv4-migration bench: autotune+latest-image winner (+1.2~3.2% all 5 pts, A/B/A confirmed), PCIe-IPC pack negative (-0.4~-3.5%), page-mark kernel N/A for GLM (60.1, 09-09)

This commit is contained in:
yy-fighting 2026-09-10 03:13:02 +08:00
parent f15003d2f5
commit e3476aff86
29 changed files with 16877 additions and 1 deletions

View File

@ -7,7 +7,7 @@
| 机器 | 在役 | 口径 / 归属 | 对应 profile | | 机器 | 在役 | 口径 / 归属 | 对应 profile |
|---|---|---|---| |---|---|---|---|
| 60.1 (6000D-1) | `glm53-pp4`Up8 卡满载,:30000 | **方案 D 生产**。09-08 PD 压测窗口停机 ~2.5h 后已恢复并核验 | `profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4.env` | | 60.1 (6000D-1) | `glm53-pp4`Up8 卡满载,: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 phaserun_phase.sh 实验链进行中) | **他人实验进行中,勿动**(方案 F PD 链已拆除GPU7 曾有外部裸金属任务)。动卡前仍先核实归属 | `profiles/pro6000/glm53_nvfp4_pro6000_pd_{prefill→decode 侧}.env` | | 60.2 (6000D-2) | `glm53-nvfp4` 实验容器09-08 晚 TP1PP8 phaserun_phase.sh 实验链进行中) | **他人实验进行中,勿动**(方案 F PD 链已拆除GPU7 曾有外部裸金属任务)。动卡前仍先核实归属 | `profiles/pro6000/glm53_nvfp4_pro6000_pd_{prefill→decode 侧}.env` |
| 60.3 | 无容器,但 8 卡被外部裸金属训练占用(`/data/mas/larm`09-08 晚实测) | 外部任务,勿动(此前台账漏记) | — | | 60.3 | 无容器,但 8 卡被外部裸金属训练占用(`/data/mas/larm`09-08 晚实测) | 外部任务,勿动(此前台账漏记) | — |
| 60.4 | `glm53-nvfp4`Up:30000restart=unless-stopped | **TP8+EAGLE+custom-AR 1stageE7b 配方)在役**09-09 下午部署EAGLE 4/1/5/mem0.90/MRR16/chunk8192/ctxlen270336/fp8KV+hicache3/decode 图桶 1-8/双 parserCAR 补丁三处全注入、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/ 补丁快照) | | 60.4 | `glm53-nvfp4`Up:30000restart=unless-stopped | **TP8+EAGLE+custom-AR 1stageE7b 配方)在役**09-09 下午部署EAGLE 4/1/5/mem0.90/MRR16/chunk8192/ctxlen270336/fp8KV+hicache3/decode 图桶 1-8/双 parserCAR 补丁三处全注入、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/ 补丁快照) |

View File

@ -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 GBtactic 缓冲 ≈3.9GB/卡),
# KV 池不缩1,040,384
# - PCIe-IPC AllReduce 包在同一实验中五点全降(-0.4~-3.5%)判负勿叠用:
# TP2 单对端 NCCL AR 同 switch P2P 已近最优
# - 上生产须补 --tool-call-parser glm47 与 --reasoning-parser glm45D 系既有缺口)
# - 部署器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}"

View File

@ -0,0 +1,81 @@
# DSV4-Flash 优化点迁移实验报告 — GLM-5.3-NVFP4 @ 60.16000D-1
日期2026-09-09 机器174.1.60.18×RTX 6000DSM120无 NVLink 纯 PCIe
基线:方案 DTP2 PP4 mono60.1 生产 glm53-pp4 原样配方)
口径i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐高者优(用户判定规则)
源报告飞书《DSV4-Flash 优化日报》(A3Fmw6YePikH9lkRx8scN1UYngw) +
《DSV4-Flash 单机八卡推理优化完整报告》(Bby6wiQ9Bi8yGjkB6EIcG1CqnJe)
## 一、迁移性判定总表(先判后测)
| DSV4 优化点 | GLM-5.3 迁移判定 | 依据 |
|---|---|---|
| P0 native-headsdeepseek_v4.py | **N/A不实验** | GLM 已原生直通 flashinfer8/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.0cc16/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.50.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 <host>/fi_jit_cache:/root/.cache/flashinfer \
-v <host>/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}

View File

@ -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)

View File

@ -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"
}

View File

@ -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

View File

@ -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)

View File

@ -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)

View File

@ -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"
}

View File

@ -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

View File

@ -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` —— **GLMGlmMoeDsaForCausalLM的入口**
`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 无增益未实验

View File

@ -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}")

View File

@ -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

View File

@ -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])

View File

@ -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",
]

View File

@ -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 <tvm/ffi/container/array.h>
#include <cstdint>
#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<int>(world_size), max_numel, static_cast<int>(elem_size),
static_cast<int>(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<fptr_t> ipc_ptrs, int64_t rank, int64_t max_numel, int64_t elem_size,
int64_t max_blocks) {
const int world_size = static_cast<int>(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<int>(elem_size),
static_cast<int>(max_blocks));
handle->views = fi::make_peer_views(ptrs, world_size, static_cast<int>(rank), handle->layout);
handle->rank = static_cast<int>(rank);
handle->world_size = world_size;
handle->max_blocks = static_cast<int>(max_blocks);
handle->max_numel = max_numel;
handle->elem_size = static_cast<int>(elem_size);
cudaError_t err = cudaMemset(reinterpret_cast<void*>(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<fptr_t>(handle);
}
void pcie_ipc_dispose(fptr_t handle) { delete reinterpret_cast<PcieIpcHandle*>(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<PcieIpcHandle*>(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<size_t>(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<fi::Variant>(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<nv_bfloat16>(static_cast<const nv_bfloat16*>(inp.data_ptr()),
static_cast<nv_bfloat16*>(out.data_ptr()), numel, h->views,
h->rank, h->world_size, h->max_blocks, h->max_numel,
static_cast<int>(blocks), static_cast<int>(threads), algo,
enable_pdl, stream);
break;
case float16_code:
err = fi::all_reduce<half>(
static_cast<const half*>(inp.data_ptr()), static_cast<half*>(out.data_ptr()), numel,
h->views, h->rank, h->world_size, h->max_blocks, h->max_numel, static_cast<int>(blocks),
static_cast<int>(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);

View File

@ -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",
],
)

View File

@ -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"
}

View File

@ -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: ...

View File

@ -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

View File

@ -0,0 +1,51 @@
# 原始压测数据 — DSV4 优化点迁移实验60.1 / 6000D-12026-09-09
口径i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐 tok/sbench_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-latest1轮 | 102.57 | 147.01 | 190.94 | 200.58 | 207.80 |
0b cc8 r1=87.85 为该臂首轮冷 JITfi_jit_cache/sglang_cache 挂载此臂首次填充);
arm5 同为全新容器首轮但缓存已热 → 102.57,直接落在 0b r2/r3 暖机水平。
## 中位数对比vs 0b latest 基线)
| 臂 | cc8 | cc16 | cc32 | cc40 | cc64 | 判定 |
|---|---|---|---|---|---|---|
| 0anightly | +0.7% | +0.4% | +0.0% | +0.2% | +0.3% | 持平 → 按用户规则转 latest |
| **arm2autotune** | **+3.2%** | **+2.3%** | **+2.1%** | **+1.3%** | **+1.2%** | **五点全胜,唯一胜者** |
| arm3IPC | 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_memstock 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非本实验变量。

View File

@ -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"

View File

@ -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

View File

@ -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"