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:
parent
f15003d2f5
commit
e3476aff86
@ -7,7 +7,7 @@
|
||||
|
||||
| 机器 | 在役 | 口径 / 归属 | 对应 profile |
|
||||
|---|---|---|---|
|
||||
| 60.1 (6000D-1) | `glm53-pp4`(Up,8 卡满载,:30000) | **方案 D 生产**。09-08 PD 压测窗口停机 ~2.5h 后已恢复并核验 | `profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4.env` |
|
||||
| 60.1 (6000D-1) | `glm53-pp4`(Up,8 卡满载,:30000) | **方案 D 生产**。09-09 DSV4 优化点迁移实验窗口停机 ~7h 后已恢复并核验(09-08 PD 窗口 ~2.5h 前史)。**升级路径已交付未执行**:latest 镜像(d6e72886=sglang0.5.19+fi0.6.18)+开 autotune+持久 SGLANG_CACHE_DIR/fi_jit_cache 挂载,16k/512 cc8-64 五点 +1.2~+3.2%(A/B/A 回切确认因果),代价 ~3.9GB/卡 tactic 缓冲、KV 池不缩;profile `..._tp2pp4_latest_autotune.env`,实验全量 `experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/`。同轮判决:PCIe-IPC AR 包五点全降判负勿叠用;page-mark 向量化对 GLM 双证据判死(GLM 入口不经过该 kernel)。`dsv4_scan`(DSV4 团队扫描容器)经授权保持停止,还原说明 `/root/dsv4_scan_restore_note.txt` | `profiles/pro6000/glm53_nvfp4_pro6000_sglang_tp2pp4.env`(现役);升级候选 `..._tp2pp4_latest_autotune.env` |
|
||||
| 60.2 (6000D-2) | `glm53-nvfp4` 实验容器(09-08 晚 TP1PP8 phase,run_phase.sh 实验链进行中) | **他人实验进行中,勿动**(方案 F PD 链已拆除;GPU7 曾有外部裸金属任务)。动卡前仍先核实归属 | `profiles/pro6000/glm53_nvfp4_pro6000_pd_{prefill→decode 侧}.env` |
|
||||
| 60.3 | 无容器,但 8 卡被外部裸金属训练占用(`/data/mas/larm`,09-08 晚实测) | 外部任务,勿动(此前台账漏记) | — |
|
||||
| 60.4 | `glm53-nvfp4`(Up,:30000,restart=unless-stopped) | **TP8+EAGLE+custom-AR 1stage(E7b 配方)在役**(09-09 下午部署:EAGLE 4/1/5/mem0.90/MRR16/chunk8192/ctxlen270336/fp8KV+hicache3/decode 图桶 1-8/双 parser;CAR 补丁三处全注入、8 rank `SSKJ_CAR_PATCH_ACTIVE` 确认,health 200、生成冒烟、质量门 7/7 见 `/root/qg_604_car.log`)。启动 `/root/deploy_glm53_604_exp.sh`(=60.7 实验版逐字拷贝,`RESTART=yes CAR_PATCH=1`),补丁 `/root/patches/`(md5 与仓库 car_patch 归档一致)。当日早间曾短暂部署 TP2PP4 D 配方复刻(deploy_s2_test_604.sh 留盘可切回)后被本方案替换;同日经授权清退外部 vllm 评测流水线(tmux `mas` 的 run_multiseed.sh 链,--resume 可续跑) | `experiments/pro6000/glm53_nvfp4_pro6000d_sglang_dual_scenario_bench/scripts/`(deploy_glm53_604_exp.sh + car_patch/ 补丁快照) |
|
||||
|
||||
@ -0,0 +1,49 @@
|
||||
# GLM-5.3-NVFP4 SGLang TP=2 PP=4 + latest镜像 + autotune profile(方案 D 升级版)。
|
||||
# 2026-09-09 DSV4 优化点迁移实验优胜配置:五点(16k/512 cc8-64)全胜 +1.2~+3.2%,
|
||||
# A/B/A 回切确认因果成立。见 experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/。
|
||||
# 已入档未执行;60.1 生产仍跑 nightly 版(glm53-pp4,见同目录 tp2pp4.env)。
|
||||
#
|
||||
# 关键点(实测,勿随意改):
|
||||
# - 相对 tp2pp4.env 仅两处变化:镜像 nightly-dev-20260828 → latest(d6e72886);
|
||||
# 删 --disable-flashinfer-autotune。启动参数其余逐字相同
|
||||
# - 镜像中性实测:latest 与 nightly 同配置五点差 ≤0.7%(换镜像无风险无收益)
|
||||
# - autotune 增益钉在 tactic 缓存上:SGLANG_CACHE_DIR 与 flashinfer JIT 缓存
|
||||
# 必须挂宿主持久盘,否则重部署重抽签(增益消失/不可复现,DSV4 §12.6 同款教训)
|
||||
# - 代价:available_gpu_mem 15.64→11.79 GB(tactic 缓冲 ≈3.9GB/卡),
|
||||
# KV 池不缩(1,040,384)
|
||||
# - PCIe-IPC AllReduce 包在同一实验中五点全降(-0.4~-3.5%)判负勿叠用:
|
||||
# TP2 单对端 NCCL AR 同 switch P2P 已近最优
|
||||
# - 上生产须补 --tool-call-parser glm47 与 --reasoning-parser glm45(D 系既有缺口)
|
||||
# - 部署器:experiments/pro6000/glm53_nvfp4_dsv4_migration_bench/scripts/deploy_glm53_exp.sh
|
||||
# (EXP_NAME=<名> EXP_IMAGE=lmsysorg/sglang:latest AUTOTUNE=1)
|
||||
|
||||
PLATFORM=pro6000
|
||||
EXPERIMENT=glm53_nvfp4_pro6000_sglang_tp2pp4_latest_autotune
|
||||
MODEL_NAME=GLM-5.3-NVFP4
|
||||
ENGINE=sglang
|
||||
RUNTIME=docker
|
||||
DOCKER_IMAGE=lmsysorg/sglang:latest
|
||||
DOCKER_IMAGE_DIGEST=sha256:d6e7288627be8b02be88e4bba38e73f6d50e2826869f753c13a4c4385ab3eda9
|
||||
CONTAINER_NAME=glm53-pp4-autotune
|
||||
MODEL_PATH=/data/hf_models/GLM-5.3-NVFP4
|
||||
SERVED_MODEL_NAME=/data/hf_models/GLM-5.3-NVFP4
|
||||
PORT=30000
|
||||
HEALTH_PATH=/health
|
||||
HEALTH_WAIT_S=600
|
||||
CONTAINER_PYTHON=python3
|
||||
|
||||
TP=2
|
||||
PP=4
|
||||
MEM_FRACTION_STATIC=0.85
|
||||
MAX_RUNNING_REQUESTS=48
|
||||
CHUNKED_PREFILL_SIZE=16384
|
||||
|
||||
DEVICE_VARS="CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7"
|
||||
ENGINE_ENV="PYTHONUNBUFFERED=1 HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1 SGLANG_CACHE_DIR=/root/.cache/sglang"
|
||||
|
||||
DOCKER_FLAGS="--gpus all --shm-size 64g --ipc=host -p ${PORT}:${PORT}"
|
||||
VOLUMES="/data/hf_models:/data/hf_models /data/glm53_exp/fi_jit_cache:/root/.cache/flashinfer /data/glm53_exp/sglang_cache:/root/.cache/sglang /data/glm53_exp/triton_cache:/root/.triton"
|
||||
|
||||
BOOTSTRAP="python3 -m sglang.launch_server ${LAUNCH_ARGS}"
|
||||
|
||||
LAUNCH_ARGS="--model-path ${MODEL_PATH} --tp-size ${TP} --pp-size ${PP} --mem-fraction-static ${MEM_FRACTION_STATIC} --max-running-requests ${MAX_RUNNING_REQUESTS} --disable-radix-cache --disable-shared-experts-fusion --moe-runner-backend flashinfer_cutlass --disable-custom-all-reduce --chunked-prefill-size ${CHUNKED_PREFILL_SIZE} --host 0.0.0.0 --port ${PORT} --json-model-override-args {\"index_topk_freq\": 4}"
|
||||
@ -0,0 +1,81 @@
|
||||
# DSV4-Flash 优化点迁移实验报告 — GLM-5.3-NVFP4 @ 60.1(6000D-1)
|
||||
|
||||
日期:2026-09-09 | 机器:174.1.60.1(8×RTX 6000D,SM120,无 NVLink 纯 PCIe)
|
||||
基线:方案 D(TP2 PP4 mono,60.1 生产 glm53-pp4 原样配方)
|
||||
口径:i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐高者优(用户判定规则)
|
||||
源报告:飞书《DSV4-Flash 优化日报》(A3Fmw6YePikH9lkRx8scN1UYngw) +
|
||||
《DSV4-Flash 单机八卡推理优化完整报告》(Bby6wiQ9Bi8yGjkB6EIcG1CqnJe)
|
||||
|
||||
## 一、迁移性判定总表(先判后测)
|
||||
|
||||
| DSV4 优化点 | GLM-5.3 迁移判定 | 依据 |
|
||||
|---|---|---|
|
||||
| P0 native-heads(deepseek_v4.py) | **N/A,不实验** | GLM 已原生直通 flashinfer(8/16/32-head 模板齐全),无中间层可剥 |
|
||||
| P1 page-mark 向量化(BLOCK=1024) | **判死(双证据),压测前撤臂** | GLM 入口 `flashinfer_sparse_mla_forward:607` 直调 trtllm_batch_decode_with_kv_cache_mla,不经过 `_split_kv_pages_to_64`(page-mark 唯一入口,DSV4 SWA 专用);运行时 GLM serving 0 个 Triton kernel 启动。证据见 `patches/page_mark_verdict_evidence.md` |
|
||||
| pdi/SWA 相关整族 | **N/A** | GLM-5.3 无滑动窗口注意力 |
|
||||
| chunk2048+radix 命中 | **N/A** | 本场景 16k 独立输入无共享前缀,radix OFF 为既定口径 |
|
||||
| flashinfer autotune + 持久 tactic 缓存 | **可迁移 → 实验(arm2)** | DSV4 E2 实测 autotune-on + 持久缓存胜 autotune-off ~7%;GLM 同为 flashinfer 栈 |
|
||||
| PCIe-IPC AllReduce 包(PR#34528 + fi#4393) | **可迁移 → 实验(arm3)** | 同代 SM120 无 NVLink 平台同痛点;DSV4 TP4 实测 +1.8% |
|
||||
| 最新镜像(sglang 0.5.19 + flashinfer 0.6.18) | **可迁移 → 实验(0b)** | 用户加测;DSV4 团队验证过、IPC manifest 指定镜像 |
|
||||
| DeepGEMM MoE runner / dsa-topk flashinfer / AR-fusion | **勿碰** | DSV4 报告负结果 + 本集群已知 SM120 门禁(见双场景报告勿碰清单) |
|
||||
|
||||
## 二、实验臂与判决
|
||||
|
||||
| 臂 | 配置 | 判决 |
|
||||
|---|---|---|
|
||||
| 0a | D 配置 @ nightly-dev-20260828(生产镜像) | 基线带 102.2/148.2/191.1/198.4/208.2,与 mono 战役发表 D 数字交叉一致 |
|
||||
| 0b | D 配置 @ latest(d6e72886, 0.5.19+fi0.6.18) | 与 0a 差 ≤0.7% 全点持平 → **按用户规则转 latest 跑后续臂**(镜像中性,无回归) |
|
||||
| 1 | page-mark 向量化 | **压测前撤销**(双证据判死,省一整轮),见判定总表 |
|
||||
| 2 | latest + autotune-on + 持久 SGLANG_CACHE_DIR/fi_jit_cache | **五点全胜 +1.2~+3.2%(唯一胜者)**:104.8/151.1/195.0/200.6/210.0;cc16/32/64 分布与基线不重叠;TPOT 全点改善(−2~−4.6%);代价 available_gpu_mem 15.64→11.79 GB(≈3.9GB/卡 tactic 缓冲),KV 池不缩(1,040,384);门 6/7 不变 |
|
||||
| 3 | latest + PCIe-IPC 12 文件包 + 双 env 开关 | **五点全降判负**:97.9/143.2/187.1/197.3/204.5(−0.4~−3.5%),cc8/16/32/64 分布不重叠(方向为劣),TPOT 全点变差;8 rank 均确认 `FlashInfer PCIe-IPC all-reduce enabled (world=2)`——激活正确、真输。机制:TP2 单对端 NCCL AR 走同 switch P2P 已近最优,IPC 工作区拷贝+跨块同步纯倒贴;DSV4 的 +1.8% 本就是 TP4 口径,TP 度越低收益越小,TP2 翻负。与既有结论「TP2 AR≈free」互证 |
|
||||
| 4 | 冠军组合 | 不需要(唯一胜者 autotune 已完整测量=arm2) |
|
||||
| 5 | 回切确认:stock-latest 全新部署 1 轮 | **PASS**:102.6/147.0/190.9/200.6/207.8 全落 0b 基线带(缓存已热,无 0b r1 式冷 JIT 压低),且低于 arm2 → A/B/A 闭环,**autotune 增益判定为因果非环境漂移** |
|
||||
|
||||
原始数据:`results/results_raw.md`。
|
||||
|
||||
## 三、优胜配置交付(已入档未执行)
|
||||
|
||||
**推荐部署 = 方案 D 配方升级两处:镜像换 latest + 开 autotune(持久缓存)**
|
||||
|
||||
```
|
||||
镜像: lmsysorg/sglang:latest (sha256:d6e7288627be8b02be88e4bba38e73f6d50e2826869f753c13a4c4385ab3eda9
|
||||
= sglang 0.5.19 + flashinfer 0.6.18;生产现役 nightly-dev-20260828 同 digest 已本地 tag 为 dev)
|
||||
启动: python3 -m sglang.launch_server --model-path /data/hf_models/GLM-5.3-NVFP4 \
|
||||
--tp-size 2 --pp-size 4 --mem-fraction-static 0.85 --max-running-requests 48 \
|
||||
--disable-radix-cache --disable-shared-experts-fusion \
|
||||
--moe-runner-backend flashinfer_cutlass \
|
||||
--disable-custom-all-reduce --chunked-prefill-size 16384 \
|
||||
--host 0.0.0.0 --port 30000 --json-model-override-args '{"index_topk_freq": 4}'
|
||||
# 相对现役:镜像 nightly-dev→latest;删 --disable-flashinfer-autotune;其余逐字不变
|
||||
挂载(新增,持久化必需): -v <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}
|
||||
@ -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)
|
||||
File diff suppressed because it is too large
Load Diff
@ -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"
|
||||
}
|
||||
@ -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
|
||||
@ -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)
|
||||
@ -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)
|
||||
@ -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"
|
||||
}
|
||||
@ -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
|
||||
|
||||
@ -0,0 +1,44 @@
|
||||
# page-mark 向量化补丁对 GLM-5.3-NVFP4 判死:双证据
|
||||
|
||||
DSV4 报告的 E2 优化点「`_page_mark_kernel` 标量→向量化(BLOCK=1024)」在 DSV4-Flash 上
|
||||
实测 prefill +3.0%。评估迁移到 GLM-5.3-NVFP4 时,静态 + 运行时双证据证明 GLM 路径
|
||||
**根本不执行该 kernel**,补丁零增益,实验臂在压测前撤销(省一整轮部署+扫描)。
|
||||
|
||||
## 证据一(静态,调用链闭环)
|
||||
|
||||
目标文件 `flash_mla_sm120.py`(两镜像 nightly-dev-20260828 / latest md5 完全相同:
|
||||
`9df4d50d44df9b9f6db36538f2d8b52f`,基础版已归档为本目录 `base_flash_mla_sm120.py`)。
|
||||
|
||||
全文件仅有的调用点(grep 实证):
|
||||
|
||||
- `:444` `_page_mark_kernel[...]` 唯一发射点,位于 `_split_kv_pages_to_64`(`:373` 定义)内
|
||||
- `:515` `_split_kv_pages_to_64(...)` 唯一调用点,位于 `_flash_mla_flashinfer`(`:477` 定义)内
|
||||
- `:607` `flashinfer_sparse_mla_forward` —— **GLM(GlmMoeDsaForCausalLM)的入口**
|
||||
|
||||
`flashinfer_sparse_mla_forward`(`:607-644`)函数体直接调用
|
||||
`flashinfer.mla.trtllm_batch_decode_with_kv_cache_mla`,入参
|
||||
`kv_cache.view(-1, page_size, kv_cache_dim)`、`block_tables=indices.unsqueeze(1)`,
|
||||
**不经过 `_split_kv_pages_to_64`、不碰 `_page_mark_kernel`**。
|
||||
|
||||
即:page-mark / page-split 整套机制是 DSV4-Flash 的 SWA(滑动窗口)页分裂路径
|
||||
(`_flash_mla_flashinfer`)专用;GLM-5.3 走 sparse-MLA 直接入口,与该路径无交集。
|
||||
|
||||
## 证据二(运行时,Triton 零启动)
|
||||
|
||||
`_page_mark_kernel` 是 Triton kernel。若 GLM serving 路径会执行它,容器内必然产生
|
||||
Triton JIT 缓存或 Triton kernel 加载。实测(60.1,方案 D 配置容器):
|
||||
|
||||
- 容器环境无 `TRITON_CACHE_DIR` 覆盖,检查 `/root/.triton`、`/root/.cache/triton`、`/tmp`
|
||||
——GLM serving 全程 **0 个 Triton kernel 编译/加载**
|
||||
- 对比:同机 DSV4-Flash 容器相同路径下有大量 Triton 缓存
|
||||
|
||||
结论:GLM-5.3-NVFP4 的 serving 热路径不 launch 任何 Triton kernel,
|
||||
`_page_mark_kernel` 从未执行 —— 与证据一互相印证。
|
||||
|
||||
## 归档内容
|
||||
|
||||
- `base_flash_mla_sm120.py`:两镜像共有的基础模块(md5 9df4d50d…)
|
||||
- `page_mark_vectorized/`:DSV4 E2 形态的向量化补丁(BLOCK=1024, mask other=-1、
|
||||
idx>=0 幂等写 1、grid=cdiv(N,1024)),对 GLM 无效但留作 DSV4 侧资产
|
||||
- `native_heads/`:P0 优化点,同理判 N/A —— GLM 已原生直通 flashinfer
|
||||
(8/16/32-head 模板齐全),`deepseek_v4.py` 补丁对 GLM 无增益未实验
|
||||
@ -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}")
|
||||
File diff suppressed because it is too large
Load Diff
@ -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
|
||||
@ -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])
|
||||
@ -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",
|
||||
]
|
||||
@ -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);
|
||||
File diff suppressed because it is too large
Load Diff
@ -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",
|
||||
],
|
||||
)
|
||||
@ -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"
|
||||
}
|
||||
@ -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: ...
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@ -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
|
||||
@ -0,0 +1,51 @@
|
||||
# 原始压测数据 — DSV4 优化点迁移实验(60.1 / 6000D-1,2026-09-09)
|
||||
|
||||
口径:i16384 / o512 / cc∈{8,16,32,40,64},输出吞吐 tok/s(bench_corpus.py,
|
||||
PG19 真实语料固定窗口 run-id 9311-9313 池C + 9314/9315 pool-override,
|
||||
radix OFF 输入逐 token 相同)。每臂 3 轮取中位;arm5 回切确认为 1 轮。
|
||||
cc40/64 点 nreq=80/128(沿用方案 D 口径,表注见双场景报告)。
|
||||
|
||||
## 输出吞吐(out tok/s)逐轮 + 中位
|
||||
|
||||
| 臂 | cc8 (r1/r2/r3) | cc16 | cc32 | cc40 | cc64 |
|
||||
|---|---|---|---|---|---|
|
||||
| **0a** nightly-dev-20260828 基线 | 96.38/104.25/**102.22** | 148.08/148.98/**148.21** | 192.32/191.13/**191.13** | 202.00/198.41/**198.41** | 208.07/208.18/**208.18** |
|
||||
| **0b** latest(0.5.19+fi0.6.18) 基线 | 87.85/105.79/**101.51** | 147.71/146.66/**147.64** | 190.64/191.07/**191.07** | 200.54/198.01/**198.01** | 207.16/207.74/**207.61** |
|
||||
| **arm2** autotune+持久缓存 (latest) | 104.78/103.33/**104.78** | 152.53/150.89/**151.08** | 195.03/193.96/**195.03** | 200.63/200.27/**200.63** | 210.01/209.83/**210.01** |
|
||||
| **arm3** PCIe-IPC 12文件包 (latest) | 96.89/98.57/**97.91** | 143.18/143.30/**143.18** | 187.00/187.08/**187.06** | 198.78/195.17/**197.30** | 204.52/204.34/**204.52** |
|
||||
| **arm5** 回切确认 stock-latest(1轮) | 102.57 | 147.01 | 190.94 | 200.58 | 207.80 |
|
||||
|
||||
注:0b cc8 r1=87.85 为该臂首轮冷 JIT(fi_jit_cache/sglang_cache 挂载此臂首次填充);
|
||||
arm5 同为全新容器首轮但缓存已热 → 102.57,直接落在 0b r2/r3 暖机水平。
|
||||
|
||||
## 中位数对比(vs 0b latest 基线)
|
||||
|
||||
| 臂 | cc8 | cc16 | cc32 | cc40 | cc64 | 判定 |
|
||||
|---|---|---|---|---|---|---|
|
||||
| 0a(nightly) | +0.7% | +0.4% | +0.0% | +0.2% | +0.3% | 持平 → 按用户规则转 latest |
|
||||
| **arm2(autotune)** | **+3.2%** | **+2.3%** | **+2.1%** | **+1.3%** | **+1.2%** | **五点全胜,唯一胜者** |
|
||||
| arm3(IPC) | −3.5% | −3.0% | −2.1% | −0.4% | −1.5% | 五点全降,判负 |
|
||||
| arm5(回切确认) | +1.0% | −0.4% | −0.1% | +1.3% | +0.1% | 全落基线带,A/B/A 闭环 |
|
||||
|
||||
分布不重叠性:arm2 vs 0b 在 cc16/32/64 轮间分布不重叠(arm2 min > 0b max);
|
||||
arm3 vs 0b 在 cc8/16/32/64 分布不重叠(方向为劣)。
|
||||
|
||||
## TPOT 副指标(r1 轮 mean,秒/token)
|
||||
|
||||
| 臂 | cc8 | cc16 | cc64 |
|
||||
|---|---|---|---|
|
||||
| 0b 基线 | 0.0592 | 0.0778 | 0.1430 |
|
||||
| arm2 autotune | 0.0565 (−4.6%) | 0.0745 (−4.2%) | 0.1401 (−2.0%) |
|
||||
| arm3 IPC | 0.0628 (+6.1%) | 0.0813 (+4.5%) | 0.1455 (+1.7%) |
|
||||
|
||||
## 显存与 KV 池
|
||||
|
||||
全臂 KV 池恒定 **max_total_num_tokens=1,040,384**(无容量漂移)。
|
||||
latest 镜像下 available_gpu_mem:stock 15.64 GB → autotune 11.79 GB
|
||||
(**autotune tactic 缓冲代价 ≈3.9 GB/卡**;池不缩)。
|
||||
|
||||
## 质量门(quality_gate_605.sh)
|
||||
|
||||
全臂(0a/0b/arm2/arm3)一致 **6/7**:GSM8K×5 + 中文推理全过;
|
||||
tool-call FAIL 为方案 D 无 --tool-call-parser 的配置缺口(生产上需补
|
||||
--tool-call-parser glm47 与 --reasoning-parser glm45),非本实验变量。
|
||||
@ -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"
|
||||
@ -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
|
||||
@ -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"
|
||||
Loading…
x
Reference in New Issue
Block a user