From 91ab3fe32d48ab4fbe108c48fd9e35ae6b2c46ae Mon Sep 17 00:00:00 2001 From: Zhiyi Hong <2497491955@qq.com> Date: Mon, 31 Aug 2026 17:10:00 +0800 Subject: [PATCH] feat: add Kimi-K3 PP8 DFlash PD integration and warmup regression --- README.md | 2 + docs/Kimi-K3_PP与DFlash迁移审计.md | 119 + .../Dockerfile | 32 + .../README.md | 105 + .../bench_gsm8k_acceptance.py | 204 ++ .../config.env | 36 + .../deploy_pd_dflash.sh | 261 ++ .../flashinfer_4460_b460bc0.patch | 739 ++++++ .../pp_dflash_files.json | 81 + .../pp_dflash_integration.patch | 2203 +++++++++++++++++ .../results/model_inventory_20260831.jsonl | 8 + .../metadata/preflight.txt | 8 + .../results/pwarm1_602_cpu.log | 28 + .../results/pwarm1_602_inspect.json | 185 ++ .../results/pwarm1_delta_metadata.json | 93 + 15 files changed, 4104 insertions(+) create mode 100644 docs/Kimi-K3_PP与DFlash迁移审计.md create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/Dockerfile create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/README.md create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/bench_gsm8k_acceptance.py create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/config.env create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/deploy_pd_dflash.sh create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/flashinfer_4460_b460bc0.patch create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_files.json create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_integration.patch create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/model_inventory_20260831.jsonl create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pd-dflash-pwarm1-preflight-20260831/metadata/preflight.txt create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_602_cpu.log create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_602_inspect.json create mode 100644 experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_delta_metadata.json diff --git a/README.md b/README.md index 4948d67..7e8e2fc 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,7 @@ # sskj — 多平台大模型推理性能基准测试项目 +**更新(2026-08-31 17:05:00 CST)**:新增 Kimi-K3 PP8 + DFlash PD 适配验证入口。基于 SGLang PR #33863 固定源码,接通 PP 分段 hidden 投影、P 侧 prompt draft KV 生成、D 侧输入生命周期与 TP4→TP32 的 draft GQA KV 传输,保留 Kimi SM120 FlashInfer MXFP4 接入。修复 P 普通预热误带 DFlash verify metadata 的启动问题,23 项 CPU 回归通过;修复镜像已同步 601–608,服务级验证正在进行,尚无 GSM8K 结果。配置为 P TP4/PP8/EP4、D TP32/PP1/EP4、BF16 KV、8K Chunk,计划固定 64 题 C1/C8。详见 `experiments/pro6000/kimi3_pro6000_pd_dflash_validation/README.md`。 + **更新(2026-08-27 13:53:26 CST)**:完成 Kimi-K3 八节点标准 PD 第一阶段。P 组 601-604 使用 PP8×TP4×EP4、FlashInfer MXFP4、Chunk 8K,D 组 605-608 使用 PP1×TP32×EP32、Marlin,通过 Mooncake 0.3.12.post1 和 4 Rail RDMA 传输;统一 P/D `page_size=64` 后,16K→1 与 16K→512 的 C1/C8 共 91/91 请求成功。代表结果:16K→1 C8 Input TPS 6364.31、TTFT P50/P95 20.555/21.345 秒;16K→512 C8 TPOT P50/P95 63.20/66.55 ms。详见 `experiments/pro6000/kimi3_pro6000_pd_pp8_standard/README.md`。 **更新(2026-08-20 10:13:01 CST)**:完成 Kimi-K3 Prefill TP Reduce Scatter 可行性审计并停止该方向。K3 的 MLA 输出门控仍依赖完整 7168 维 hidden,且 69/93 层为 KDA;保留或恢复 gate hidden 后,原 hidden All-Reduce 无法消除并新增 2112 维 latent All-Gather,估算通信量反增约 14.7%。研究原型仅保留为否决证据,不进入四节点实验或上游 PR;后续转向 MoE A2A 与 Pipeline Parallelism。详见 `experiments/pro6000/kimi3_pro6000_sglang_tp_reduce_scatter_prefill/README.md`。 diff --git a/docs/Kimi-K3_PP与DFlash迁移审计.md b/docs/Kimi-K3_PP与DFlash迁移审计.md new file mode 100644 index 0000000..34cc422 --- /dev/null +++ b/docs/Kimi-K3_PP与DFlash迁移审计.md @@ -0,0 +1,119 @@ +# Kimi-K3:从 PP + DSpark 迁移到 PP + DFlash + +更新:2026-08-31;PP、传输与 MoE 集成首版。 + +## 1. 决策与范围 + +**优先推进 DFlash 的 PP + PD 适配。** #33863 的分段投影设计能复用于 DFlash,但需要补齐 DFlash worker、调度入口和异构 TP 下的 draft KV 传输。 + +当前已实现 DFlash PP worker、Kimi 跨 stage capture、PD 输入衔接、异构 TP 的 draft KV 传输,以及 Kimi FlashInfer 布局/SiTU 接入。首轮 P 启动暴露普通预热误带 verify metadata 的问题,已修复并增加回归,23 项 CPU 测试通过。SM120 编译及 Kimi SiTU GPU 集成回归已通过。修复镜像已同步全部八节点,Run `pd-dflash-pwarm1-20260831-1652` 正在重新验证服务启动;尚无 PD 请求和 GSM8K 结果。详见 [实现进度与命令](evidence/kimi_k3_pd_dflash/README.md)。 + +暂停 EAGLE3 baseline。KV cache 保持 BF16,chunk 保持 8192。FlashInfer 使用官方 0.6.18 加已合并 #4460 的显式 backport;原版 0.6.18 wheel 尚未包含所需 CUTLASS SiTU 接口。 + +目标拓扑:P 使用 601–604 的 TP4/PP8/EP4 与已适配的 FlashInfer MXFP4;D 使用 605–608 的 TP32/PP1,接入 DFlash。这里适配的是 **P 侧 PP 生成 draft 上下文**,D 侧投机执行仍为 PP1。 + +审计固定版本: + +| 对象 | 版本 | +|---|---| +| SGLang main | `3139ceaeec50868a441d68dc231663f3777e0d93` | +| PR #33863 head | `6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5` | +| PR 状态 | Open;维护者要求拆分 DSv4 PP/PD、DSpark、Kimi/线性注意力工作 | + +PR 公布的准确性数据来自 DSv4 Flash / H20,不能代替本次 Kimi-K3 / 6000D / P-TP4 与 D-TP32 验证。[PR #33863](https://github.com/sgl-project/sglang/pull/33863) + +## 2. 为什么可以迁移 + +DSparkDraftModel 继承 DFlashDraftModel。普通 Kimi draft 的 prompt 上下文处理都包含:采集多个 target 层的 hidden,拼接后线性投影,再归一化,生成 draft KV。 + +#33863 将这一步改为: + +```text +各 P/PP stage:采集自己负责的 hidden → 使用对应权重列做局部投影 +PP stage 间:传递并累加投影结果 [token 数, 7168] +最后一个 P/PP stage:统一 RMSNorm → 各 draft 层 K/V 投影、K norm、RoPE → 写 draft KV +P → D:传 target KV、KDA 状态及 draft KV +D:使用已有 DFlash proposer / verify / commit 流程 +``` + +依据是线性运算 `concat(h_i) × W^T = sum(h_i × W_i^T)`。归一化必须在求和后做。BF16 分段累加会改变舍入顺序,因此需要比较中间张量、logits 和 token 接受行为。 + +关键接口已经加在共用 DFlash 模型中:[project_target_hidden_partial](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/models/dflash.py#L683)。最后阶段的归一化及 KV 写入实现位于 [write_projected_context_kv](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/models/dspark.py#L751)。 + +这一路线保留 P 侧 PP8 的收益,不要求把各层原始 hidden 通过 PD 网络传给 D。 + +## 3. 需要补哪些代码 + +以下路径均相对于 SGLang 的 `python/sglang/srt/`。 + +| 模块 | 已有基础 | 本次增量 | +|---|---|---| +| `models/kimi_k3.py` | PR 已支持各 PP stage 局部采集 DSpark 特征 | 接通 DFlash capture setter,按 DFlash checkpoint 的层号及 Kimi residual 语义采集 | +| `models/dflash.py` | PR 已有局部线性投影 | 复用投影;最后阶段归一化后复用 DFlash 的逐层 KV 写入路径 | +| `speculative/dflash_worker_v2.py` | 非 PP prefill 的 hidden → draft KV | 接收、转发 PP proxy;仅最后阶段写完整 draft KV;处理非末 stage 没有 logits/next token 的情况 | +| `managers/scheduler_pp_mixin.py` | PR 的成功交集、失败并集及一致释放机制 | 将目前 DSpark 专用的 proxy/input 衔接扩展到 DFlash,保持请求顺序和完成条件一致 | +| `speculative/spec_info.py` | 当前 PD 输入分派只有 EAGLE、DSpark | 接入 DFlash 输入构造,并验证 overlap FutureMap 与 idle batch 生命周期 | +| `disaggregation/prefill.py`、`utils.py`、`mooncake/conn.py` | 最后 PP stage 传 draft KV、按层号配对 | 加入 draft GQA 的 head 分片/复制与收发端独立 stride,保留 target MLA/KDA 的已有路径 | +| `arg_groups/speculative_hook.py`、`validation_hook.py` | 当前 main 拒绝 PP + DFlash | 仅放开已实现的 PD-prefill + PP 组合;D 端 PP1 范围不变 | + +关键证据: + +- [Kimi 局部层采集](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/models/kimi_k3.py#L2912)。 +- [DSpark PP 上下文处理](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py#L655)。 +- [PP scheduler 的 DSpark 专用入口](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/managers/scheduler_pp_mixin.py#L1391)。 +- [main 的 PD 输入分派](https://github.com/sgl-project/sglang/blob/3139ceaeec50868a441d68dc231663f3777e0d93/python/sglang/srt/speculative/spec_info.py#L174):DFlash 仍返回 None。仓库里的 `dflash_disaggregation.py` 辅助函数尚未被这个入口调用,不能仅凭文件存在认定 PD 已支持。 + +相关 [#36140](https://github.com/sgl-project/sglang/issues/36140) 记录了输入缺失和 draft 状态缺失问题;[#36277](https://github.com/sgl-project/sglang/pull/36277) 是提前拒绝不支持配置的保护性改动,不是 DFlash PD 的实现。 + +## 4. 异构 TP 的明确风险 + +605 的 `/data/hf_models/Kimi-K3-DFlash/config.json` 为 6 层 draft、8 个 KV head、head_dim=128、BF16、4096 sliding window;target capture 层是 `[19,37,54,66,78,90]`。 + +[DFlashAttention](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/models/dflash.py#L142) 使用 `max(1, total_kv_heads // tp_size)`: + +| draft 布局 | 每 rank KV heads | BF16、page=64 时单层 K 的页大小 | +|---|---:|---:| +| P TP4 | 2 | `64 × 2 × 128 × 2 = 32768 bytes` | +| D TP32 | 1,按 head 复制到多个 rank | `64 × 1 × 128 × 2 = 16384 bytes` | + +PR 当前 [Mooncake flat 传输分支](https://github.com/sgl-project/sglang/blob/6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5/python/sglang/srt/disaggregation/mooncake/conn.py#L686) 按 layer ID 配对后,对源偏移、目标偏移和复制字节数都使用源端 `item_len`。Kimi hybrid MLA 会进入这条路径。层号正确只能确认“哪层到哪层”,没有解决“该层哪些 head 到哪个 rank”。 + +### CPU 验证结果 + +从该 PR 源码直接用 AST 提取 `_send_kvcache_generic` 和 `build_transfer_entry_pairs`,使用 synthetic 地址、单页索引和记录型传输函数执行。未调用 GPU、Mooncake 或 RDMA。 + +[地址规划结果与源文件 SHA256](evidence/kimi_k3_pp_dflash_pr33863_transfer_plan.json)。 + +| 用例 | 地址/长度是否符合目标布局 | +|---|---| +| 收发双方均为 1 KV head | 符合 | +| P 为 2 head、D 为 1 head | 不符合:目标偏移和复制长度均按源端大条目计算 | + +具体用例:D buffer 起始地址设为 2000000,目标 page ID=2。正确目标地址为 `2000000 + 2×16384 = 2032768`,当前函数生成 `2065536`,复制 32768 bytes 而不是单 head 页的 16384 bytes。 + +这是特定布局的地址规划复现,尚未启动完整服务复现。迁移时必须加入每条目的布局校验,并实现 GQA head-aware 传输;不能直接复用 flat copy。还需确认 D 接收到全部 target/draft 组件后才进入首轮 draft。 + +## 5. 实现与验收顺序 + +1. **局部数学与协议验证。** 检查各 stage capture 层覆盖、feature 顺序、空 capture stage、单次 RMSNorm;比较拼接投影与分段投影误差。测试 TP4→TP32 的全部 8 个 draft KV head、K/V、6 层、页索引以及失败清理。源/目标布局不兼容时在传输前报错。 +2. **最小 PD 请求。** 保留 P TP4/PP8、D TP32/PP1,先 C1。核对真实 buffer metadata、draft KV 到达情况、首轮 draft logits,以及成功/失败后的资源释放。 +3. **GSM8K 固定 64 题。** 相同 prompt/template、采样设置,依次测试 C1/C8;保留准确率、完成率、接受长度、接受直方图、Output TPS、总耗时与显存状态。当前客户端非流式,不把总耗时换算成 TTFT/ITL。输入不裁剪,输出上限 512,报告截断,不新增 16K synthetic 测试。历史 5-shot 与前 64 题有重叠,本轮沿用以便部署回归,不能作为独立无泄漏的模型准确率评估。 +4. **验收后决定继续投入。** DFlash 正确性与部署通过且有实际 Decode 收益,继续 PD 调优;若关键适配无法通过,或接受行为同样异常,切换到下面的 D-only 诊断。 + +真实实验代码继续维护在 601 的 `/data/hzy/sskj` 工作区,不改同事部署目录,不再创建分支。本轮核实当前检出名为 `hzy-kimi3-pd-pp8-standard`,保持现状;独立 SGLang 源码位于 `/data/hzy/src/sglang-kimi-pp-dflash-33863`。 + +## 6. DSpark 接受长度诊断的备用路径 + +DSpark 低接受长度的问题先独立于 PD 排查。PP + PD 适配本身不会自动修复 D-only 已存在的问题。 + +| 检查项 | 目的 | +|---|---| +| checkpoint、tokenizer、mask token、target_layer_ids、RoPE | 排除 draft/target 配置失配及错误 hidden 来源 | +| 相同 token 前缀下,逐位置 draft token 与 target logits | 找到拒绝从第几个位置开始,区分首 token 错位与持续预测质量问题 | +| verify 的 KDA/SSM 状态更新、回滚与 token 位置 | 判断是否首次可用、随后状态漂移 | +| fused/replay 路径与参考执行比较 | 将 kernel/状态管理错误与 drafter 本身质量分开 | +| 接受统计定义、实际 proposal 数、bonus token | 统一接受长度口径,防止指标解释错误 | + +605 的 DSpark checkpoint 为 block=7、5 层、capture `[7,23,51,67,83]`;DFlash 为 block=16、6 层、另一组 capture 与 RoPE。二者的原始平均接受长度不能单独判断实现异常,需同时看每次实际提议数、逐位置接受比例和实际吞吐。 + +本轮不重跑 no-spec 性能基线,保留此前结果作为参考。当前尚未确定 DSpark 接受长度异常的根因。 diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/Dockerfile b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/Dockerfile new file mode 100644 index 0000000..fe34edd --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/Dockerfile @@ -0,0 +1,32 @@ +FROM lmsysorg/sglang@sha256:28e0d26073161e49ca56eba808d264a4223804a212020f1dfe1b2405b9f8a399 + +# Use the official core and compile SM120 CUTLASS through its supported JIT path. +# Remove old companion packages instead of suppressing their version checks. +COPY results/official_flashinfer_0618/flashinfer_python-0.6.18-py3-none-any.whl /opt/kimi-dflash-wheels/ +RUN python3 -m pip uninstall -y flashinfer-jit-cache flashinfer-cubin && \ + python3 -m pip install --no-index --no-deps --force-reinstall \ + /opt/kimi-dflash-wheels/flashinfer_python-0.6.18-py3-none-any.whl + +# #4460 is merged on main but absent from the v0.6.18 release branch. +# Apply the unchanged official commit to Python and wheel-bundled C++ sources. +COPY flashinfer_4460_b460bc0.patch /opt/kimi-dflash-wheels/ +RUN cd "$(python3 -c 'import sysconfig; print(sysconfig.get_paths()["purelib"])')" && \ + git apply --check --include='flashinfer/**' /opt/kimi-dflash-wheels/flashinfer_4460_b460bc0.patch && \ + git apply --check --directory=flashinfer/data --include='flashinfer/data/csrc/**' /opt/kimi-dflash-wheels/flashinfer_4460_b460bc0.patch && \ + git apply --include='flashinfer/**' /opt/kimi-dflash-wheels/flashinfer_4460_b460bc0.patch && \ + git apply --directory=flashinfer/data --include='flashinfer/data/csrc/**' /opt/kimi-dflash-wheels/flashinfer_4460_b460bc0.patch + +COPY results/sglang-kimi-pp-dflash-33863-integrated.tar.gz /opt/kimi-dflash-source.tar.gz +RUN mkdir -p /opt/kimi-dflash && \ + tar -xzf /opt/kimi-dflash-source.tar.gz -C /opt/kimi-dflash + +ENV PYTHONPATH=/opt/kimi-dflash/python +ENV PYTHONDONTWRITEBYTECODE=1 +ENV FLASHINFER_DISABLE_VERSION_CHECK="" +ENV SGLANG_SOURCE_ROOT=/opt/kimi-dflash + +# Build has no GPU: check the API, not SGLang's CUDA-device availability predicate. +RUN python3 -c "import inspect, flashinfer; from flashinfer.fused_moe import cutlass_fused_moe; from flashinfer.tllm_enums import ActivationType; from sglang.srt.speculative.dflash_worker_v2 import DFlashWorkerV2; from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager; assert hasattr(ActivationType, 'Situ'); assert {'situ_beta', 'situ_linear_beta'} <= set(inspect.signature(cutlass_fused_moe).parameters); print('RUNTIME_IMPORT_OK', flashinfer.__version__, 'upstream SiTU backport b460bc0')" +RUN python3 -m unittest discover -s /opt/kimi-dflash/test/registered/disaggregation -v + +ENTRYPOINT ["python3", "-m", "sglang.launch_server"] diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/README.md b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/README.md new file mode 100644 index 0000000..c33d265 --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/README.md @@ -0,0 +1,105 @@ +# Kimi-K3 PP + DFlash 适配进度 + +2026-08-31:PP worker、PD 输入衔接、异构 TP 传输和 Kimi FlashInfer 接入已完成首版;SM120 编译及 Kimi SiTU GPU 集成回归通过。首轮八节点部署在 P 侧预热阶段失败,已修复普通预热误带 DFlash verify metadata 的问题;本地、601 构建镜像及 602 增量导入容器的 23 项 CPU 回归通过。正在重新验证服务启动,尚无 PD 请求或 GSM8K 结果。 + +## 源码和环境 + +- 基于官方 PR #33863 的 `6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5`。 +- 上游 Python 目录的 4619 个文件均与 Git blob SHA1 校验一致。 +- `pp_dflash_integration.patch` 包含完整 SGLang 增量;`pp_dflash_files.json` 记录文件 SHA256。`mixed_kv_transfer.patch` 保留传输层首版记录。 +- 601 独立源码路径:`/data/hzy/src/sglang-kimi-pp-dflash-33863`。 +- 601 实验路径:`/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_pd_dflash_validation`。 +- 不修改原 EAGLE3/DSpark 部署;不创建新分支。 + +## 本轮实现 + +| 部分 | 实现 | +|---|---| +| Kimi hidden 采集 | 接入 DFlash checkpoint 的 capture 层;PP 边界由下一 stage 使用实际 residual 聚合权重采集 | +| 分段投影 | 各 stage 计算对应权重列的局部投影,通过 PP proxy 累加;末 stage 统一 RMSNorm 并写 6 层 draft KV | +| P 侧 draft 执行 | 非末 stage 只保留最小 KV pool;P 不执行 draft decode CUDA Graph,D 的原有初始化路径保留 | +| Scheduler 与 PD | 接通 DFlash proxy、next draft input、FutureMap 发布与 idle 生命周期;仅放开 P 侧 PP | +| Mooncake 传输 | 注册每 entry 的 dtype/head/page stride/容量;target MLA 整页传输,draft GQA 按 head 交集逐 token 切片;写入前完成全部边界检查 | +| FlashInfer | 迁移已验证的 Kimi SM120 布局与 SiTU 参数适配,同时保留新上游 SwigluStep 行为 | + +目标范围:Kimi hybrid MLA + 普通 NHD DFlash draft KV、Mooncake、CP=DCP=1,无 staging/unified KV/HiSparse。P TP4/PP8/EP4;D TP32/PP1;BF16 KV;chunk=8192。 + +## FlashInfer 依赖 + +官方 `#4460` 已于 2026-08-21 合并,提交为 `b460bc00cb373541102d2155aec35bd626e522ce`。但 `v0.6.18` 属于另一条发布分支,实际下载的官方 wheel 没有 CUTLASS `situ_beta/situ_linear_beta` API。 + +因此本镜像使用 **官方 0.6.18 wheel + 未改写的 #4460 合并补丁**。Dockerfile 对 2 个 Python 和 5 个 C++ 文件先执行 `git apply --check`,再应用上游补丁。它不是未经修改的官方 0.6.18,也不是旧的私人 SiTU kernel。 + +移除镜像内旧版 `flashinfer-cubin`、`flashinfer-jit-cache`,保留版本校验;SM120 CUTLASS 通过官方 JIT 路径编译。[安装说明](https://docs.flashinfer.ai/installation.html)、[#4460](https://github.com/flashinfer-ai/flashinfer/pull/4460)。 + +## CPU 验证 + +```bash +# 独立传输测试,无需 GPU 或 torch +cd /data/hzy/src/sglang-kimi-pp-dflash-33863 +python3 test/registered/disaggregation/test_mixed_kv_entry_layout.py + +# 全部 CPU 回归需要带 torch 的容器;镜像构建时自动执行 +python3 -m unittest discover -s test/registered/disaggregation -v +``` + +| 回归 | 数量 | 主要覆盖 | +|---|---:|---| +| 传输 | 7 | 144 组 KV-head/TP 组合、实际 P-TP4→D-TP32 全 32 rank、逐字节复制、越界/重复写拒绝、注册往返 | +| PP 与 PD 输入 | 12 | PP1/2/4/8/16 capture 归属、边界 residual、空 capture stage、BF16/FP32 分段投影、单次 norm、末 stage 写 KV、输入生命周期、P 普通预热不创建 verify metadata | +| Kimi MoE 合并 | 4 | 两类 gate/up 布局、SiTU 参数、activation 白名单、非连续输入、保留 SwigluStep、API 能力检查 | + +23 项在本地、601 修复镜像和 602 导入容器中通过,并完成真实 DFlashWorkerV2、MooncakeKVManager 导入。8 条启动命令的 CLI 解析及两类服务的参数后处理已检查。CPU 测试对实际方法作 AST 提取,使用 CPU tensor 或记录型 engine,核验数学与调用契约,不替代服务级验证。 + +601 GPU6 的 `test_kimi_k3_sm120_situ_layout_and_noncontiguous_input` 已通过:检查 Kimi gate/up 及 scale 布局、SiTU(4,25)、非连续输入,以及 SGLang adapter 与直接 FlashInfer 调用的输出一致性。这是小 shape 的集成回归;端到端正确性由后续 PD/GSM8K 检验。 + +修复前镜像 `local/sglang:kimi-k3-pp-dflash-33863-fi0618-situ4460` 的独立 GPU 复测通过,Mooncake CUDA engine 导入成功,GPU 测试 1 passed、25.28 秒(复用 JIT 缓存)。 + +当前镜像为 `local/sglang:kimi-k3-pp-dflash-33863-fi0618-situ4460-pwarm1`,只追加 P 预热修复,不改变 GPU kernel。Linux/amd64 manifest 为 `sha256:281ccb2666a38acd539e6fbb5d55d3682eb2af9ac9089f67e9bf92a8ddd822eb`,image config 为 `sha256:8ab5eee9902556bcba2a2b439a9fa1b14d93e0554503754fb7ea74dad6c3cf79`;OCI index 为 `sha256:d30d68757014b973c2edd5d061cc3fd55660a4c68cc3d0af9050fcb7e64c4a78`。区分这三个摘要,不将 Docker 不同模式显示的 ID 当成代码不一致。 + +601 证据均位于实验目录 `results/`: + +- `pp_and_transport_cpu_601_20260831.log`:18 项 CPU 回归。 +- `patched_worker_import_20260831.log`、`patched_launch_help_20260831.log`:真实模块导入与 CLI。 +- `image_build_situ4460_cpucheck_20260831.log`:集成镜像的完整 API/回归检查。 +- `command_parse_20260831.log`:P/D 全部 8 条命令通过 `ServerArgs` argparse 检查;该检查不执行 ServerArgs 后处理或服务初始化。 +- `args_resolve_20260831.log`、`resolved_p_20260831.json`、`resolved_d_20260831.json`:P/D 真实模型配置通过 `resolve_once()` 后处理;保持 BF16 KV、8K chunk 和 FlashInfer。Kimi 投机验证自动选择 `nv_cutedsl`,P 侧 PP 自动关闭 overlap scheduler;该检查不加载模型权重。 +- `sm120_cutlass_compile_20260831.log`:SM120 预编译进度。首次构建被外部 Docker SIGKILL 终止,退出 137;Docker 事件为显式 kill,无 OOM 事件。已编译的对象文件保留在独立缓存。 +- `sm120_cutlass_compile_resume_20260831.log`:恢复后编译成功,`SM120_CUTLASS_BUILD_OK`,退出 0。 +- `gpu_situ_smoke_detached_20260831.log`:首次完整单测发现 SiTU activation 白名单遗漏;已修复并增加 CPU 回归。 +- `gpu_situ_smoke_fix_20260831.log`:修复后 GPU 回归通过,1 passed;首次 JIT 在内的总用时 771.10 秒。 +- `image_build_activation_fix_20260831.log`:固化修复后的镜像构建及 22 项 CPU 回归。 +- `final_image_gpu_verify_20260831.log`、`final_image_gpu_container_20260831.json`:最终镜像独立 GPU 复测、Mooncake 导入、镜像 ID 和挂载证据。 +- `pd-dflash-20260831-1617/logs/p_0.log` 至 `p_3.log`:首轮 P 预热报错原始证据;该 Run 未进入请求测试,退出码 1。 +- `image_build_pwarm1_20260831.log`:预热修复镜像构建及 23 项 CPU 回归。 +- `pwarm1_602_cpu.log`、`pwarm1_602_inspect.json`:约 49 MiB 增量包导入后的容器回归与平台摘要。 + +已通过 compileall、Black 和 Ruff 的未定义变量/语法检查。 + +## 实验入口 + +代码位于 601 的 `/data/hzy/sskj/experiments/pro6000/kimi3_pro6000_pd_dflash_validation`。只运行 `deploy_pd_dflash.sh`,它负责 P/D 启动、健康检查、Router、smoke 和评测。`bench_gsm8k_acceptance.py` 沿用原客户端的 prompt 与请求设置,将结果标签改为 DFlash,并增加原始 `meta_info`、verify 次数和输出结束原因归档。 + +```bash +# 601:命令检查,不需要 GPU 或 sudo +cd /data/hzy/sskj/experiments/pro6000/kimi3_pro6000_pd_dflash_validation +DRY_RUN=1 RUN_ID=command-audit-20260831 bash deploy_pd_dflash.sh all +``` + +首轮已执行 `all`,在 P 侧启动阶段退出。修复镜像同步完成后使用新 Run ID 重试。正式顺序为 `start` → `smoke` → `bench` → `logs` → `stop`,也可用 `all` 串行执行。所有操作使用同一个 `RUN_ID`。失败时保留本任务容器和日志,入口不会自动杀其他实验。 + +配置选择:P 为 TP4/PP8/EP4,D 为 TP32/PP1/EP4,两侧 FlashInfer;KV 为 BF16,chunk=8192,page=64。为 C1/C8 评测将活跃请求与 Decode Graph 上限设为 8;未额外改变模型 context 上限。 + +GSM8K 为原数据集前 64 题,沿用历史客户端的 5-shot、temperature=0、输出上限 512,分别运行 C1/C8。输入不裁剪。历史 5-shot 也取自测试集前 5 题,和本次 64 题有重叠:结果适用于与历史流程的部署回归,不作为独立无泄漏的模型准确率评估。逐题输出和截断情况保留。 + +## 后续 + +确认 8 节点模型、镜像与网络一致,然后运行 C1 smoke 和固定 GSM8K 64 题 C1/C8。保留逐题答案、截断、接受直方图、总耗时、吞吐、原始 `meta_info` 与服务日志;不重跑 EAGLE3/no-spec 基线。当前客户端为非流式请求,不把客户端总耗时当作 TTFT 或 ITL。 + +运行时设置 `SGLANG_CACHE_DIR=/cache`,FlashInfer、Triton、PyTorch 扩展和 CUDA 编译缓存也指向该挂载目录,`TMPDIR=/cache/tmp`,宿主路径见 `config.env` 的 `JIT_CACHE`,避免容器可写层占满根盘;预检会创建所需临时目录。 + +## 首轮启动修复 + +`base_runner.py::_dummy_run` 已将 PD Prefill target 设为普通 Decode 预热,但随后仍创建 `DFlashVerifyInput`。普通 Triton Attention 读取 `kv_indptr` 时因此报 `AttributeError`。修复让该分支的 `spec_info=None`,D 侧真正的 TARGET_VERIFY 路径不变。没有强行添加字段,也没有修改 DFlash 接受算法。 + +镜像同步须检查 Docker 所在根盘,而不只检查模型盘 `/data`。本镜像层展开约 34.4GB,压缩内容约 15GB;共享层会减少增量占用。607 已按用户授权删除两个无容器引用的 vLLM 镜像,根盘恢复到约 107GB,原镜像元数据保存在 `results/607_vllm_images_before_delete_20260831.json`。用户清理 606 后其根盘恢复到约 91GB。601 使用 `ctr images export` 将既有 OCI 压缩层直接导出到 `/data`,其旧 vLLM 镜像尚未删除。 diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/bench_gsm8k_acceptance.py b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/bench_gsm8k_acceptance.py new file mode 100644 index 0000000..59202d3 --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/bench_gsm8k_acceptance.py @@ -0,0 +1,204 @@ +#!/usr/bin/env python3 +"""Official SGLang GSM8K prompt semantics with speculative telemetry.""" + +import argparse +import ast +import json +import re +import time +from pathlib import Path + +import numpy as np +import sglang as sgl +from sglang.lang.api import set_default_backend +from sglang.test.test_utils import ( + add_common_sglang_args_and_parse, + select_sglang_backend, +) +from sglang.utils import read_jsonl + +INVALID = -9999999 + + +def get_one_example(lines, index, include_answer): + text = "Question: " + lines[index]["question"] + "\nAnswer:" + if include_answer: + text += " " + lines[index]["answer"] + return text + + +def get_answer_value(answer): + numbers = re.findall(r"\d+", answer.replace(",", "")) + if not numbers: + return INVALID + try: + return ast.literal_eval(numbers[-1]) + except (SyntaxError, ValueError): + return INVALID + + +def histogram_accept_length(histogram): + if not histogram: + return None + if isinstance(histogram, dict): + pairs = ((int(key), int(value)) for key, value in histogram.items()) + else: + pairs = enumerate(histogram) + total_steps = 0 + accepted_drafts = 0 + for accepted, count in pairs: + total_steps += count + accepted_drafts += accepted * count + return 1.0 + accepted_drafts / total_steps if total_steps else None + + +def main(args): + set_default_backend(select_sglang_backend(args)) + lines = list(read_jsonl(args.data_path)) + few_shot = "".join( + get_one_example(lines, index, True) + "\n\n" + for index in range(args.num_shots) + ) + questions = [] + labels = [] + for index in range(args.num_questions): + prompt = few_shot + get_one_example(lines, index, False) + questions.append({"question": prompt}) + labels.append(get_answer_value(lines[index]["answer"])) + + @sgl.function + def few_shot_gsm8k(s, question): + s += question + s += sgl.gen( + "answer", + max_tokens=args.max_new_tokens, + stop=["Question", "Assistant:", "<|separator|>"], + ) + + start = time.perf_counter() + states = few_shot_gsm8k.run_batch( + questions, + temperature=args.temperature, + top_p=args.top_p, + num_threads=args.parallel, + progress_bar=True, + ) + duration = time.perf_counter() - start + + rows = [] + for index, state in enumerate(states): + state_error = state.error() + if state_error is not None: + rows.append( + { + "prompt_id": index, + "output": None, + "correct": False, + "error": repr(state_error), + "completion_tokens": 0, + "spec_accept_length": None, + "spec_accept_length_from_histogram": None, + "spec_accept_rate": None, + "spec_accepted_drafts": None, + "spec_proposed_drafts": None, + "spec_accept_histogram": None, + } + ) + continue + output = state["answer"] + metadata = state.get_meta_info("answer") + reported = metadata.get("spec_accept_length") + histogram = metadata.get( + "spec_correct_drafts_histogram", metadata.get("spec_accept_histogram") + ) + reconstructed = histogram_accept_length(histogram) + rows.append( + { + "prompt_id": index, + "output": output, + "correct": get_answer_value(output) == labels[index], + "error": None, + "completion_tokens": metadata.get("completion_tokens"), + "spec_accept_length": reported, + "spec_accept_length_from_histogram": reconstructed, + "spec_accept_rate": metadata.get("spec_accept_rate"), + "spec_accepted_drafts": metadata.get("spec_accepted_drafts"), + "spec_proposed_drafts": metadata.get("spec_proposed_drafts"), + "spec_accept_histogram": histogram, + "spec_verify_ct": metadata.get("spec_verify_ct"), + "finish_reason": metadata.get("finish_reason"), + "meta_info": metadata, + } + ) + + Path(args.output_file).write_text( + "\n".join(json.dumps(row, ensure_ascii=False) for row in rows) + "\n", + encoding="utf-8", + ) + successful_rows = [row for row in rows if row["error"] is None] + failed_rows = [row for row in rows if row["error"] is not None] + accept_lengths = [row["spec_accept_length"] for row in successful_rows] + if not successful_rows: + raise RuntimeError("all GSM8K requests failed") + if any(value is None for value in accept_lengths): + raise RuntimeError("speculative acceptance metadata is missing from responses") + total_output_tokens = sum( + row["completion_tokens"] or 0 for row in successful_rows + ) + summary = { + "questions": len(rows), + "successful_requests": len(successful_rows), + "failed_requests": len(failed_rows), + "failed_prompt_ids": [row["prompt_id"] for row in failed_rows], + "num_shots": args.num_shots, + "max_new_tokens": args.max_new_tokens, + "parallel": args.parallel, + "temperature": args.temperature, + "top_p": args.top_p, + "prompt_format": "sglang_official_raw_five_shot", + "speculative_algorithm": args.speculative_algorithm, + "length_limited_requests": sum( + isinstance(row.get("finish_reason"), dict) + and row["finish_reason"].get("type") == "length" + for row in successful_rows + ), + "accuracy_all_questions": float( + np.mean([row["correct"] for row in rows]) + ), + "accuracy_successful_requests": float( + np.mean([row["correct"] for row in successful_rows]) + ), + "mean_accept_length_equal_weight_per_question": float( + np.mean(accept_lengths) + ), + "median_accept_length_per_question": float(np.median(accept_lengths)), + "min_accept_length_per_question": float(np.min(accept_lengths)), + "max_accept_length_per_question": float(np.max(accept_lengths)), + "duration_s": duration, + "output_throughput": total_output_tokens / duration, + } + Path(args.summary_file).write_text( + json.dumps(summary, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + print(json.dumps(summary, ensure_ascii=False, indent=2)) + if failed_rows: + raise RuntimeError( + f"{len(failed_rows)} of {len(rows)} GSM8K requests failed; " + f"see {args.output_file}" + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--data-path", required=True) + parser.add_argument("--num-shots", type=int, default=5) + parser.add_argument("--num-questions", type=int, default=64) + parser.add_argument("--max-new-tokens", type=int, default=512) + parser.add_argument("--temperature", type=float, default=0.0) + parser.add_argument("--top-p", type=float, default=1.0) + parser.add_argument("--speculative-algorithm", default="DFLASH") + parser.add_argument("--output-file", required=True) + parser.add_argument("--summary-file", required=True) + args = add_common_sglang_args_and_parse(parser) + main(args) diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/config.env b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/config.env new file mode 100644 index 0000000..9e4183b --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/config.env @@ -0,0 +1,36 @@ +EXPERIMENT=kimi3_pro6000_pd_dflash_validation +PD_IMAGE=local/sglang:kimi-k3-pp-dflash-33863-fi0618-situ4460-pwarm1 +MODEL_PATH=/data/hf_models/Kimi-K3 +DRAFT_MODEL_PATH=/data/hf_models/Kimi-K3-DFlash +P_NODES=(174.1.60.1 174.1.60.2 174.1.60.3 174.1.60.4) +D_NODES=(174.1.60.5 174.1.60.6 174.1.60.7 174.1.60.8) +SSH_USER=user +SSH_OPTIONS=(-o BatchMode=yes -o ConnectTimeout=10) +PORT=30000 +DIST_PORT=20000 +BOOTSTRAP_PORT=28800 +ROUTER_PORT=31000 +P_TP=4 +P_PP=8 +P_EP=4 +D_TP=32 +D_PP=1 +D_EP=4 +P_MEM=0.88 +D_MEM=0.86 +P_MAMBA_RATIO=0.36 +D_MAMBA_RATIO=0.21 +CHUNK=8192 +PAGE_SIZE=64 +KV_DTYPE=bfloat16 +DRAFT_TOKENS=16 +# GSM8K C1/C8 only; cap capture and scheduling to the measured concurrency. +MAX_RUNNING=8 +GRAPH_BS=8 +IB_DEVICES=mlx5_0,mlx5_1,mlx5_2,mlx5_3 +JIT_CACHE=/data/hzy/cache/kimi-dflash-fi0618-situ4460 +HEALTH_WAIT_S=2400 +GSM8K="${SCRIPT_DIR}/../../../datasets/gsm8k/test.jsonl" +GSM8K_SHA256=3730d312f6e3440559ace48831e51066acaca737f6eabec99bccb9e4b3c39d14 +# Reuse the established client and its prompt/telemetry definitions. +BENCH_CLIENT="${SCRIPT_DIR}/bench_gsm8k_acceptance.py" diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/deploy_pd_dflash.sh b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/deploy_pd_dflash.sh new file mode 100644 index 0000000..32f16d9 --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/deploy_pd_dflash.sh @@ -0,0 +1,261 @@ +#!/usr/bin/env bash +set -Eeuo pipefail +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +source "${SCRIPT_DIR}/config.env" +ACTION="${1:-status}" +DRY_RUN="${DRY_RUN:-0}" +RUN_ID="${RUN_ID:-kimi3-pd-dflash-$(date +%Y%m%d-%H%M%S)}" +[[ "$RUN_ID" =~ ^[a-zA-Z0-9_.-]+$ ]] || exit 2 +RESULT_ROOT="${SCRIPT_DIR}/results/${RUN_ID}" +SUDO_PASSWORD_FILE="${SUDO_PASSWORD_FILE:-/data/hzy/.sudo_password}" +ALL_NODES=("${P_NODES[@]}" "${D_NODES[@]}") +mkdir -p "${RESULT_ROOT}"/{commands,logs,bench,metadata} + +remote() { + local host="$1" command + shift + if [[ "$host" == "${P_NODES[0]}" ]]; then + "$@" + else + printf -v command '%q ' "$@" + ssh "${SSH_OPTIONS[@]}" "${SSH_USER}@${host}" "$command" + fi +} + +root() { + local host="$1" password command + shift + IFS= read -r password <"$SUDO_PASSWORD_FILE" + if [[ "$host" == "${P_NODES[0]}" ]]; then + printf '%s\n' "$password" | sudo -S -p '' -- "$@" + else + printf -v command '%q ' "$@" + printf '%s\n' "$password" | ssh "${SSH_OPTIONS[@]}" \ + "${SSH_USER}@${host}" "sudo -S -p '' -- ${command}" + fi +} + +name() { printf '%s_%s_%s' "$EXPERIMENT" "$1" "$2"; } + +build_command() { + local group="$1" rank="$2" tp pp ep mem ratio head + if [[ "$group" == p ]]; then + HOST="${P_NODES[$rank]}"; head="${P_NODES[0]}" + tp="$P_TP"; pp="$P_PP"; ep="$P_EP"; mem="$P_MEM"; ratio="$P_MAMBA_RATIO" + MODE=prefill + else + HOST="${D_NODES[$rank]}"; head="${D_NODES[0]}" + tp="$D_TP"; pp="$D_PP"; ep="$D_EP"; mem="$D_MEM"; ratio="$D_MAMBA_RATIO" + MODE=decode + fi + CMD=(docker run -d --name "$(name "$group" "$rank")" + --gpus all --network host --ipc=host --ulimit memlock=-1 + --device /dev/infiniband + -v "${MODEL_PATH}:${MODEL_PATH}:ro" + -v "${DRAFT_MODEL_PATH}:${DRAFT_MODEL_PATH}:ro" + -v "${JIT_CACHE}:/cache" + -e "SGLANG_HOST_IP=${HOST}" -e PYTHONUNBUFFERED=1 + -e HF_HUB_OFFLINE=1 -e TRANSFORMERS_OFFLINE=1 + -e GLOO_SOCKET_IFNAME=bond0 -e NCCL_SOCKET_IFNAME=bond1 + -e "NCCL_IB_HCA=${IB_DEVICES}" -e NCCL_IB_GID_INDEX=3 + -e NCCL_IB_TIMEOUT=22 -e NCCL_IB_RETRY_CNT=7 -e NCCL_CUMEM_ENABLE=1 + -e SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK=0 + -e SGLANG_MOE_FUSED_GATE_RADIX=1 + -e FLASHINFER_WORKSPACE_BASE=/cache -e FLASHINFER_CUDA_ARCH_LIST=12.0f + -e SGLANG_CACHE_DIR=/cache + -e XDG_CACHE_HOME=/cache -e TRITON_CACHE_DIR=/cache/triton + -e TORCH_EXTENSIONS_DIR=/cache/torch_extensions -e CUDA_CACHE_PATH=/cache/cuda + -e TMPDIR=/cache/tmp + -e MAX_JOBS=4 --entrypoint python3 "$PD_IMAGE" -m sglang.launch_server + --model-path "$MODEL_PATH" --served-model-name kimi-k3 --trust-remote-code + --tp-size "$tp" --pp-size "$pp" --ep-size "$ep" + --nnodes 4 --node-rank "$rank" --dist-init-addr "${head}:${DIST_PORT}" + --moe-runner-backend flashinfer_mxfp4 --moe-a2a-backend none + --kv-cache-dtype "$KV_DTYPE" --speculative-draft-kv-cache-dtype "$KV_DTYPE" + --chunked-prefill-size "$CHUNK" --page-size "$PAGE_SIZE" + --mem-fraction-static "$mem" --mamba-full-memory-ratio "$ratio" + --mamba-radix-cache-strategy extra_buffer_lazy --disable-radix-cache + --max-running-requests "$MAX_RUNNING" --cuda-graph-max-bs-decode "$GRAPH_BS" + --dist-timeout 3600 --disaggregation-transfer-backend mooncake + --disaggregation-mode "$MODE" --disaggregation-bootstrap-port "$BOOTSTRAP_PORT" + --disaggregation-ib-device "$IB_DEVICES" + --speculative-algorithm DFLASH --speculative-draft-model-path "$DRAFT_MODEL_PATH" + --speculative-num-draft-tokens "$DRAFT_TOKENS" + --enable-metrics --host 0.0.0.0 --port "$PORT") +} + +emit() { + local key="$1" host="$2" + shift 2 + printf '%q ' "$@" >"${RESULT_ROOT}/commands/${key}.cmd.txt" + printf '\n' >>"${RESULT_ROOT}/commands/${key}.cmd.txt" + printf '%s %s: ' "$key" "$host" + cat "${RESULT_ROOT}/commands/${key}.cmd.txt" + [[ "$DRY_RUN" == 1 ]] || root "$host" "$@" +} + +preflight() { + local host gpu image reference="" hash + for host in "${ALL_NODES[@]}"; do + image="$(root "$host" docker image inspect --format '{{.Id}}' "$PD_IMAGE")" + if [[ -z "$reference" ]]; then reference="$image"; fi + [[ "$image" == "$reference" ]] || { echo "image mismatch: $host" >&2; return 1; } + remote "$host" test -r "${MODEL_PATH}/config.json" + remote "$host" test -r "${DRAFT_MODEL_PATH}/model.safetensors" + hash="$(remote "$host" sha256sum "${DRAFT_MODEL_PATH}/config.json")" + [[ "${hash%% *}" == 92e2928e57f417921cd1c031a18840834c55ed13ed0d722acfa8f41b01080717 ]] + remote "$host" test -e /dev/infiniband/uverbs0 + root "$host" mkdir -p "${JIT_CACHE}/tmp" + gpu="$(remote "$host" nvidia-smi --query-compute-apps=pid --format=csv,noheader)" + [[ -z "$gpu" ]] || { echo "GPU occupied: $host $gpu" >&2; return 1; } + remote "$host" python3 -c 'import socket,sys +for p in map(int,sys.argv[1:]): + with socket.socket() as s: + s.setsockopt(socket.SOL_SOCKET,socket.SO_REUSEADDR,1) + try: + s.bind(("0.0.0.0",p)); s.listen(1) + except OSError as exc: + raise SystemExit(f"TCP port {p} unavailable: {exc}")' "$PORT" "$DIST_PORT" "$BOOTSTRAP_PORT" "$ROUTER_PORT" + printf '%s %s\n' "$host" "$image" + done >"${RESULT_ROOT}/metadata/preflight.txt" + hash="$(sha256sum "$GSM8K")" + [[ "${hash%% *}" == "$GSM8K_SHA256" ]] + [[ -r "$BENCH_CLIENT" ]] +} + +logs() { + local group rank host + for group in p d; do + for rank in 0 1 2 3; do + build_command "$group" "$rank"; host="$HOST" + root "$host" docker logs "$(name "$group" "$rank")" \ + >"${RESULT_ROOT}/logs/${group}_${rank}.log" 2>&1 || true + done + done + root "${P_NODES[0]}" docker logs "$(name router 0)" \ + >"${RESULT_ROOT}/logs/router.log" 2>&1 || true +} + +wait_health() { + local host="$1" port="$2" deadline=$((SECONDS + HEALTH_WAIT_S)) + local checked=$SECONDS group rank running + [[ "$DRY_RUN" == 1 ]] && return 0 + while (( SECONDS < deadline )); do + if curl -fsS --max-time 3 "http://${host}:${port}/health" >/dev/null 2>&1; then return 0; fi + if (( SECONDS - checked >= 30 )); then + if [[ "$port" == "$ROUTER_PORT" ]]; then + running="$(root "$host" docker inspect --format '{{.State.Running}}' "$(name router 0)")" + [[ "$running" == true ]] || return 1 + else + if [[ "$host" == "${P_NODES[0]}" ]]; then group=p; else group=d; fi + for rank in 0 1 2 3; do + build_command "$group" "$rank" + running="$(root "$HOST" docker inspect --format '{{.State.Running}}' "$(name "$group" "$rank")")" + [[ "$running" == true ]] || { echo "server exited: $HOST" >&2; return 1; } + done + fi + checked=$SECONDS + fi + sleep 5 + done + echo "health timeout: $host:$port" >&2 + return 1 +} + +start() { + local group rank + [[ "$DRY_RUN" == 1 ]] || preflight + for group in p d; do + for rank in 1 2 3 0; do + build_command "$group" "$rank" + emit "${group}_${rank}" "$HOST" "${CMD[@]}" + done + if [[ "$group" == p ]]; then wait_health "${P_NODES[0]}" "$PORT" + else wait_health "${D_NODES[0]}" "$PORT"; fi + done + emit router "${P_NODES[0]}" docker run -d --name "$(name router 0)" --network host \ + --entrypoint python3 "$PD_IMAGE" -m sglang_router.launch_router \ + --pd-disaggregation --mini-lb \ + --prefill "http://${P_NODES[0]}:${PORT}" "$BOOTSTRAP_PORT" \ + --decode "http://${D_NODES[0]}:${PORT}" --host 0.0.0.0 --port "$ROUTER_PORT" + wait_health "${P_NODES[0]}" "$ROUTER_PORT" +} + +stop() { + local group rank + root "${P_NODES[0]}" docker rm -f "$(name router 0)" >/dev/null 2>&1 || true + for group in d p; do + for rank in 0 1 2 3; do + build_command "$group" "$rank" + root "$HOST" docker rm -f "$(name "$group" "$rank")" >/dev/null 2>&1 || true + done + done +} + +smoke() { + curl -fsS --max-time 600 -H 'Content-Type: application/json' \ + -d '{"text":"Question: Janet has 3 apples and buys 2 more. How many apples does she have?\nAnswer:","sampling_params":{"temperature":0,"max_new_tokens":128}}' \ + "http://${P_NODES[0]}:${ROUTER_PORT}/generate" >"${RESULT_ROOT}/bench/smoke.json" + python3 - "${RESULT_ROOT}/bench/smoke.json" <<'PY' +import json, sys +result = json.load(open(sys.argv[1])) +assert result.get("text"), result +assert result.get("meta_info", {}).get("completion_tokens", 0) > 0, result +print(json.dumps(result, ensure_ascii=False)) +PY + logs +} + +bench() { + local concurrency + for concurrency in 1 8; do + emit "gsm8k_c${concurrency}" "${P_NODES[0]}" docker run --rm --network host \ + -v "${GSM8K}:${GSM8K}:ro" -v "${BENCH_CLIENT}:/bench.py:ro" \ + -v "${RESULT_ROOT}:/results" --entrypoint python3 "$PD_IMAGE" /bench.py \ + --data-path "$GSM8K" --num-questions 64 --num-shots 5 \ + --max-new-tokens 512 --temperature 0 --top-p 1 --parallel "$concurrency" \ + --speculative-algorithm DFLASH \ + --host "${P_NODES[0]}" --port "$ROUTER_PORT" --backend srt \ + --output-file "/results/bench/gsm8k_c${concurrency}.jsonl" \ + --summary-file "/results/bench/gsm8k_c${concurrency}_summary.json" \ + >"${RESULT_ROOT}/bench/gsm8k_c${concurrency}.log" 2>&1 + done + [[ "$DRY_RUN" == 1 ]] || logs +} + +on_error() { + local rc=$? + trap - ERR + logs || true + echo "FAILED rc=$rc; logs: $RESULT_ROOT (containers retained for diagnosis)" >&2 + exit "$rc" +} + +if [[ "$DRY_RUN" == 1 ]]; then + case "$ACTION" in start|all|bench) ;; *) echo 'DRY_RUN only supports start/all/bench' >&2; exit 2;; esac +else + ip -o -4 addr show | grep -q " ${P_NODES[0]}/" || { + echo 'Run this entry on Head 601 only' >&2; exit 2; + } + trap on_error ERR +fi + +case "$ACTION" in + preflight) preflight ;; + start) start ;; + stop) stop ;; + smoke) smoke ;; + bench) bench ;; + logs) logs ;; + status) + for host in "${ALL_NODES[@]}"; do + printf '%s\n' "$host" + root "$host" docker ps --filter "name=${EXPERIMENT}" --format '{{.Names}} {{.Status}}' + done ;; + all) + start + if [[ "$DRY_RUN" != 1 ]]; then smoke; fi + bench + if [[ "$DRY_RUN" != 1 ]]; then logs; stop; fi ;; + *) echo 'Usage: deploy_pd_dflash.sh {preflight|start|stop|status|logs|smoke|bench|all}' >&2; exit 2 ;; +esac diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/flashinfer_4460_b460bc0.patch b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/flashinfer_4460_b460bc0.patch new file mode 100644 index 0000000..09898ae --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/flashinfer_4460_b460bc0.patch @@ -0,0 +1,739 @@ +From b460bc00cb373541102d2155aec35bd626e522ce Mon Sep 17 00:00:00 2001 +From: Xuanyu Zhang +Date: Fri, 21 Aug 2026 20:39:35 +0100 +Subject: [PATCH] feat(moe): add SiTU-GLU activation to the CUTLASS fused-MoE + backend (#4460) +MIME-Version: 1.0 +Content-Type: text/plain; charset=UTF-8 +Content-Transfer-Encoding: 8bit + + + +## 📌 Description + +Adds SiTU-GLU activation support to the CUTLASS fused-MoE backend, +covering all SM variants (SM89/90/100/103/120) via the shared activation +kernel code. + +- Adds `ActivationType::Situ = 10` enum value (consistent with the +TRT-LLM Gen path in #4180) +- Implements `SituAdaptor` with `beta` (default 4.0) and `linear_beta` +(default 25.0) per the SiTU paper (Kimi-K3) +- Uses `2·sigmoid(2z)−1` for tanh (matching the CuTe-DSL path in #4009) +— avoids `tanh.approx.f32` error amplification at `linear_beta=25` + - Supports per-expert `situ_beta` / `situ_linear_beta` tensors +- Refactors per-expert activation param dispatch into +`setPerExpertActivationParams()` / `hasPerExpertActivationParams()` +helpers (reduces duplication across `doGatedActivationKernel` and +`doActivationKernel`) + - Tests both default and per-expert parameters in BF16 and FP8 + +## 🚀 Pull Request Checklist + +Thank you for contributing to FlashInfer! Before we review your pull +request, please make sure the following items are complete. + +### ✅ Pre-commit Checks + +- [x] I have installed `pre-commit` by running `pip install pre-commit` +(or used your preferred method). +- [x] I have installed the hooks with `pre-commit install`. +- [x] I have run the hooks manually with `pre-commit run --all-files` +and fixed any reported issues. + +> If you are unsure about how to set up `pre-commit`, see [the +pre-commit documentation](https://pre-commit.com/). + +## 🧪 Tests + +- [x] Tests have been added or updated as needed. +- [x] `pytest tests/moe/test_trtllm_cutlass_fused_moe.py` — SiTU cases +in both `test_moe` (BF16) and `test_moe_fp8` + +## Reviewer Notes + + + + + +## Summary by CodeRabbit + +- **New Features** +- Added SiTU-GLU activation support for fused Mixture-of-Experts +operations. + - Added optional global or per-expert SiTU scaling parameters. +- Added support across standard, low-latency, and FP8 MoE execution +paths. +- Added default SiTU scaling values when custom parameters are not +provided. + - Added validation for per-expert scaling inputs. + +- **Tests** +- Added coverage for default and per-expert SiTU scales, including FP8 +execution. + + +--------- + +Co-authored-by: Mickael Seznec +Co-authored-by: Claude +--- + benchmarks/routines/moe.py | 15 +-- + .../cutlass_fused_moe_kernels.cuh | 92 ++++++++++-------- + .../flashinfer_cutlass_fused_moe_binding.cu | 93 ++++++++----------- + .../kernels/cutlass_kernels/include/common.h | 1 + + .../include/moe_gemm_kernels.h | 2 +- + .../cutlass_kernels/include/moe_kernels.h | 10 +- + flashinfer/fused_moe/core.py | 19 ++++ + flashinfer/tllm_enums.py | 5 + + tests/moe/test_trtllm_cutlass_fused_moe.py | 69 ++++++++++++-- + 9 files changed, 191 insertions(+), 115 deletions(-) + +diff --git a/benchmarks/routines/moe.py b/benchmarks/routines/moe.py +index e0dae2a5f0..ba55597280 100644 +--- a/benchmarks/routines/moe.py ++++ b/benchmarks/routines/moe.py +@@ -21,7 +21,7 @@ + cutlass_fused_moe, + fused_topk_deepseek, + ) +-from flashinfer.tllm_enums import RoutingMethodType ++from flashinfer.tllm_enums import RoutingMethodType, is_gated_activation + from flashinfer import fp4_quantize, mxfp8_quantize + from flashinfer.testing.utils import ( + bench_gpu_time, +@@ -76,19 +76,6 @@ def _activation_kwarg(fn, activation_type: ActivationType) -> dict: + return {} + + +-def is_gated_activation(activation_type: ActivationType) -> bool: +- """Whether the activation splits FC1 output into gate/up halves (FC1 weight has 2*intermediate +- rows). SwigluStep's clamp limit defaults to 7.0 (the Step-3 model value) in the kernel, so no +- swiglu_limit tensor needs to be passed. +- """ +- return activation_type in ( +- ActivationType.Swiglu, +- ActivationType.Geglu, +- ActivationType.SwigluBias, +- ActivationType.SwigluStep, +- ) +- +- + def run_moe_test(args): + """ + Run a MOE test. +diff --git a/csrc/fused_moe/cutlass_backend/cutlass_fused_moe_kernels.cuh b/csrc/fused_moe/cutlass_backend/cutlass_fused_moe_kernels.cuh +index 0b2ff69359..8fa79d09c0 100644 +--- a/csrc/fused_moe/cutlass_backend/cutlass_fused_moe_kernels.cuh ++++ b/csrc/fused_moe/cutlass_backend/cutlass_fused_moe_kernels.cuh +@@ -2245,6 +2245,47 @@ struct SwigluStepAdaptor { + } + }; + ++// SiTU-GLU (Kimi-K3 / mistral ffn_activations.situ_glu). A gated activation that transforms both ++// branches, evaluated in fp32 (bf16 rounding is visible at the tanh saturation points): ++// out = (beta * tanh(gate / beta) * sigmoid(gate)) * (linear_beta * tanh(up / linear_beta)) ++// Note the sigmoid reads the *uncapped* gate. ++struct SituAdaptor { ++ constexpr static bool IS_GLU = true; ++ float beta = 4.0f; ++ float linear_beta = 25.0f; ++ ++ template ++ __device__ T operator()(T const& gate, T const& linear) const { ++ cutlass::epilogue::thread::Sigmoid sigmoid{}; ++ // tanh(z) == 2*sigmoid(2z) - 1. CUTLASS's Sigmoid uses ::expf, whereas its Tanh lowers to ++ // tanh.approx.f32 whose 2^-11 absolute error linear_beta=25 would amplify to ~1e-2. ++ // The `+ (-1.0f)` is because cutlass::Array has no operator-(Array, scalar). ++ auto tanh = [&](T const& z) { return sigmoid(z * 2.0f) * 2.0f + (-1.0f); }; ++ return (tanh(gate * (1.0f / beta)) * sigmoid(gate) * beta) * ++ (tanh(linear * (1.0f / linear_beta)) * linear_beta); ++ } ++}; ++ ++__device__ inline bool hasPerExpertActivationParams(ActivationParams const& params) { ++ return params.swiglu_alpha || params.swiglu_beta || params.swiglu_limit || params.situ_beta || ++ params.situ_linear_beta; ++} ++ ++// Only assigns what the caller actually supplied, so each adaptor keeps its compile-time default ++// (e.g. SwigluStepAdaptor::limit == 7.0, SituAdaptor::beta == 4.0). ++template ++__device__ void setPerExpertActivationParams(ActFn& fn, ActivationParams const& params, ++ int64_t expert) { ++ if constexpr (std::is_same_v) { ++ if (params.situ_beta) fn.beta = params.situ_beta[expert]; ++ if (params.situ_linear_beta) fn.linear_beta = params.situ_linear_beta[expert]; ++ } else { ++ if (params.swiglu_alpha) fn.alpha = params.swiglu_alpha[expert]; ++ if (params.swiglu_beta) fn.beta = params.swiglu_beta[expert]; ++ if (params.swiglu_limit) fn.limit = params.swiglu_limit[expert]; ++ } ++} ++ + // ============================== Gated Activation ================================= + constexpr static int MAX_ACTIVATION_THREADS_PER_BLOCK = 256; + +@@ -2276,26 +2317,12 @@ __global__ void doGatedActivationKernel(ActivationOutputType* output, + int64_t const num_elems_in_col = inter_size / ACTIVATION_ELEM_PER_THREAD; + int64_t const inter_size_vec = inter_size / ACTIVATION_ELEM_PER_THREAD; + +- float gate_alpha = 1.0f; +- float gate_bias = 0.0f; +- float gate_limit = std::numeric_limits::infinity(); +- if (activation_type.swiglu_alpha || activation_type.swiglu_beta || activation_type.swiglu_limit) { +- int expert = findTotalEltsLessThanTarget(expert_first_token_offset, num_experts_per_node, +- (int64_t)token + 1) - +- 1; +- gate_alpha = activation_type.swiglu_alpha ? activation_type.swiglu_alpha[expert] : 1.0f; +- gate_bias = activation_type.swiglu_beta ? activation_type.swiglu_beta[expert] : 0.0f; +- gate_limit = activation_type.swiglu_limit ? activation_type.swiglu_limit[expert] +- : std::numeric_limits::infinity(); +- } +- + ActFn fn{}; +- fn.alpha = gate_alpha; +- fn.beta = gate_bias; +- // Keep the activation's compile-time default limit (e.g. 7.0 for SwigluStep) unless the caller +- // supplied a per-expert swiglu_limit tensor. +- if (activation_type.swiglu_limit) { +- fn.limit = gate_limit; ++ if (hasPerExpertActivationParams(activation_type)) { ++ int64_t const expert = findTotalEltsLessThanTarget(expert_first_token_offset, ++ num_experts_per_node, (int64_t)token + 1) - ++ 1; ++ setPerExpertActivationParams(fn, activation_type, expert); + } + for (int64_t elem_index = start_offset; elem_index < num_elems_in_col; elem_index += stride) { + auto linear_value = arrayConvert(gemm_result_vec[elem_index]); +@@ -2328,6 +2355,8 @@ void doGatedActivation(ActivationOutputType* output, GemmOutputType const* gemm_ + ? &doGatedActivationKernel + : activation_type == ActivationType::SwigluStep + ? &doGatedActivationKernel ++ : activation_type == ActivationType::Situ ++ ? &doGatedActivationKernel + : nullptr; + TLLM_CHECK_WITH_INFO(fn != nullptr, "Invalid activation type"); + fn<<>>(output, gemm_result, expert_first_token_offset, inter_size, +@@ -2391,22 +2420,13 @@ __global__ __launch_bounds__(MAX_ACTIVATION_THREADS_PER_BLOCK) void doActivation + size_t output_offset = token * inter_size; + + int64_t expert = 0; +- float gate_alpha = 1.0f; +- float gate_beta = 0.0f; +- float gate_limit = std::numeric_limits::infinity(); + if (bias_ptr || IsNVFP4 || IsMXFP8 || use_per_expert_act_scale || +- activation_params.swiglu_alpha || activation_params.swiglu_beta || +- activation_params.swiglu_limit) { ++ hasPerExpertActivationParams(activation_params)) { + expert = permuted_token_selected_experts + ? permuted_token_selected_experts[token] + : findTotalEltsLessThanTarget(expert_first_token_offset, num_experts_per_node, + token + 1) - + 1; +- +- gate_alpha = activation_params.swiglu_alpha ? activation_params.swiglu_alpha[expert] : 1.0f; +- gate_beta = activation_params.swiglu_beta ? activation_params.swiglu_beta[expert] : 0.0f; +- gate_limit = activation_params.swiglu_limit ? activation_params.swiglu_limit[expert] +- : std::numeric_limits::infinity(); + } + + size_t act_scale_idx = use_per_expert_act_scale ? expert : 0; +@@ -2444,13 +2464,7 @@ __global__ __launch_bounds__(MAX_ACTIVATION_THREADS_PER_BLOCK) void doActivation + int64_t const gated_off_vec = gated_off / ACTIVATION_ELEM_PER_THREAD; + + ActFn fn{}; +- fn.alpha = gate_alpha; +- fn.beta = gate_beta; +- // Keep the activation's compile-time default limit (e.g. 7.0 for SwigluStep) unless the caller +- // supplied a per-expert swiglu_limit tensor. +- if (activation_params.swiglu_limit) { +- fn.limit = gate_limit; +- } ++ setPerExpertActivationParams(fn, activation_params, expert); + auto compute_activation = [&](int64_t elem_index) { + GemmResultElem fc1_gemm_value; + cutlass::arch::global_load( +@@ -2670,7 +2684,11 @@ void doActivation(T* output, GemmOutputType const* gemm_result, float const* fp8 + IdentityAdaptor, + decltype(block_scaling_type)::value, + decltype(disableFP4QuantFastMathTag)::value, +- decltype(nvfp4_4over6_config_tag)> // Identity ++ decltype(nvfp4_4over6_config_tag)>, // Identity ++ &doActivationKernel // Situ + }; + return fn_list[static_cast(activation_type.activation_type)]; + }; +diff --git a/csrc/fused_moe/cutlass_backend/flashinfer_cutlass_fused_moe_binding.cu b/csrc/fused_moe/cutlass_backend/flashinfer_cutlass_fused_moe_binding.cu +index 79a7aa7757..46237da826 100644 +--- a/csrc/fused_moe/cutlass_backend/flashinfer_cutlass_fused_moe_binding.cu ++++ b/csrc/fused_moe/cutlass_backend/flashinfer_cutlass_fused_moe_binding.cu +@@ -81,6 +81,17 @@ class DtypeUtils { + DtypeUtils() = default; + }; + ++// Validates one of the optional per-expert activation scale tensors (swiglu_alpha, situ_beta, ...) ++// and returns its data pointer, or nullptr when the caller did not supply it. ++inline float const* checkedPerExpertScale(Optional const& scale, ++ int num_experts_on_rank, char const* name) { ++ if (!scale.has_value()) return nullptr; ++ CHECK_INPUT_AND_TYPE(scale.value(), dl_float32); ++ TVM_FFI_ICHECK_EQ(scale.value().size(0), num_experts_on_rank) ++ << name << " must have num_experts_on_rank elements."; ++ return static_cast(scale.value().data_ptr()); ++} ++ + class FusedMoeRunner : public tvm::ffi::ModuleObj { + public: + template < +@@ -302,6 +313,7 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + Optional fc2_expert_biases, Optional> quant_scales, + Optional input_sf, Optional swiglu_alpha, + Optional swiglu_beta, Optional swiglu_limit, ++ Optional situ_beta, Optional situ_linear_beta, + bool swizzled_input_sf, int64_t tp_size, int64_t tp_rank, int64_t ep_size, + int64_t ep_rank, int64_t cluster_size, int64_t cluster_rank, bool enable_alltoall, + bool min_latency_mode, Optional> profile_ids, bool enable_pdl, +@@ -397,21 +409,6 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + int const num_experts_on_rank = fc2_expert_weights.size(0); + auto const num_experts_total = static_cast(num_experts_on_rank * ep_size); + auto parallelism_config = kernels::MOEParallelismConfig(tp_size, tp_rank, ep_size, ep_rank); +- if (swiglu_alpha.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_alpha.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_alpha.value().size(0), num_experts_on_rank) +- << "swiglu_alpha must have num_experts_on_rank elements."; +- } +- if (swiglu_beta.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_beta.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_beta.value().size(0), num_experts_on_rank) +- << "swiglu_beta must have num_experts_on_rank elements."; +- } +- if (swiglu_limit.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_limit.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_limit.value().size(0), num_experts_on_rank) +- << "swiglu_limit must have num_experts_on_rank elements."; +- } + // Swiglu + swiglu_alpha/beta/limit selects the SwigluBias kernel; other gated activations + // (e.g. SwigluStep) keep their own kernel. + if (base_activation_type == ActivationType::Swiglu && +@@ -420,12 +417,11 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + } + auto activation_params = ActivationParams( + base_activation_type, +- reinterpret_cast(swiglu_alpha.has_value() ? swiglu_alpha.value().data_ptr() +- : nullptr), +- reinterpret_cast(swiglu_beta.has_value() ? swiglu_beta.value().data_ptr() +- : nullptr), +- reinterpret_cast(swiglu_limit.has_value() ? swiglu_limit.value().data_ptr() +- : nullptr)); ++ checkedPerExpertScale(swiglu_alpha, num_experts_on_rank, "swiglu_alpha"), ++ checkedPerExpertScale(swiglu_beta, num_experts_on_rank, "swiglu_beta"), ++ checkedPerExpertScale(swiglu_limit, num_experts_on_rank, "swiglu_limit"), ++ checkedPerExpertScale(situ_beta, num_experts_on_rank, "situ_beta"), ++ checkedPerExpertScale(situ_linear_beta, num_experts_on_rank, "situ_linear_beta")); + + setRunnerProfiles(profile_ids); + +@@ -486,7 +482,8 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + Optional fc2_expert_biases, + Optional> quant_scales, Optional input_sf, + Optional swiglu_alpha, Optional swiglu_beta, +- Optional swiglu_limit, bool swizzled_input_sf, ++ Optional swiglu_limit, Optional situ_beta, ++ Optional situ_linear_beta, bool swizzled_input_sf, + TensorView num_active_experts_per_node, TensorView experts_to_token_score, + TensorView active_expert_global_ids, int64_t tp_size, int64_t tp_rank, + int64_t ep_size, int64_t ep_rank, int64_t cluster_size, +@@ -567,21 +564,6 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + int const num_experts_on_rank = fc2_expert_weights.size(0); + auto const num_experts_total = static_cast(num_experts_on_rank * ep_size); + auto parallelism_config = kernels::MOEParallelismConfig(tp_size, tp_rank, ep_size, ep_rank); +- if (swiglu_alpha.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_alpha.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_alpha.value().size(0), num_experts_on_rank) +- << "swiglu_alpha must have num_experts_on_rank elements."; +- } +- if (swiglu_beta.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_beta.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_beta.value().size(0), num_experts_on_rank) +- << "swiglu_beta must have num_experts_on_rank elements."; +- } +- if (swiglu_limit.has_value()) { +- CHECK_INPUT_AND_TYPE(swiglu_limit.value(), dl_float32); +- TVM_FFI_ICHECK_EQ(swiglu_limit.value().size(0), num_experts_on_rank) +- << "swiglu_limit must have num_experts_on_rank elements."; +- } + // Swiglu + swiglu_alpha/beta/limit selects the SwigluBias kernel; other gated activations + // (e.g. SwigluStep) keep their own kernel. + if (base_activation_type == ActivationType::Swiglu && +@@ -590,12 +572,11 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + } + auto activation_params = ActivationParams( + base_activation_type, +- reinterpret_cast(swiglu_alpha.has_value() ? swiglu_alpha.value().data_ptr() +- : nullptr), +- reinterpret_cast(swiglu_beta.has_value() ? swiglu_beta.value().data_ptr() +- : nullptr), +- reinterpret_cast(swiglu_limit.has_value() ? swiglu_limit.value().data_ptr() +- : nullptr)); ++ checkedPerExpertScale(swiglu_alpha, num_experts_on_rank, "swiglu_alpha"), ++ checkedPerExpertScale(swiglu_beta, num_experts_on_rank, "swiglu_beta"), ++ checkedPerExpertScale(swiglu_limit, num_experts_on_rank, "swiglu_limit"), ++ checkedPerExpertScale(situ_beta, num_experts_on_rank, "situ_beta"), ++ checkedPerExpertScale(situ_linear_beta, num_experts_on_rank, "situ_linear_beta")); + + setRunnerProfiles(profile_ids); + +@@ -811,16 +792,17 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + Optional fc2_expert_biases, Optional> quant_scales, + Optional input_sf, Optional swiglu_alpha, + Optional swiglu_beta, Optional swiglu_limit, ++ Optional situ_beta, Optional situ_linear_beta, + bool swizzled_input_sf, int64_t tp_size, int64_t tp_rank, int64_t ep_size, + int64_t ep_rank, int64_t cluster_size, int64_t cluster_rank, bool enable_alltoall, + bool min_latency_mode, Optional> profile_ids, bool enable_pdl, + int64_t base_activation_type, Optional workspace_buffer) { + runMoe(output, input, token_selected_experts, token_final_scales, fc1_expert_weights, + fc1_expert_biases, fc2_expert_weights, fc2_expert_biases, quant_scales, input_sf, +- swiglu_alpha, swiglu_beta, swiglu_limit, swizzled_input_sf, tp_size, tp_rank, +- ep_size, ep_rank, cluster_size, cluster_rank, enable_alltoall, min_latency_mode, +- profile_ids, enable_pdl, static_cast(base_activation_type), +- workspace_buffer); ++ swiglu_alpha, swiglu_beta, swiglu_limit, situ_beta, situ_linear_beta, ++ swizzled_input_sf, tp_size, tp_rank, ep_size, ep_rank, cluster_size, ++ cluster_rank, enable_alltoall, min_latency_mode, profile_ids, enable_pdl, ++ static_cast(base_activation_type), workspace_buffer); + }); + } else if (name == "run_moe_min_latency") { + return Function::FromTyped( +@@ -830,20 +812,21 @@ class FusedMoeRunner : public tvm::ffi::ModuleObj { + Optional fc2_expert_biases, Optional> quant_scales, + Optional input_sf, Optional swiglu_alpha, + Optional swiglu_beta, Optional swiglu_limit, ++ Optional situ_beta, Optional situ_linear_beta, + bool swizzled_input_sf, TensorView num_active_experts_per_node, + TensorView experts_to_token_score, TensorView active_expert_global_ids, + int64_t tp_size, int64_t tp_rank, int64_t ep_size, int64_t ep_rank, + int64_t cluster_size, int64_t cluster_rank, bool enable_alltoall, + bool min_latency_mode, Optional> profile_ids, bool enable_pdl, + int64_t base_activation_type, Optional workspace_buffer) { +- runMoeMinLantency(output, input, token_selected_experts, token_final_scales, +- fc1_expert_weights, fc1_expert_biases, fc2_expert_weights, +- fc2_expert_biases, quant_scales, input_sf, swiglu_alpha, swiglu_beta, +- swiglu_limit, swizzled_input_sf, num_active_experts_per_node, +- experts_to_token_score, active_expert_global_ids, tp_size, tp_rank, +- ep_size, ep_rank, cluster_size, cluster_rank, enable_alltoall, +- min_latency_mode, profile_ids, enable_pdl, +- static_cast(base_activation_type), workspace_buffer); ++ runMoeMinLantency( ++ output, input, token_selected_experts, token_final_scales, fc1_expert_weights, ++ fc1_expert_biases, fc2_expert_weights, fc2_expert_biases, quant_scales, input_sf, ++ swiglu_alpha, swiglu_beta, swiglu_limit, situ_beta, situ_linear_beta, ++ swizzled_input_sf, num_active_experts_per_node, experts_to_token_score, ++ active_expert_global_ids, tp_size, tp_rank, ep_size, ep_rank, cluster_size, ++ cluster_rank, enable_alltoall, min_latency_mode, profile_ids, enable_pdl, ++ static_cast(base_activation_type), workspace_buffer); + }); + } else if (name == "get_workspace_size") { + return Function::FromTyped([this](int64_t num_rows, int64_t hidden_size, int64_t inter_size, +diff --git a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/common.h b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/common.h +index ce1c0df4e0..11c01a6860 100644 +--- a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/common.h ++++ b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/common.h +@@ -30,6 +30,7 @@ enum class ActivationType { + SwigluStep, + GegluTanh, + Identity, ++ Situ, + InvalidType + }; + +diff --git a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_gemm_kernels.h b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_gemm_kernels.h +index 3761e32cf1..2be7043187 100644 +--- a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_gemm_kernels.h ++++ b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_gemm_kernels.h +@@ -244,7 +244,7 @@ constexpr bool isGatedActivation(ActivationType activation_type) { + return activation_type == ActivationType::Swiglu || activation_type == ActivationType::Geglu || + activation_type == ActivationType::SwigluBias || + activation_type == ActivationType::SwigluStep || +- activation_type == ActivationType::GegluTanh; ++ activation_type == ActivationType::GegluTanh || activation_type == ActivationType::Situ; + } + + enum class Sm90Wfp4Afp8ScaleMode : uint8_t { +diff --git a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_kernels.h b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_kernels.h +index 01daea011b..dfd19e72c5 100644 +--- a/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_kernels.h ++++ b/csrc/nv_internal/tensorrt_llm/kernels/cutlass_kernels/include/moe_kernels.h +@@ -118,6 +118,9 @@ struct ActivationParams { + float const* swiglu_alpha = nullptr; + float const* swiglu_beta = nullptr; + float const* swiglu_limit = nullptr; ++ // SiTU-GLU per-expert tanh scales; nullptr uses the SituAdaptor compile-time defaults. ++ float const* situ_beta = nullptr; ++ float const* situ_linear_beta = nullptr; + + explicit ActivationParams(ActivationType activation_type) : activation_type(activation_type) { + TLLM_CHECK_WITH_INFO( +@@ -126,11 +129,14 @@ struct ActivationParams { + } + + ActivationParams(ActivationType activation_type, float const* swiglu_alpha, +- float const* swiglu_beta, float const* swiglu_limit) ++ float const* swiglu_beta, float const* swiglu_limit, ++ float const* situ_beta = nullptr, float const* situ_linear_beta = nullptr) + : activation_type(activation_type), + swiglu_alpha(swiglu_alpha), + swiglu_beta(swiglu_beta), +- swiglu_limit(swiglu_limit) {} ++ swiglu_limit(swiglu_limit), ++ situ_beta(situ_beta), ++ situ_linear_beta(situ_linear_beta) {} + + // TODO Port everything properly and get rid of these implicit conversions + operator ActivationType() const { return activation_type; } +diff --git a/flashinfer/fused_moe/core.py b/flashinfer/fused_moe/core.py +index 585ce94048..8a68e9f27b 100644 +--- a/flashinfer/fused_moe/core.py ++++ b/flashinfer/fused_moe/core.py +@@ -884,6 +884,8 @@ def cutlass_fused_moe( + swiglu_alpha: Optional[torch.Tensor] = None, + swiglu_beta: Optional[torch.Tensor] = None, + swiglu_limit: Optional[torch.Tensor] = None, ++ situ_beta: Optional[torch.Tensor] = None, ++ situ_linear_beta: Optional[torch.Tensor] = None, + swizzled_input_sf: bool = True, + tp_size: int = 1, + tp_rank: int = 0, +@@ -1015,6 +1017,8 @@ def cutlass_fused_moe( + swiglu_alpha, + swiglu_beta, + swiglu_limit, ++ situ_beta, ++ situ_linear_beta, + swizzled_input_sf, + *min_latency_output, + tp_size, +@@ -1058,6 +1062,8 @@ def _fake_cutlass_fused_moe( + swiglu_alpha: Optional[torch.Tensor] = None, + swiglu_beta: Optional[torch.Tensor] = None, + swiglu_limit: Optional[torch.Tensor] = None, ++ situ_beta: Optional[torch.Tensor] = None, ++ situ_linear_beta: Optional[torch.Tensor] = None, + swizzled_input_sf: bool = True, + tp_size: int = 1, + tp_rank: int = 0, +@@ -1206,6 +1212,9 @@ def cutlass_fused_moe( + use_fused_finalize: bool = True, + profile_ids: Optional[List[int]] = None, + workspace_buffer: Optional[torch.Tensor] = None, ++ *, ++ situ_beta: Optional[torch.Tensor] = None, ++ situ_linear_beta: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Compute a Mixture of Experts (MoE) layer using CUTLASS backend. + +@@ -1278,6 +1287,14 @@ def cutlass_fused_moe( + swiglu_limit : Optional[torch.Tensor] + Swiglu limit for swiglu activation. + ++ situ_beta : Optional[torch.Tensor] ++ Per-expert ``beta`` tanh scale for the ``Situ`` activation (float32, ++ ``[num_experts_on_rank]``). ``None`` uses ``DEFAULT_SITU_BETA``. ++ ++ situ_linear_beta : Optional[torch.Tensor] ++ Per-expert ``linear_beta`` tanh scale for the ``Situ`` activation (float32, ++ ``[num_experts_on_rank]``). ``None`` uses ``DEFAULT_SITU_LINEAR_BETA``. ++ + tp_size : int = 1 + Tensor parallelism size. Defaults to 1. + +@@ -1447,6 +1464,8 @@ def cutlass_fused_moe( + swiglu_alpha, + swiglu_beta, + swiglu_limit, ++ situ_beta, ++ situ_linear_beta, + swizzled_input_sf, + tp_size, + tp_rank, +diff --git a/flashinfer/tllm_enums.py b/flashinfer/tllm_enums.py +index c0fad19ec4..0d6d783331 100644 +--- a/flashinfer/tllm_enums.py ++++ b/flashinfer/tllm_enums.py +@@ -103,6 +103,11 @@ def is_gated(self) -> bool: + DEFAULT_SWIGLU_BETA = 0.0 + DEFAULT_SWIGLU_LIMIT = torch.finfo(torch.float32).max + ++# SiTU-GLU tanh scales. Must match the SituAdaptor defaults in ++# csrc/fused_moe/cutlass_backend/cutlass_fused_moe_kernels.cuh. ++DEFAULT_SITU_BETA = 4.0 ++DEFAULT_SITU_LINEAR_BETA = 25.0 ++ + + def normalize_activation_type( + activation_type: Union[int, ActivationType], +diff --git a/tests/moe/test_trtllm_cutlass_fused_moe.py b/tests/moe/test_trtllm_cutlass_fused_moe.py +index 3f58d5dbf7..e874c44810 100644 +--- a/tests/moe/test_trtllm_cutlass_fused_moe.py ++++ b/tests/moe/test_trtllm_cutlass_fused_moe.py +@@ -20,6 +20,7 @@ + + import pytest + from flashinfer.fused_moe.core import ActivationType ++from flashinfer.tllm_enums import DEFAULT_SITU_BETA, DEFAULT_SITU_LINEAR_BETA + import torch + from torch.nn import functional as F + +@@ -53,6 +54,17 @@ + set_nvfp4_4over6_env = moe_utils.set_nvfp4_4over6_env + + ++def make_situ_scales(num_experts): ++ """Per-expert SiTU-GLU tanh scales, deliberately different from DEFAULT_SITU_BETA / ++ DEFAULT_SITU_LINEAR_BETA so a kernel silently falling back to those would fail.""" ++ return { ++ "situ_beta": torch.full((num_experts,), 5.0, dtype=torch.float32).cuda(), ++ "situ_linear_beta": torch.full( ++ (num_experts,), 18.0, dtype=torch.float32 ++ ).cuda(), ++ } ++ ++ + def dynamic_per_tensor_fp8_quant(x: torch.tensor) -> tuple[torch.tensor, torch.tensor]: + fp8_traits_max = FLOAT8_E4M3_MAX + fp8_traits_min = -FLOAT8_E4M3_MAX +@@ -323,6 +335,8 @@ def compute_with_experts( + beta=None, + limit=None, + activation_type=ActivationType.Swiglu, ++ situ_beta=DEFAULT_SITU_BETA, ++ situ_linear_beta=DEFAULT_SITU_LINEAR_BETA, + ): + results = torch.zeros_like(x) + for expert_id in range(num_experts): +@@ -359,6 +373,22 @@ def compute_with_experts( + x2 = x2.clamp_(min=-limit, max=limit) + beta + + inter = x1_scaled * x2 ++ elif activation_type == ActivationType.Situ: ++ # SiTU-GLU, computed in fp32; see SituAdaptor in cutlass_fused_moe_kernels.cuh. ++ # situ_beta / situ_linear_beta are per-expert when given as a tensor/list. ++ sb = float( ++ situ_beta[expert_id] if hasattr(situ_beta, "__getitem__") else situ_beta ++ ) ++ slb = float( ++ situ_linear_beta[expert_id] ++ if hasattr(situ_linear_beta, "__getitem__") ++ else situ_linear_beta ++ ) ++ gate = (expert_inputs @ w1_expert.t()).float() ++ up = (expert_inputs @ w3_expert.t()).float() ++ out_glu = sb * torch.tanh(gate / sb) * torch.sigmoid(gate) ++ out_linear = slb * torch.tanh(up / slb) ++ inter = (out_glu * out_linear).to(x.dtype) + else: + inter = F.silu(expert_inputs @ w1_expert.t()) * ( + expert_inputs @ w3_expert.t() +@@ -390,12 +420,23 @@ def compute_with_experts( + @pytest.mark.parametrize("top_k", TOP_K_VALUES) + @pytest.mark.parametrize("intermediate_size", INTERMEDIATE_SIZES) + @pytest.mark.parametrize( +- "activation_type", +- [ActivationType.Swiglu, ActivationType.SwigluStep], +- ids=["swiglu", "swiglustep"], ++ "activation_type, situ_per_expert", ++ [ ++ (ActivationType.Swiglu, False), ++ (ActivationType.SwigluStep, False), ++ (ActivationType.Situ, False), ++ (ActivationType.Situ, True), ++ ], ++ ids=["swiglu", "swiglustep", "situ_default", "situ_per_expert"], + ) + def test_moe( +- batch_size, hidden_size, num_experts, top_k, intermediate_size, activation_type ++ batch_size, ++ hidden_size, ++ num_experts, ++ top_k, ++ intermediate_size, ++ activation_type, ++ situ_per_expert, + ): + # Skip invalid configurations + if top_k > num_experts: +@@ -425,6 +466,13 @@ def test_moe( + ) + + routing_weights, selected_experts = compute_routing(router_logits, top_k) ++ ++ # When situ_per_expert is off the kernel must fall back to its compile-time defaults, which is ++ # what compute_with_experts() uses by default. The per-expert scales below are deliberately ++ # non-default so the test fails if the kernel ignores the tensors and falls back anyway. ++ situ_kwargs = make_situ_scales(num_experts) if situ_per_expert else {} ++ ref_situ_kwargs = {k: v.tolist() for k, v in situ_kwargs.items()} ++ + ref_output = compute_with_experts( + num_experts, + x, +@@ -433,6 +481,7 @@ def test_moe( + selected_experts, + routing_weights, + activation_type=activation_type, ++ **ref_situ_kwargs, + ) + flash_output = torch.empty_like(ref_output) + flash_output = fused_moe.cutlass_fused_moe( +@@ -445,6 +494,7 @@ def test_moe( + output=flash_output, + quant_scales=None, + activation_type=activation_type, ++ **situ_kwargs, + ) + + torch.testing.assert_close(ref_output, flash_output[0], rtol=1e-2, atol=1e-2) +@@ -613,8 +663,8 @@ def run_unfused(): + @pytest.mark.parametrize("otype, wtype", [(torch.float16, torch.float8_e4m3fn)]) + @pytest.mark.parametrize( + "activation_type", +- [ActivationType.Swiglu, ActivationType.SwigluStep], +- ids=["swiglu", "swiglustep"], ++ [ActivationType.Swiglu, ActivationType.SwigluStep, ActivationType.Situ], ++ ids=["swiglu", "swiglustep", "situ"], + ) + def test_moe_fp8( + batch_size, +@@ -661,6 +711,11 @@ def test_moe_fp8( + w31_dequantized.data[expert_id].copy_(torch.mul(w31_quant.to(dtype=otype), s31)) + w2_dequantized.data[expert_id].copy_(torch.mul(w2_quant.to(dtype=otype), s2)) + ++ situ_kwargs = ( ++ make_situ_scales(num_experts) if activation_type == ActivationType.Situ else {} ++ ) ++ ref_situ_kwargs = {k: v.tolist() for k, v in situ_kwargs.items()} ++ + routing_weights, selected_experts = compute_routing(router_logits, top_k) + ref_output = compute_with_experts( + num_experts, +@@ -670,6 +725,7 @@ def test_moe_fp8( + selected_experts, + routing_weights, + activation_type=activation_type, ++ **ref_situ_kwargs, + ) + flash_output = torch.empty_like(ref_output) + # For fp8, the hidden_state expects quantized. +@@ -693,6 +749,7 @@ def test_moe_fp8( + quant_scales=quant_scales, + output=flash_output, + activation_type=activation_type, ++ **situ_kwargs, + ) + torch.testing.assert_close(ref_output, flash_output, rtol=1e-1, atol=1e-1) + diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_files.json b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_files.json new file mode 100644 index 0000000..8a1546e --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_files.json @@ -0,0 +1,81 @@ +{ + "base": "6465a6f3d3b6c8b7fee40fba0fdc09cf5e9ca1c5", + "files": [ + { + "path": "python/sglang/srt/arg_groups/speculative_hook.py", + "sha256": "26d656eba10219a5ff6bb4c2675ca15355f6b3d9f33639af9b3a21a8e1f8b2cb" + }, + { + "path": "python/sglang/srt/arg_groups/validation_hook.py", + "sha256": "08151771718df1f17a0df91b005d17fd62055314a642c656d7964ff7c5b11b40" + }, + { + "path": "python/sglang/srt/disaggregation/base/conn.py", + "sha256": "fbd2c0e0d8deee59f58a7a2dccb38e6137aae74a8e47fc40a23bb25d8b7e86e1" + }, + { + "path": "python/sglang/srt/disaggregation/common/kv_entry_layout.py", + "sha256": "97f559279f1f87fb4769ef923e47b697476df85f324d7594a97f3364f9875042" + }, + { + "path": "python/sglang/srt/disaggregation/decode.py", + "sha256": "b00e588cb3d0b6706a586ef4b488454834d4891b46c0679c2d41781b1a4f3e28" + }, + { + "path": "python/sglang/srt/disaggregation/mooncake/conn.py", + "sha256": "08c8bf93e5c0c04e581b585fe355284e2967a623b5d99761d212a67c66500f49" + }, + { + "path": "python/sglang/srt/disaggregation/prefill.py", + "sha256": "d0175b3df8eddd58fbfa41805d0928d614e412cf74055e57aade82578b174bfd" + }, + { + "path": "python/sglang/srt/disaggregation/utils.py", + "sha256": "40a1b51fd9b549b4208f7dc4b72061f99d96fad9096bc9f270d88b6fd556512b" + }, + { + "path": "python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py", + "sha256": "128bc0d46cc1c7b3c5258a423437a0d79b2f02c15ad2fcfc3cbf26b782921c86" + }, + { + "path": "python/sglang/srt/layers/quantization/mxfp4.py", + "sha256": "bf072608ac84646c932c5dcbad3ee5f9cda7f2a73cbee97801be06aa64c2f2c7" + }, + { + "path": "python/sglang/srt/managers/scheduler_pp_mixin.py", + "sha256": "46ae2313973ba3769382da4376c59961e547f8f7e6a204c27e593cc51ef712db" + }, + { + "path": "python/sglang/srt/model_executor/runner/base_runner.py", + "sha256": "95aee38376e2cfc178e9bd3d675a3489268d4aa3c41934d22a27339f16810c7e" + }, + { + "path": "python/sglang/srt/models/kimi_k3.py", + "sha256": "d4d1c13050ce11e63e05c6aa589f7d0a7431cb5994138a83846bfa1f57890ba4" + }, + { + "path": "python/sglang/srt/speculative/dflash_pp.py", + "sha256": "efb9f02236c9c5eabe8e27c14eedffb1cbc297cc20a7abd8aca9d748e0096b3a" + }, + { + "path": "python/sglang/srt/speculative/dflash_worker_v2.py", + "sha256": "73d9dd199605a58e3090c2b8f993b851089fd21d32cf78978061704569f7c80c" + }, + { + "path": "python/sglang/srt/speculative/spec_info.py", + "sha256": "c188dfd0bf2abba18e24fee9a44897c3b564dffcfb228bd079ea900be6b1fc50" + }, + { + "path": "test/registered/disaggregation/test_dflash_pp_context.py", + "sha256": "3f3022afa4acd7f66e14caf6bbbfb88b1f2e9320d087300aa78c9d9e20d3a321" + }, + { + "path": "test/registered/disaggregation/test_flashinfer_kimi_merge.py", + "sha256": "0cdb173b8d073c3ba2ce5e07eee75f203bcb2313156c7e7af3c27bb24b4ea9e2" + }, + { + "path": "test/registered/disaggregation/test_mixed_kv_entry_layout.py", + "sha256": "d0dabf551d27f5b80ffdb5129cec42c4fc628887e6faa73537981872705a1d8e" + } + ] +} diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_integration.patch b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_integration.patch new file mode 100644 index 0000000..83cf672 --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/pp_dflash_integration.patch @@ -0,0 +1,2203 @@ +diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py +--- a/python/sglang/srt/arg_groups/speculative_hook.py ++++ b/python/sglang/srt/arg_groups/speculative_hook.py +@@ -196,9 +196,9 @@ + "Currently DFLASH speculative decoding does not support dp attention." + ) + +- if cfg.pp_size != 1: +- raise ValueError( +- "Currently DFLASH speculative decoding only supports pp_size == 1." ++ if cfg.pp_size != 1 and cfg.disaggregation_mode != "prefill": ++ raise ValueError( ++ "DFLASH with pp_size > 1 is only supported on a PD prefill server." + ) + + if cfg.speculative_draft_model_path is None: +diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py +--- a/python/sglang/srt/arg_groups/validation_hook.py ++++ b/python/sglang/srt/arg_groups/validation_hook.py +@@ -48,12 +48,13 @@ + assert ( + cfg.disable_overlap_schedule + ), "Pipeline parallelism is not compatible with overlap schedule" +- pp_dspark_prefill = ( +- cfg.speculative_algorithm or "" +- ).upper() == "DSPARK" and cfg.disaggregation_mode == "prefill" ++ pp_dspark_prefill = (cfg.speculative_algorithm or "").upper() in ( ++ "DSPARK", ++ "DFLASH", ++ ) and cfg.disaggregation_mode == "prefill" + assert cfg.speculative_algorithm is None or pp_dspark_prefill, ( + "Pipeline parallelism with speculative decoding is only supported " +- "for DSPARK on a PD prefill server" ++ "for DSPARK/DFLASH on a PD prefill server" + ) + assert cfg.min_free_slots_delay is None, ( + "--min-free-slots-delay is not supported with pipeline " +diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py +--- a/python/sglang/srt/disaggregation/base/conn.py ++++ b/python/sglang/srt/disaggregation/base/conn.py +@@ -46,6 +46,7 @@ + kv_data_lens: List[int] + kv_item_lens: List[int] + kv_layer_ids: List[int] ++ kv_entry_layouts: Optional[List[dict]] = None + kv_cache_dtype_str: str + aux_data_ptrs: List[int] + aux_data_lens: List[int] +diff --git a/python/sglang/srt/disaggregation/common/kv_entry_layout.py b/python/sglang/srt/disaggregation/common/kv_entry_layout.py +new file mode 100644 +--- /dev/null ++++ b/python/sglang/srt/disaggregation/common/kv_entry_layout.py +@@ -0,0 +1,157 @@ ++"""Validated byte plans for mixed replicated MLA and head-sharded draft KV.""" ++ ++from __future__ import annotations ++ ++from dataclasses import asdict, dataclass ++from typing import Sequence ++ ++ ++@dataclass(frozen=True) ++class KVEntryLayout: ++ kind: str ++ page_size: int ++ item_len: int ++ buffer_nbytes: int ++ dtype: str ++ dtype_bytes: int ++ total_heads: int = 0 ++ head_dim: int = 0 ++ ++ def __post_init__(self): ++ if self.kind not in ("flat", "nhd"): ++ raise ValueError(f"Unsupported PD KV entry layout: {self.kind}") ++ for name in ("page_size", "item_len", "buffer_nbytes", "dtype_bytes"): ++ value = getattr(self, name) ++ if type(value) is not int or value <= 0: ++ raise ValueError(f"{name} must be a positive integer") ++ if not isinstance(self.dtype, str) or not self.dtype: ++ raise ValueError("KV dtype must be specified") ++ if self.kind == "nhd": ++ for name in ("total_heads", "head_dim"): ++ value = getattr(self, name) ++ if type(value) is not int or value <= 0: ++ raise ValueError(f"{name} must be a positive integer") ++ ++ def to_wire(self): ++ return asdict(self) ++ ++ @classmethod ++ def from_wire(cls, record): ++ if not isinstance(record, dict): ++ raise ValueError("KV entry metadata must be an object") ++ return cls(**record) ++ ++ ++def _head_range(layout: KVEntryLayout, tp_size: int, rank: int): ++ if type(tp_size) is not int or tp_size <= 0: ++ raise ValueError("TP size must be a positive integer") ++ if type(rank) is not int or not 0 <= rank < tp_size: ++ raise ValueError("TP rank is outside its group") ++ heads = layout.total_heads ++ if heads % tp_size and tp_size % heads: ++ raise ValueError("KV heads and TP size must divide one another") ++ local_heads = max(1, heads // tp_size) ++ replication = max(1, tp_size // heads) ++ start = (rank // replication) * local_heads ++ expected = layout.page_size * local_heads * layout.head_dim * layout.dtype_bytes ++ if layout.item_len != expected: ++ raise ValueError(f"KV entry stride mismatch: {layout.item_len} != {expected}") ++ return start, start + local_heads ++ ++ ++def plan_kv_entry_transfer( ++ *, ++ src: KVEntryLayout, ++ dst: KVEntryLayout, ++ src_ptr: int, ++ dst_ptr: int, ++ src_pages: Sequence[int], ++ dst_pages: Sequence[int], ++ src_tp_size: int, ++ dst_tp_size: int, ++ src_rank: int, ++ dst_rank: int, ++) -> list[tuple[int, int, int]]: ++ """Return physical source/destination spans, with no network side effects. ++ ++ GQA rank ordering matches QKVParallelLinear and send_kvcache_slice: ++ consecutive TP ranks replicate the same head when TP exceeds KV heads. ++ The caller retains bootstrap's existing rank fan-out and completion count. ++ """ ++ if len(src_pages) != len(dst_pages): ++ raise ValueError("Source and destination page counts differ") ++ if (src.kind, src.page_size, src.dtype, src.dtype_bytes) != ( ++ dst.kind, ++ dst.page_size, ++ dst.dtype, ++ dst.dtype_bytes, ++ ): ++ raise ValueError("Incompatible KV entry kind, page size or dtype") ++ if src_ptr <= 0 or dst_ptr <= 0: ++ raise ValueError("KV buffer pointers must be positive") ++ for pages, layout in ((src_pages, src), (dst_pages, dst)): ++ for page in pages: ++ if type(page) is not int or page < 0: ++ raise ValueError("KV page IDs must be non-negative integers") ++ if (page + 1) * layout.item_len > layout.buffer_nbytes: ++ raise ValueError("KV page exceeds its registered buffer") ++ ++ src_offset = dst_offset = 0 ++ length = src.item_len ++ outer_count = 1 ++ src_row_stride, dst_row_stride = src.item_len, dst.item_len ++ if src.kind == "flat": ++ if src.item_len != dst.item_len: ++ raise ValueError("Replicated KV entries require equal strides") ++ else: ++ if (src.total_heads, src.head_dim) != (dst.total_heads, dst.head_dim): ++ raise ValueError("Source and destination KV head geometry differ") ++ src_start, src_end = _head_range(src, src_tp_size, src_rank) ++ dst_start, dst_end = _head_range(dst, dst_tp_size, dst_rank) ++ if max(src_tp_size, dst_tp_size) % min(src_tp_size, dst_tp_size): ++ raise ValueError("Source and destination TP sizes must divide one another") ++ if src_tp_size <= dst_tp_size: ++ if src_rank != dst_rank // (dst_tp_size // src_tp_size): ++ raise ValueError("Source rank does not match the PD bootstrap mapping") ++ else: ++ writers = src_tp_size // dst_tp_size ++ if not dst_rank * writers <= src_rank < (dst_rank + 1) * writers: ++ raise ValueError("Source rank does not match the PD bootstrap mapping") ++ # Preserve acknowledgements from replicated source ranks, but avoid ++ # concurrent writes of the same head from equivalent replicas. ++ replicas = min(max(1, src_tp_size // src.total_heads), writers) ++ if src_rank % replicas: ++ return [] ++ start, end = max(src_start, dst_start), min(src_end, dst_end) ++ if start >= end: ++ raise ValueError("Registered PD ranks own disjoint KV heads") ++ head_bytes = src.head_dim * src.dtype_bytes ++ src_offset, dst_offset = (start - src_start) * head_bytes, ( ++ start - dst_start ++ ) * head_bytes ++ length = (end - start) * head_bytes ++ outer_count = src.page_size ++ src_row_stride, dst_row_stride = ( ++ src.item_len // outer_count, ++ dst.item_len // outer_count, ++ ) ++ ++ blocks = [] ++ for src_page, dst_page in zip(src_pages, dst_pages): ++ for token in range(outer_count): ++ source = ( ++ src_ptr + src_page * src.item_len + token * src_row_stride + src_offset ++ ) ++ destination = ( ++ dst_ptr + dst_page * dst.item_len + token * dst_row_stride + dst_offset ++ ) ++ if ( ++ blocks ++ and blocks[-1][0] + blocks[-1][2] == source ++ and blocks[-1][1] + blocks[-1][2] == destination ++ ): ++ previous = blocks[-1] ++ blocks[-1] = (previous[0], previous[1], previous[2] + length) ++ else: ++ blocks.append((source, destination, length)) ++ return blocks +diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py +--- a/python/sglang/srt/disaggregation/decode.py ++++ b/python/sglang/srt/disaggregation/decode.py +@@ -51,6 +51,7 @@ + ReqToMetadataIdxAllocator, + TransferBackend, + _is_fake_transfer, ++ build_dflash_kv_entry_layouts, + build_kv_layer_ids, + build_staging_slot_metadata, + get_dsv4_c128_state_indices, +@@ -86,6 +87,7 @@ + ) + from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool + from sglang.srt.mem_cache.memory_pool import ( ++ HybridLinearKVPool, + HybridReqToTokenPool, + KVCache, + ReqToTokenPool, +@@ -558,6 +560,21 @@ + kv_args.kv_data_ptrs = kv_data_ptrs + kv_args.kv_data_lens = kv_data_lens + kv_args.kv_item_lens = kv_item_lens ++ if self.scheduler.spec_algorithm.is_dflash() and isinstance( ++ self.token_to_kv_pool, HybridLinearKVPool ++ ): ++ if ( ++ self.transfer_backend != TransferBackend.MOONCAKE ++ or self.scheduler.enable_hisparse ++ ): ++ raise ValueError( ++ "Kimi DFlash mixed KV transfer requires Mooncake without HiSparse" ++ ) ++ kv_args.kv_entry_layouts = build_dflash_kv_entry_layouts( ++ self.token_to_kv_pool, ++ self.draft_token_to_kv_pool, ++ self.scheduler.draft_worker.draft_model_runner.model_config.get_total_num_kv_heads(), ++ ) + kv_args.kv_layer_ids = build_kv_layer_ids( + token_to_kv_pool=self.token_to_kv_pool, + draft_token_to_kv_pool=self.draft_token_to_kv_pool, +diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py +--- a/python/sglang/srt/disaggregation/mooncake/conn.py ++++ b/python/sglang/srt/disaggregation/mooncake/conn.py +@@ -2,6 +2,7 @@ + + import concurrent.futures + import dataclasses ++import json + import logging + import os + import struct +@@ -22,6 +23,10 @@ + CommonKVReceiver, + CommonKVSender, + KVTransferError, ++) ++from sglang.srt.disaggregation.common.kv_entry_layout import ( ++ KVEntryLayout, ++ plan_kv_entry_transfer, + ) + from sglang.srt.disaggregation.common.staging_handler import ( + STAGING_WATERMARK_WAIT_S, +@@ -147,6 +152,7 @@ + dst_kv_layer_ids: List[int] + dst_state_layer_ids: List[List[int]] + dst_state_types: List[StateType] = dataclasses.field(default_factory=list) ++ dst_kv_entry_layouts: Optional[List[dict]] = None + dst_dcp_size: int = 1 + dst_dcp_rank: int = 0 + requires_dcp_relayout: bool = False +@@ -185,6 +191,9 @@ + else [] + ), + dst_state_types=unpack_state_types(msg[19]) if len(msg) > 19 else [], ++ dst_kv_entry_layouts=( ++ json.loads(msg[20]) if len(msg) > 20 and msg[20] else None ++ ), + staging_base_ptr=( + struct.unpack("Q", msg[14])[0] + if len(msg) > 14 and len(msg[14]) == 8 +@@ -218,6 +227,16 @@ + self.init_engine() + self.register_buffer_to_engine() + self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() ++ if self.kv_args.kv_entry_layouts is not None: ++ if ( ++ self.enable_staging ++ or self.dcp_size != 1 ++ or self.attn_cp_size != 1 ++ or get_memory().enable_unified_memory ++ ): ++ raise ValueError( ++ "Mixed draft KV transfer requires unstaged, non-unified KV with CP=DCP=1" ++ ) + self.enable_trace = server_args.enable_trace + if self.disaggregation_mode == DisaggregationMode.PREFILL: + self.start_prefill_thread() +@@ -874,7 +893,24 @@ + dst_device_kv_indices: Optional[npt.NDArray[np.int32]] = None, + dst_kv_item_len: Optional[int] = None, + dst_attn_tp_size: Optional[int] = None, ++ dst_tp_rank: Optional[int] = None, ++ dst_kv_entry_layouts: Optional[List[dict]] = None, + ): ++ if self.kv_args.kv_entry_layouts is not None: ++ if dst_device_kv_indices is not None: ++ raise ValueError( ++ "Mixed draft KV transfer does not support device-indexed compressed KV" ++ ) ++ return self._send_mixed_kv_entries( ++ mooncake_session_id, ++ prefill_kv_indices, ++ dst_kv_ptrs, ++ dst_kv_indices, ++ dst_layer_ids, ++ dst_attn_tp_size, ++ dst_tp_rank, ++ dst_kv_entry_layouts, ++ ) + self._validate_envelope_kv_layout( + dst_kv_ptrs, dst_kv_item_len, dst_attn_tp_size + ) +@@ -902,6 +938,62 @@ + dst_device_data_indices=dst_device_kv_indices, + dst_device_data_ptrs=dst_device_kv_ptrs, + ) ++ ++ def _send_mixed_kv_entries( ++ self, ++ session_id, ++ src_pages, ++ dst_ptrs, ++ dst_pages, ++ dst_layer_ids, ++ dst_tp_size, ++ dst_rank, ++ dst_layouts, ++ ): ++ src_layouts = self.kv_args.kv_entry_layouts ++ if dst_layouts is None or dst_tp_size is None or dst_rank is None: ++ raise ValueError( ++ "Decode peer must advertise mixed draft KV geometry before transfer" ++ ) ++ if len(src_layouts) != len(self.kv_args.kv_data_ptrs) or len( ++ dst_layouts ++ ) != len(dst_ptrs): ++ raise ValueError("KV layout metadata length must match registered entries") ++ pairs = build_transfer_entry_pairs( ++ self.kv_args.kv_layer_ids, ++ dst_layer_ids or [], ++ len(src_layouts), ++ len(dst_layouts), ++ allow_positional_fallback=False, ++ ) ++ source_pages = src_pages.tolist() ++ destination_pages = dst_pages.tolist() ++ blocks = [] ++ for i, j in pairs: ++ source_layout = KVEntryLayout.from_wire(src_layouts[i]) ++ destination_layout = KVEntryLayout.from_wire(dst_layouts[j]) ++ if ( ++ source_layout.item_len != self.kv_args.kv_item_lens[i] ++ or source_layout.buffer_nbytes != self.kv_args.kv_data_lens[i] ++ ): ++ raise ValueError("Source KV registration changed after layout capture") ++ blocks.extend( ++ plan_kv_entry_transfer( ++ src=source_layout, ++ dst=destination_layout, ++ src_ptr=self.kv_args.kv_data_ptrs[i], ++ dst_ptr=dst_ptrs[j], ++ src_pages=source_pages, ++ dst_pages=destination_pages, ++ src_tp_size=self.attn_tp_size, ++ dst_tp_size=dst_tp_size, ++ src_rank=self.kv_args.engine_rank % self.attn_tp_size, ++ dst_rank=dst_rank % dst_tp_size, ++ ) ++ ) ++ # Validate the complete chunk before issuing any write, including when ++ # a later draft entry fails after valid target MLA entries were planned. ++ return self._transfer_data(session_id, blocks) + + def send_kvcache_dcp( + self, +@@ -1852,6 +1944,8 @@ + dst_device_kv_indices=chunked_dst_device_kv_indice, + dst_kv_item_len=target_rank_registration_info.dst_kv_item_len, + dst_attn_tp_size=target_rank_registration_info.dst_attn_tp_size, ++ dst_tp_rank=target_rank_registration_info.dst_tp_rank, ++ dst_kv_entry_layouts=target_rank_registration_info.dst_kv_entry_layouts, + ) + elif ( + self.enable_staging +@@ -2551,6 +2645,9 @@ + dst_dcp_rank, + packed_staging_slot_layer_ids, + packed_state_types, ++ json.dumps(self.kv_mgr.kv_args.kv_entry_layouts).encode( ++ "utf-8" ++ ), + ] + ) + except zmq.ZMQError: +diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py +--- a/python/sglang/srt/disaggregation/prefill.py ++++ b/python/sglang/srt/disaggregation/prefill.py +@@ -43,6 +43,7 @@ + MetadataBuffers, + ReqToMetadataIdxAllocator, + TransferBackend, ++ build_dflash_kv_entry_layouts, + build_kv_layer_ids, + build_staging_slot_metadata, + get_dsv4_c128_state_indices, +@@ -252,6 +253,18 @@ + kv_args.kv_data_ptrs = kv_data_ptrs + kv_args.kv_data_lens = kv_data_lens + kv_args.kv_item_lens = kv_item_lens ++ if self.scheduler.spec_algorithm.is_dflash() and isinstance( ++ self.token_to_kv_pool, HybridLinearKVPool ++ ): ++ if self.transfer_backend != TransferBackend.MOONCAKE: ++ raise ValueError( ++ "Kimi DFlash mixed KV transfer currently requires Mooncake" ++ ) ++ kv_args.kv_entry_layouts = build_dflash_kv_entry_layouts( ++ self.token_to_kv_pool, ++ draft_kv_pool, ++ self.scheduler.draft_worker.draft_model_runner.model_config.get_total_num_kv_heads(), ++ ) + kv_args.kv_layer_ids = build_kv_layer_ids( + token_to_kv_pool=self.token_to_kv_pool, + draft_token_to_kv_pool=draft_kv_pool, +diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py +--- a/python/sglang/srt/disaggregation/utils.py ++++ b/python/sglang/srt/disaggregation/utils.py +@@ -915,6 +915,61 @@ + return [(i, i) for i in range(n_src)] + + ++def build_dflash_kv_entry_layouts(target_pool, draft_pool, total_draft_heads: int): ++ """Describe Kimi's replicated MLA entries and the DFlash draft's GQA entries.""" ++ from sglang.srt.disaggregation.common.kv_entry_layout import KVEntryLayout ++ from sglang.srt.mem_cache.memory_pool import ( ++ HybridLinearKVPool, ++ MHATokenToKVPool, ++ MLATokenToKVPool, ++ ) ++ ++ if not isinstance(target_pool, HybridLinearKVPool): ++ raise ValueError("Mixed DFlash KV transfer requires a hybrid MLA target pool") ++ target = target_pool.full_kv_pool ++ if not isinstance(target, MLATokenToKVPool): ++ raise ValueError( ++ "Mixed DFlash KV transfer requires replicated MLA target entries" ++ ) ++ _, sizes, strides = target_pool.get_contiguous_buf_infos() ++ records = [ ++ KVEntryLayout( ++ "flat", ++ target_pool.page_size, ++ stride, ++ size, ++ str(target.dtype), ++ target.store_dtype.itemsize, ++ ).to_wire() ++ for size, stride in zip(sizes, strides) ++ ] ++ if draft_pool is None: ++ return records ++ if type(draft_pool) is not MHATokenToKVPool or draft_pool.kv_cache_layout != "nhd": ++ raise ValueError("DFlash PD draft KV requires the ordinary NHD MHA pool") ++ _, sizes, strides = draft_pool.get_contiguous_buf_infos() ++ num_layers = len(draft_pool.k_buffer) ++ if len(sizes) != 2 * num_layers or len(strides) != len(sizes): ++ raise ValueError( ++ "DFlash draft KV must expose separate K/V entries for every layer" ++ ) ++ for index, (size, stride) in enumerate(zip(sizes, strides)): ++ dim = draft_pool.head_dim if index < num_layers else draft_pool.v_head_dim ++ records.append( ++ KVEntryLayout( ++ "nhd", ++ draft_pool.page_size, ++ stride, ++ size, ++ str(draft_pool.dtype), ++ draft_pool.store_dtype.itemsize, ++ total_draft_heads, ++ dim, ++ ).to_wire() ++ ) ++ return records ++ ++ + def build_kv_layer_ids( + *, + token_to_kv_pool, +diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py +--- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py ++++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutlass.py +@@ -8,6 +8,7 @@ + + from __future__ import annotations + ++import inspect + from dataclasses import dataclass + from typing import TYPE_CHECKING, Optional + +@@ -89,6 +90,10 @@ + swiglu_beta: Optional[torch.Tensor] = None + swiglu_limit: Optional[torch.Tensor] = None + ++ # Kimi uses the public CUTLASS SiTU API from FlashInfer #4460. ++ situ_beta: Optional[torch.Tensor] = None ++ situ_linear_beta: Optional[torch.Tensor] = None ++ + # Bailing clamps after SiLU, which the kernel only implements in its + # SwigluStep variant. + use_swiglu_step: bool = False +@@ -115,6 +120,19 @@ + return cutlass_fused_moe, ActivationType + + ++def flashinfer_cutlass_supports_situ() -> bool: ++ try: ++ fused_moe, activation_type = _flashinfer_cutlass_fused_moe() ++ parameters = inspect.signature(fused_moe).parameters ++ except (ImportError, RuntimeError, TypeError, ValueError): ++ return False ++ return ( ++ hasattr(activation_type, "Situ") ++ and "situ_beta" in parameters ++ and "situ_linear_beta" in parameters ++ ) ++ ++ + def _activation_type(runner_config: MoeRunnerConfig): + from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import get_activation_type + +@@ -131,6 +149,8 @@ + ActivationType.Relu2, + ActivationType.Identity, + } ++ if hasattr(ActivationType, "Situ"): ++ supported.add(ActivationType.Situ) + assert activation in supported, ( + f"Activation {runner_config.activation!r} " + f"(is_gated={runner_config.is_gated}) maps to {activation.name}, " +@@ -346,6 +366,7 @@ + if weight_global_scale is not None: + from flashinfer import mxfp8_quantize + ++ x = x.contiguous() + x, input_sf = mxfp8_quantize( + x, + is_sf_swizzled_layout=True, +@@ -369,6 +390,22 @@ + output_dtype = torch.bfloat16 + with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()): + out = torch.empty(x.shape[0], out_hidden, dtype=output_dtype, device=x.device) ++ ++ situ_kwargs = {} ++ if runner_config.activation == "situ": ++ if not flashinfer_cutlass_supports_situ(): ++ raise RuntimeError("FlashInfer CUTLASS SiTU public API is unavailable") ++ activation_type = ActivationType.Situ ++ situ_kwargs = { ++ "situ_beta": quant_info.situ_beta, ++ "situ_linear_beta": quant_info.situ_linear_beta, ++ } ++ else: ++ activation_type = ( ++ ActivationType.SwigluStep ++ if quant_info.use_swiglu_step ++ else ActivationType.Swiglu ++ ) + + flashinfer_cutlass_fused_moe( + input=x, +@@ -390,14 +427,11 @@ + ep_rank=quant_info.moe_ep_rank, + use_w4_group_scaling=not use_mxfp8_act_scaling, + use_mxfp8_act_scaling=use_mxfp8_act_scaling, +- activation_type=( +- ActivationType.SwigluStep +- if quant_info.use_swiglu_step +- else ActivationType.Swiglu +- ), ++ activation_type=activation_type, + tune_max_num_tokens=next_power_of_2(x.shape[0]), + output=out, + use_fused_finalize=envs.SGLANG_FLASHINFER_MOE_FUSED_FINALIZE.get(), ++ **situ_kwargs, + ) + + if do_pad: +diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py +--- a/python/sglang/srt/layers/quantization/mxfp4.py ++++ b/python/sglang/srt/layers/quantization/mxfp4.py +@@ -1209,15 +1209,14 @@ + torch.cuda.empty_cache() + + def _process_weights_for_sm120_cutlass(self, layer): +- """Prepare GPT-OSS MXFP4 experts for FlashInfer CUTLASS on SM120. +- +- GPT-OSS stores gate/up rows pair-wise as +- ``[gate_0, up_0, gate_1, up_1, ...]``. FlashInfer's fused MoE consumes +- two contiguous halves in ``[up; gate]`` order. Build that layout after +- checkpoint loading so padding cannot move the split, pad both GEMMs to +- CUTLASS's 128-element alignment, and swizzle the native E8M0 scales for +- the SM120 MXFP8-by-MXFP4 kernels. Packed FP4 weight bytes themselves do +- not need an SM120 permutation. ++ """Prepare MXFP4 experts for FlashInfer CUTLASS on SM120. ++ ++ FlashInfer consumes two contiguous halves in ``[up; gate]`` order. ++ GPT-OSS checkpoints store pair-interleaved ``[gate_i, up_i]`` rows, ++ while Kimi-K3 stores contiguous ``[gate; up]`` halves. Normalize either ++ layout after loading, pad both GEMMs to CUTLASS's 128-element alignment, ++ and swizzle the native E8M0 scales. Packed FP4 bytes need no additional ++ SM120 permutation. + """ + from flashinfer import block_scale_interleave + +@@ -1229,9 +1228,15 @@ + E = layer.num_local_experts + device = layer.w13_weight.device + ++ gate_up_interleaved = self.moe_runner_config.gate_up_interleaved ++ ++ def _split_gate_up(unpadded): ++ if gate_up_interleaved: ++ return unpadded[:, 0::2, :], unpadded[:, 1::2, :] ++ return unpadded[:, :N_un, :], unpadded[:, N_un : 2 * N_un, :] ++ + def _stack_up_gate_w13(unpadded, last_pad, last_un): +- gate_rows = unpadded[:, 0::2, :] +- up_rows = unpadded[:, 1::2, :] ++ gate_rows, up_rows = _split_gate_up(unpadded) + out = torch.zeros( + E, 2 * N_pad, last_pad, dtype=unpadded.dtype, device=device + ) +@@ -1248,8 +1253,14 @@ + + bias_dtype = layer.w13_weight_bias.dtype + w13_bias_padded = torch.zeros(E, 2 * N_pad, dtype=bias_dtype, device=device) +- w13_bias_padded[:, :N_un] = layer.w13_weight_bias.data[:, 1::2] +- w13_bias_padded[:, N_pad : N_pad + N_un] = layer.w13_weight_bias.data[:, 0::2] ++ if gate_up_interleaved: ++ gate_bias = layer.w13_weight_bias.data[:, 0::2] ++ up_bias = layer.w13_weight_bias.data[:, 1::2] ++ else: ++ gate_bias = layer.w13_weight_bias.data[:, :N_un] ++ up_bias = layer.w13_weight_bias.data[:, N_un : 2 * N_un] ++ w13_bias_padded[:, :N_un] = up_bias ++ w13_bias_padded[:, N_pad : N_pad + N_un] = gate_bias + + def _pad_w2_3d(unpadded, last_pad, last_un): + out = torch.zeros(E, K_pad, last_pad, dtype=unpadded.dtype, device=device) +@@ -1281,17 +1292,59 @@ + layer.w13_weight_bias = Parameter(w13_bias_padded, requires_grad=False) + layer.w2_weight_bias = Parameter(w2_bias_padded, requires_grad=False) + +- layer.swiglu_alpha = Parameter( +- torch.full((E,), 1.702, dtype=torch.float32, device=device), +- requires_grad=False, +- ) +- layer.swiglu_beta = Parameter( +- torch.ones(E, dtype=torch.float32, device=device), +- requires_grad=False, +- ) +- layer.swiglu_limit = Parameter( +- torch.full((E,), 7.0, dtype=torch.float32, device=device), +- requires_grad=False, ++ activation = self.moe_runner_config.activation ++ alpha = self.moe_runner_config.gemm1_alpha ++ beta = self.moe_runner_config.gemm1_beta ++ limit = self.moe_runner_config.gemm1_clamp_limit ++ if activation == "situ": ++ situ_beta = 4.0 if alpha is None else alpha ++ situ_linear_beta = 25.0 if limit is None else limit ++ alpha = beta = limit = None ++ else: ++ situ_beta = situ_linear_beta = None ++ alpha = 1.702 if alpha is None else alpha ++ beta = 1.0 if beta is None else beta ++ limit = 7.0 if limit is None else limit ++ ++ layer.swiglu_alpha = ( ++ None ++ if alpha is None ++ else Parameter( ++ torch.full((E,), alpha, dtype=torch.float32, device=device), ++ requires_grad=False, ++ ) ++ ) ++ layer.swiglu_beta = ( ++ None ++ if beta is None ++ else Parameter( ++ torch.full((E,), beta, dtype=torch.float32, device=device), ++ requires_grad=False, ++ ) ++ ) ++ layer.swiglu_limit = ( ++ None ++ if limit is None ++ else Parameter( ++ torch.full((E,), limit, dtype=torch.float32, device=device), ++ requires_grad=False, ++ ) ++ ) ++ layer.situ_beta = ( ++ None ++ if situ_beta is None ++ else Parameter( ++ torch.full((E,), situ_beta, dtype=torch.float32, device=device), ++ requires_grad=False, ++ ) ++ ) ++ layer.situ_linear_beta = ( ++ None ++ if situ_linear_beta is None ++ else Parameter( ++ torch.full((E,), situ_linear_beta, dtype=torch.float32, device=device), ++ requires_grad=False, ++ ) + ) + # The MXFP4 ABI uses a neutral global weight scale for each GEMM. + layer.mxfp4_weight_global_scale = Parameter( +@@ -1334,6 +1387,22 @@ + "cutlass_sm90", + "cutlass_sm120", + ): ++ if ( ++ self._fi_kernel == "cutlass_sm120" ++ and moe_runner_config.activation == "situ" ++ ): ++ from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import ( ++ flashinfer_cutlass_supports_situ, ++ ) ++ ++ if not flashinfer_cutlass_supports_situ(): ++ raise RuntimeError( ++ "Kimi-K3 FlashInfer MXFP4 MoE on SM120 requires a " ++ "FlashInfer build with CUTLASS SiTU support " ++ "(ActivationType.Situ plus situ_beta/situ_linear_beta). " ++ "Upgrade FlashInfer or restart with " ++ "--moe-runner-backend marlin." ++ ) + # Register the fused func at runner construction so the FusedOpPool + # lookup at `MoeRunner.__init__` finds it. + import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401 +@@ -1394,7 +1463,7 @@ + ) + + def _apply_sm120_cutlass(self, layer, dispatch_output): +- """SM120 GPT-OSS MXFP8 x MXFP4 MoE via FlashInfer CUTLASS.""" ++ """SM120 MXFP8 x MXFP4 MoE via FlashInfer CUTLASS.""" + from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import ( + FlashInferCutlassMxfp4MoeQuantInfo, + ) +@@ -1405,11 +1474,13 @@ + w13_weight_scale=layer.w13_weight_scale, + w2_weight_scale=layer.w2_weight_scale, + mxfp4_weight_global_scale=layer.mxfp4_weight_global_scale, +- w13_bias=layer.w13_weight_bias, +- w2_bias=layer.w2_weight_bias, ++ w13_bias=layer.w13_weight_bias if self.with_bias else None, ++ w2_bias=layer.w2_weight_bias if self.with_bias else None, + swiglu_alpha=layer.swiglu_alpha, + swiglu_beta=layer.swiglu_beta, + swiglu_limit=layer.swiglu_limit, ++ situ_beta=layer.situ_beta, ++ situ_linear_beta=layer.situ_linear_beta, + moe_tp_size=layer.moe_tp_size, + moe_tp_rank=layer.moe_tp_rank, + moe_ep_size=layer.moe_ep_size, +diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py +--- a/python/sglang/srt/managers/scheduler_pp_mixin.py ++++ b/python/sglang/srt/managers/scheduler_pp_mixin.py +@@ -1207,7 +1207,7 @@ + logits_output.auxiliary_device_output = auxiliary_output + next_token_ids = pp_outputs["next_token_ids"].to(torch.int64) + next_draft_input = None +- if isinstance(batch, ScheduleBatch) and batch.spec_algorithm.is_dspark(): ++ if isinstance(batch, ScheduleBatch) and batch.spec_algorithm.is_dflash_family(): + next_token_ids = next_token_ids.to( + device=batch.device, + dtype=torch.int64, +@@ -1216,6 +1216,11 @@ + from sglang.srt.speculative.dspark_components.dspark_draft import ( + make_next_draft_input, + ) ++ ++ if batch.spec_algorithm.is_dflash(): ++ from sglang.srt.speculative.draft_worker_common import ( ++ make_draft_input_v2 as make_next_draft_input, ++ ) + + next_draft_input = make_next_draft_input( + bonus_tokens=next_token_ids, +@@ -1388,14 +1393,14 @@ + "set_run_batch_cpu_start_time", + trace_only=True, + ) +- if cur_batch.spec_algorithm.is_dspark(): ++ if cur_batch.spec_algorithm.is_dflash_family(): + self.model_worker.set_pp_proxy_tensors_for_next_forward( + pp_proxy_tensors + ) + try: + result = self.run_batch(cur_batch, pp_proxy_tensors) + finally: +- if cur_batch.spec_algorithm.is_dspark(): ++ if cur_batch.spec_algorithm.is_dflash_family(): + self.model_worker.set_pp_proxy_tensors_for_next_forward(None) + set_time_batch( + cur_batch.reqs, +diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py +--- a/python/sglang/srt/model_executor/runner/base_runner.py ++++ b/python/sglang/srt/model_executor/runner/base_runner.py +@@ -567,12 +567,15 @@ + 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, +- ) ++ # PD prefill warms up without speculative state or verify metadata. ++ spec_info = None ++ if not _is_pd_prefill_target: ++ 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() + ): +diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py +--- a/python/sglang/srt/models/kimi_k3.py ++++ b/python/sglang/srt/models/kimi_k3.py +@@ -2755,6 +2755,15 @@ + ) + sp_sharded = False + aux_hidden_states = [] ++ if ( ++ self.dspark_layers_to_capture is not None ++ and self.start_layer - 1 in self.dspark_layers_to_capture ++ ): ++ aux_hidden_states.append( ++ self._dspark_capture_stream( ++ self.start_layer - 1, hidden_states, residual, attn_res ++ ) ++ ) + for i in range(self.start_layer, self.end_layer): + if sp_sharded and not self.layers[i]._sp_moe: + hidden_states = _sp_all_gather_rows(hidden_states) +@@ -2919,6 +2928,18 @@ + for layer_id in layer_ids + if self.model.start_layer <= int(layer_id) < self.model.end_layer + ] ++ self.capture_aux_hidden_states = bool(local_layer_ids) ++ self.model.dspark_layers_to_capture = local_layer_ids or None ++ ++ def set_dflash_layers_to_capture(self, layer_ids: list[int]) -> None: ++ from sglang.srt.speculative.dflash_pp import kimi_pp_capture_layer_ids ++ ++ local_layer_ids = kimi_pp_capture_layer_ids( ++ layer_ids, ++ self.model.start_layer, ++ self.model.end_layer, ++ self.config.num_hidden_layers, ++ ) + self.capture_aux_hidden_states = bool(local_layer_ids) + self.model.dspark_layers_to_capture = local_layer_ids or None + +@@ -3403,6 +3424,11 @@ + ) + self.language_model.set_dspark_layers_to_capture(layer_ids) + ++ def set_dflash_layers_to_capture(self, layer_ids: list[int]) -> None: ++ if self.language_model is None: ++ raise AttributeError("DFLASH capture is unavailable in encoder-only mode") ++ self.language_model.set_dflash_layers_to_capture(layer_ids) ++ + def preprocess_mm_for_encoder( + self, + mm_data, +diff --git a/python/sglang/srt/speculative/dflash_pp.py b/python/sglang/srt/speculative/dflash_pp.py +new file mode 100644 +--- /dev/null ++++ b/python/sglang/srt/speculative/dflash_pp.py +@@ -0,0 +1,19 @@ ++def kimi_pp_capture_layer_ids(layer_ids, start_layer, end_layer, num_layers): ++ """Own stream captures where their next consumer's weights are local. ++ ++ Kimi's post-layer stream uses the next layer's attention-residual weights. ++ A PP-boundary capture therefore belongs to the next stage, before its first ++ decoder layer, rather than to the stage that just produced the raw stream. ++ """ ++ ids = list(layer_ids) ++ if not ids or any(type(i) is not int for i in ids): ++ raise ValueError("DFLASH requires explicit integer capture layer IDs") ++ if len(ids) != len(set(ids)) or min(ids) < 0 or max(ids) >= num_layers: ++ raise ValueError("DFLASH capture layer IDs must be unique and in range") ++ if ids != sorted(ids): ++ raise ValueError("Kimi DFLASH capture layer IDs must follow model layer order") ++ if not 0 <= start_layer < end_layer <= num_layers: ++ raise ValueError("Invalid Kimi PP layer range") ++ return sorted( ++ i for i in ids if start_layer <= min(i + 1, num_layers - 1) < end_layer ++ ) +diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py +--- a/python/sglang/srt/speculative/dflash_worker_v2.py ++++ b/python/sglang/srt/speculative/dflash_worker_v2.py +@@ -34,6 +34,7 @@ + compute_position, + ) + from sglang.srt.runtime_context import ( ++ get_disagg, + get_exec, + get_schedule, + get_spec, +@@ -43,6 +44,7 @@ + from sglang.srt.speculative.base_spec_worker import BaseSpecWorker + from sglang.srt.speculative.dflash_info import DFlashVerifyInput + from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 ++from sglang.srt.speculative.dflash_pp import kimi_pp_capture_layer_ids + from sglang.srt.speculative.dflash_utils import ( + apply_dflash_simulated_acceptance, + apply_dflash_verify_logits_adjustments, +@@ -292,6 +294,13 @@ + self.nccl_port = nccl_port + self._target_worker = target_worker + self.model_runner = target_worker.model_runner ++ self._is_pd_prefill = get_disagg().disaggregation_mode == "prefill" ++ self._is_context_only_pp_prefill_rank = ( ++ self._is_pd_prefill and ps.pp_rank < ps.pp_size - 1 ++ ) ++ self._next_pp_proxy_tensors = None ++ self._pp_context_feature_indices = [] ++ self._pp_expects_incoming_context = False + self._need_mamba_verify_commit = False + self.page_size = get_schedule().page_size + # Normalized in arg_groups.speculative_hook.handle_speculative_decoding. +@@ -315,6 +324,8 @@ + self.draft_model_runner = bundle.draft_model_runner + self._draft_sampler = None + self.draft_model = bundle.draft_model ++ if ps.pp_size > 1: ++ self._init_pp_context_features() + self.selector = self.draft_model.candidate_selector + draft_config = parse_dflash_draft_config( + draft_hf_config=self.draft_model_runner.model_config.hf_config +@@ -396,7 +407,9 @@ + self._draft_greedy_rank_index_buf: Optional[torch.Tensor] = None + self._draft_greedy_selected_ids_buf: Optional[torch.Tensor] = None + self._draft_greedy_index_cap: int = 0 +- self._use_fused_kv_materialize = is_cuda() or is_hip() ++ self._use_fused_kv_materialize = ( ++ is_cuda() or is_hip() ++ ) and not self._is_pd_prefill + self._fused_kv_helper: Optional[object] = None + if self._use_fused_kv_materialize: + self._init_fused_kv_helper() +@@ -423,10 +436,43 @@ + # EagleDraftWorkerBase draft/draft_extend split to wrap it in. + return self._draft_worker + ++ def _init_pp_context_features(self): ++ target_model = self.model_runner.model ++ if ( ++ not self._is_pd_prefill ++ or not type(target_model).__name__.startswith("KimiK3") ++ or type(self.draft_model).__name__ != "DFlashDraftModel" ++ ): ++ raise ValueError( ++ "PP DFLASH currently requires a Kimi-K3 PD prefill target " ++ "and a DFlashDraftModel checkpoint" ++ ) ++ layer_ids = self.model_runner.spec_aux_config.dflash_target_layer_ids ++ if not layer_ids or len(layer_ids) != self.draft_model.num_context_features: ++ raise ValueError("DFLASH capture count does not match the draft projection") ++ info = self.model_runner.layer_info ++ num_layers = self.model_runner.model_config.num_hidden_layers ++ local_ids = kimi_pp_capture_layer_ids( ++ layer_ids, info.start_layer, info.end_layer, num_layers ++ ) ++ self._pp_context_feature_indices = [layer_ids.index(i) for i in local_ids] ++ self._pp_expects_incoming_context = any( ++ min(i + 1, num_layers - 1) < info.start_layer for i in layer_ids ++ ) ++ logger.info( ++ "DFLASH PP rank %s: capture layers=%s, projection columns=%s, incoming=%s", ++ self.ps.pp_rank, ++ local_ids, ++ self._pp_context_feature_indices, ++ self._pp_expects_incoming_context, ++ ) ++ + @property + def spec_v2_attn_backends(self) -> tuple: + # Every attn backend a spec_v2 forward touches; consumed by + # decide_needs_cpu_seq_lens to gate the seq_lens_cpu D2H. ++ if self._is_context_only_pp_prefill_rank: ++ return (self._target_worker.model_runner.attn_backend,) + return ( + self._target_worker.model_runner.attn_backend, + self.draft_model_runner.attn_backend, +@@ -443,6 +489,21 @@ + # enabled, the draft worker keeps a private compact req->token table + # over the same global KV index space, so radix-cache/prefix-hit KV + # remains reusable while draft attention sees only the recent window. ++ if memory_pool_config is not None and self._is_context_only_pp_prefill_rank: ++ memory_pool_config = replace( ++ memory_pool_config, ++ max_total_num_tokens=self.page_size, ++ full_max_total_num_tokens=( ++ self.page_size ++ if memory_pool_config.full_max_total_num_tokens ++ else memory_pool_config.full_max_total_num_tokens ++ ), ++ swa_max_total_num_tokens=( ++ self.page_size ++ if memory_pool_config.swa_max_total_num_tokens ++ else memory_pool_config.swa_max_total_num_tokens ++ ), ++ ) + self._draft_worker.alloc_memory_pool( + memory_pool_config=memory_pool_config, + req_to_token_pool=( +@@ -452,6 +513,8 @@ + ) + + def init_attention_backends(self): ++ if self._is_context_only_pp_prefill_rank: ++ return + self._draft_worker.init_attention_backends() + self._need_mamba_verify_commit = mambaish_config( + self.model_runner.model_config +@@ -461,6 +524,9 @@ + ) + + def init_cuda_graphs(self): ++ # P writes context KV only; it never proposes or verifies draft blocks. ++ if self._is_pd_prefill: ++ return + capture_decode_cuda_graph = ( + get_exec().graph.cuda_graph_config.decode.backend != Backend.DISABLED + ) +@@ -1657,6 +1723,99 @@ + ) -> DFlashDraftInputV2: + return make_draft_input_v2(bonus_tokens=bonus_tokens, new_seq_lens=new_seq_lens) + ++ def set_pp_proxy_tensors_for_next_forward(self, pp_proxy_tensors): ++ self._next_pp_proxy_tensors = pp_proxy_tensors ++ ++ @torch.no_grad() ++ def _forward_pp_prefill(self, batch, on_publish, pp_proxy_tensors): ++ result = self.target_worker.forward_batch_generation( ++ batch, ++ pp_proxy_tensors=pp_proxy_tensors, ++ capture_hidden_mode=CaptureHiddenMode.FULL, ++ ) ++ output = result.pp_hidden_states_proxy_tensors ++ logits = result.logits_output ++ target_hidden = ( ++ logits.hidden_states ++ if logits is not None ++ else ( ++ output.tensors.get("dspark_aux_hidden_states") ++ if output is not None ++ else None ++ ) ++ ) ++ incoming = ( ++ pp_proxy_tensors.tensors.get("dflash_ctx_acc") ++ if pp_proxy_tensors is not None ++ else None ++ ) ++ if (incoming is not None) != self._pp_expects_incoming_context: ++ raise RuntimeError("DFLASH PP context missing or unexpectedly duplicated") ++ if batch.extend_lens is None or batch.prefix_lens is None: ++ raise RuntimeError("DFLASH PP prefill requires extend_lens and prefix_lens") ++ if batch.out_cache_loc is None: ++ raise RuntimeError("DFLASH PP prefill requires out_cache_loc") ++ tokens = sum(batch.extend_lens) ++ shape = (tokens, self.draft_model.config.hidden_size) ++ local = None ++ if self._pp_context_feature_indices: ++ if target_hidden is None or target_hidden.shape[0] != tokens: ++ raise RuntimeError( ++ "DFLASH PP local hidden capture is missing or truncated" ++ ) ++ local = self.draft_model.project_target_hidden_partial( ++ target_hidden, self._pp_context_feature_indices ++ ) ++ elif logits is None and target_hidden is not None and target_hidden.numel(): ++ raise RuntimeError("DFLASH PP captured unassigned layer features") ++ for context in (incoming, local): ++ if context is not None and tuple(context.shape) != shape: ++ raise RuntimeError("DFLASH PP accumulated context shape mismatch") ++ context = incoming ++ if local is not None: ++ context = local if incoming is None else incoming.to(local) + local ++ ++ if self.ps.pp_rank < self.ps.pp_size - 1: ++ if output is None: ++ raise RuntimeError( ++ "DFLASH non-final PP stage did not return proxy tensors" ++ ) ++ output.tensors.pop("dspark_aux_hidden_states", None) ++ if context is not None: ++ output.tensors["dflash_ctx_acc"] = context ++ else: ++ if output is not None or context is None or result.next_token_ids is None: ++ raise RuntimeError( ++ "DFLASH final PP stage lacks complete context or logits" ++ ) ++ prefixes = torch.tensor( ++ batch.prefix_lens, dtype=torch.int32, device=self.device ++ ) ++ extends = torch.tensor( ++ batch.extend_lens, dtype=torch.int32, device=self.device ++ ) ++ positions, _ = compute_position( ++ self.model_runner.prefill_attention_backend_str, ++ prefixes, ++ extends, ++ tokens, ++ ) ++ # The linear partials are summed before applying RMSNorm exactly once. ++ self._append_target_hidden_sequential( ++ ctx_hidden=self.draft_model.hidden_norm(context), ++ ctx_positions=positions.to(dtype=torch.int64), ++ ctx_cache_loc=batch.out_cache_loc.to(dtype=torch.int64), ++ ) ++ result.next_draft_input = self._make_next_draft_input_prefill( ++ bonus_tokens=result.next_token_ids, seq_lens=batch.seq_lens ++ ) ++ if logits is not None: ++ logits.hidden_states = None ++ result.new_seq_lens = batch.seq_lens ++ if on_publish is not None: ++ on_publish(result.new_seq_lens) ++ return result ++ + def forward_batch_generation( + self, + batch: ScheduleBatch, +@@ -1664,9 +1823,14 @@ + grammar_barrier=None, + pp_proxy_tensors=None, + ) -> GenerationBatchResult: ++ if pp_proxy_tensors is None: ++ pp_proxy_tensors = self._next_pp_proxy_tensors ++ self._next_pp_proxy_tensors = None + self._validate_phase1_sampling_support(batch) + + if batch.forward_mode.is_extend() or batch.is_extend_in_batch: ++ if self.ps.pp_size > 1: ++ return self._forward_pp_prefill(batch, on_publish, pp_proxy_tensors) + # Target prefill: capture DFlash aux hidden states for prompt tokens. + batch_output = self.target_worker.forward_batch_generation( + batch, +diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py +--- a/python/sglang/srt/speculative/spec_info.py ++++ b/python/sglang/srt/speculative/spec_info.py +@@ -188,6 +188,14 @@ + ) + + return build_dspark_disagg_draft_input( ++ batch, last_tokens_tensor, future_map ++ ) ++ if self.is_dflash(): ++ from sglang.srt.speculative.dflash_disaggregation import ( ++ build_dflash_family_disagg_draft_input, ++ ) ++ ++ return build_dflash_family_disagg_draft_input( + batch, last_tokens_tensor, future_map + ) + return None +diff --git a/test/registered/disaggregation/test_dflash_pp_context.py b/test/registered/disaggregation/test_dflash_pp_context.py +new file mode 100644 +--- /dev/null ++++ b/test/registered/disaggregation/test_dflash_pp_context.py +@@ -0,0 +1,431 @@ ++"""CPU contract tests of the production PP methods; no model/RDMA emulation claim.""" ++ ++import ast ++import logging ++import os ++from dataclasses import make_dataclass, replace ++from pathlib import Path ++import sys ++from types import SimpleNamespace ++import unittest ++from unittest.mock import patch ++ ++import torch ++import torch.nn.functional as F ++ ++ROOT = Path(os.environ.get("SGLANG_SOURCE_ROOT", Path(__file__).resolve().parents[3])) ++SRT = ROOT / "python/sglang/srt" ++ ++ ++def load_function(path, name, globals_dict=None, owner=None): ++ tree = ast.parse((SRT / path).read_text()) ++ scope = tree.body ++ if owner is not None: ++ scope = next( ++ n for n in scope if isinstance(n, ast.ClassDef) and n.name == owner ++ ).body ++ node = next(n for n in scope if isinstance(n, ast.FunctionDef) and n.name == name) ++ module = ast.Module( ++ body=[ ++ ast.ImportFrom( ++ module="__future__", names=[ast.alias(name="annotations")], level=0 ++ ), ++ node, ++ ], ++ type_ignores=[], ++ ) ++ namespace = {"torch": torch, "F": F, **(globals_dict or {})} ++ exec(compile(ast.fix_missing_locations(module), str(SRT / path), "exec"), namespace) ++ return namespace[name] ++ ++ ++capture_ids = load_function("speculative/dflash_pp.py", "kimi_pp_capture_layer_ids") ++partial = load_function( ++ "models/dflash.py", "project_target_hidden_partial", owner="DFlashDraftModel" ++) ++ ++ ++def positions_for_test(backend, prefix, extend, tokens): ++ positions = torch.cat( ++ [torch.arange(p, p + e) for p, e in zip(prefix.tolist(), extend.tolist())] ++ ) ++ assert positions.numel() == tokens ++ return positions, None ++ ++ ++forward_pp = load_function( ++ "speculative/dflash_worker_v2.py", ++ "_forward_pp_prefill", ++ owner="DFlashWorkerV2", ++ globals_dict={ ++ "CaptureHiddenMode": SimpleNamespace(FULL="full"), ++ "compute_position": positions_for_test, ++ }, ++) ++init_features = load_function( ++ "speculative/dflash_worker_v2.py", ++ "_init_pp_context_features", ++ owner="DFlashWorkerV2", ++ globals_dict={ ++ "kimi_pp_capture_layer_ids": capture_ids, ++ "logger": logging.getLogger(__name__), ++ }, ++) ++ ++ ++class CountingNorm(torch.nn.Module): ++ def __init__(self): ++ super().__init__() ++ self.calls = 0 ++ ++ def forward(self, x): ++ self.calls += 1 ++ return F.rms_norm(x, (x.shape[-1],), eps=1e-6) ++ ++ ++class DFlashDraftModel: ++ def __init__(self, weight): ++ self.num_context_features = weight.shape[1] // weight.shape[0] ++ self.config = SimpleNamespace(hidden_size=weight.shape[0]) ++ self.fc = SimpleNamespace(weight=weight) ++ self.hidden_norm = CountingNorm() ++ ++ def project_target_hidden_partial(self, hidden, indices): ++ return partial(self, hidden, indices) ++ ++ ++class TestPPContext(unittest.TestCase): ++ def test_prefill_dummy_forward_has_no_verify_metadata(self): ++ path = SRT / "model_executor/runner/base_runner.py" ++ tree = ast.parse(path.read_text()) ++ runner = next( ++ n ++ for n in tree.body ++ if isinstance(n, ast.ClassDef) and n.name == "BaseRunner" ++ ) ++ dummy = next( ++ n ++ for n in runner.body ++ if isinstance(n, ast.FunctionDef) and n.name == "_dummy_run" ++ ) ++ start = next( ++ i ++ for i, n in enumerate(dummy.body) ++ if isinstance(n, ast.Assign) ++ and any(isinstance(t, ast.Name) and t.id == "spec_info" for t in n.targets) ++ ) ++ block = ast.Module(body=dummy.body[start : start + 2], type_ignores=[]) ++ code = compile(ast.fix_missing_locations(block), str(path), "exec") ++ for prefill_target in (True, False): ++ calls = [] ++ expected = object() ++ ++ def create(*args): ++ calls.append(args) ++ return expected ++ ++ namespace = { ++ "_is_pd_prefill_target": prefill_target, ++ "create_dummy_verify_input": create, ++ "mr": SimpleNamespace(spec_algorithm="DFLASH", is_draft_worker=False), ++ "buffers": SimpleNamespace(custom_mask=None), ++ "num_tokens_per_req": 16, ++ } ++ exec(code, namespace) ++ if prefill_target: ++ self.assertIsNone(namespace["spec_info"]) ++ self.assertEqual(calls, []) ++ else: ++ self.assertIs(namespace["spec_info"], expected) ++ self.assertEqual(calls, [("DFLASH", None, 16, False)]) ++ ++ def test_capture_ownership_including_boundaries(self): ++ for layers in (92, 94, 96): ++ for pp in (1, 2, 4, 8, 16): ++ bounds = [i * layers // pp for i in range(pp + 1)] ++ for ids in ([19, 37, 54, 66, 78, 90], list(range(layers))): ++ assigned = [] ++ for start, end in zip(bounds, bounds[1:]): ++ owned = capture_ids(ids, start, end, layers) ++ assigned.extend(owned) ++ if end < layers: ++ self.assertNotIn(end - 1, owned) ++ self.assertEqual(assigned, sorted(ids)) ++ ++ def test_invalid_capture_configuration(self): ++ for ids in ([], [1, 1], [-1], [96], ["1"], [4, 2]): ++ with self.assertRaises(ValueError): ++ capture_ids(ids, 0, 12, 96) ++ ++ def test_kimi_boundary_uses_next_stage_weights(self): ++ calls = [] ++ aggregate = lambda *args: calls.append(args) or args[0] + 5 ++ capture = load_function( ++ "models/kimi_k3.py", ++ "_dspark_capture_stream", ++ owner="KimiK3LinearModel", ++ globals_dict={ ++ "aggregate_stream": aggregate, ++ "_cdiv": lambda x, y: (x + y - 1) // y, ++ }, ++ ) ++ next_layer = SimpleNamespace( ++ self_attention_res_proj="next_proj", ++ self_attention_res_norm="next_norm", ++ prev_valid_blocks=4, ++ ) ++ # No output_attn_res_* on a non-final PP stage, exactly as in Kimi's constructor. ++ stage = SimpleNamespace(end_layer=24, layers={12: next_layer}) ++ h, residual, bank = ( ++ torch.ones(2, 4), ++ torch.full((2, 4), 2.0), ++ torch.zeros(4, 2, 4), ++ ) ++ out = capture(stage, 11, h, residual, SimpleNamespace(block_residual=bank)) ++ torch.testing.assert_close(out, h + residual + 5) ++ self.assertEqual(calls[0][2:], (4, "next_proj", "next_norm")) ++ self.assertIs(calls[0][1], bank) ++ ++ def _run_pipeline(self, dtype, ids, drop_incoming=False): ++ torch.manual_seed(17) ++ hidden_size, tokens, pp, layers = 32, 5, 8, 96 ++ h = torch.randn(tokens, len(ids), hidden_size, dtype=dtype) ++ weight = torch.randn(hidden_size, len(ids) * hidden_size, dtype=dtype) / 8 ++ batch = SimpleNamespace( ++ extend_lens=[2, 3], ++ prefix_lens=[3, 7], ++ out_cache_loc=torch.tensor([9, 10, 21, 22, 23]), ++ seq_lens=torch.tensor([5, 10]), ++ ) ++ incoming, writes, norms, publications = None, [], [], [] ++ for rank in range(pp): ++ start, end = rank * 12, (rank + 1) * 12 ++ local_ids = capture_ids(ids, start, end, layers) ++ indices = [ids.index(i) for i in local_ids] ++ local_hidden = h[:, indices, :].reshape(tokens, -1) if indices else None ++ proxy = ( ++ None ++ if rank == pp - 1 ++ else SimpleNamespace( ++ tensors={"hidden_states": torch.zeros(tokens, hidden_size)} ++ ) ++ ) ++ if proxy is not None and local_hidden is not None: ++ proxy.tensors["dspark_aux_hidden_states"] = local_hidden ++ logits = ( ++ SimpleNamespace( ++ hidden_states=( ++ local_hidden if indices else torch.zeros(tokens, hidden_size) ++ ) ++ ) ++ if rank == pp - 1 ++ else None ++ ) ++ result = SimpleNamespace( ++ pp_hidden_states_proxy_tensors=proxy, ++ logits_output=logits, ++ next_token_ids=torch.tensor([101, 102]) if rank == pp - 1 else None, ++ ) ++ draft = DFlashDraftModel(weight) ++ norms.append(draft.hidden_norm) ++ worker = SimpleNamespace( ++ ps=SimpleNamespace(pp_rank=rank, pp_size=pp), ++ device="cpu", ++ _is_pd_prefill=True, ++ draft_model=draft, ++ model_runner=SimpleNamespace( ++ model=type("KimiK3ForConditionalGeneration", (), {})(), ++ spec_aux_config=SimpleNamespace(dflash_target_layer_ids=ids), ++ layer_info=SimpleNamespace(start_layer=start, end_layer=end), ++ model_config=SimpleNamespace(num_hidden_layers=layers), ++ prefill_attention_backend_str="test", ++ ), ++ target_worker=SimpleNamespace( ++ forward_batch_generation=lambda *a, **k: result ++ ), ++ _append_target_hidden_sequential=lambda **kwargs: writes.append(kwargs), ++ _make_next_draft_input_prefill=lambda **kwargs: SimpleNamespace( ++ **kwargs ++ ), ++ ) ++ init_features(worker) ++ if drop_incoming and worker._pp_expects_incoming_context: ++ incoming = None ++ output = forward_pp(worker, batch, publications.append, incoming) ++ if proxy is not None: ++ self.assertNotIn("dspark_aux_hidden_states", proxy.tensors) ++ self.assertEqual(len(writes), 0) ++ incoming = proxy ++ self.assertEqual(len(writes), 1) ++ self.assertEqual(sum(n.calls for n in norms), 1) ++ self.assertEqual(norms[-1].calls, 1) ++ self.assertEqual(len(publications), pp) ++ self.assertIsNone(output.logits_output.hidden_states) ++ torch.testing.assert_close( ++ output.next_draft_input.bonus_tokens, torch.tensor([101, 102]) ++ ) ++ torch.testing.assert_close( ++ writes[0]["ctx_positions"], torch.tensor([3, 4, 7, 8, 9]) ++ ) ++ torch.testing.assert_close(writes[0]["ctx_cache_loc"], batch.out_cache_loc) ++ reference = F.rms_norm(F.linear(h.flatten(1), weight), (hidden_size,), eps=1e-6) ++ actual = writes[0]["ctx_hidden"] ++ if dtype == torch.float32: ++ torch.testing.assert_close(actual, reference, atol=2e-6, rtol=2e-5) ++ else: ++ relative_l2 = torch.linalg.vector_norm( ++ actual.float() - reference.float() ++ ) / torch.linalg.vector_norm(reference.float()) ++ self.assertLess(relative_l2.item(), 0.02) ++ ++ def test_pp8_projection_empty_stages_float32_and_bf16(self): ++ for dtype in (torch.float32, torch.bfloat16): ++ self._run_pipeline(dtype, [19, 37, 54, 66, 78, 90]) ++ ++ def test_pp8_boundary_capture(self): ++ self._run_pipeline(torch.float32, [11, 23, 47, 71, 83, 90]) ++ ++ def test_last_stage_without_capture_uses_incoming_context(self): ++ self._run_pipeline(torch.float32, [1, 2, 11, 23, 47, 71]) ++ ++ def test_missing_incoming_context_fails_before_kv_write(self): ++ with self.assertRaisesRegex(RuntimeError, "context missing"): ++ self._run_pipeline( ++ torch.float32, [19, 37, 54, 66, 78, 90], drop_incoming=True ++ ) ++ ++ def test_pd_input_builder_and_spec_dispatch(self): ++ class Spec: ++ is_eagle = lambda self: False ++ is_dspark = lambda self: False ++ is_dflash = lambda self: True ++ ++ calls = [] ++ relay = SimpleNamespace ++ make = lambda **kw: SimpleNamespace(**kw) ++ builder = load_function( ++ "speculative/dflash_disaggregation.py", ++ "build_dflash_family_disagg_draft_input", ++ globals_dict={"RelayPayload": relay, "make_draft_input_v2": make}, ++ ) ++ dispatch = load_function( ++ "speculative/spec_info.py", ++ "build_disagg_draft_input", ++ owner="SpeculativeAlgorithm", ++ ) ++ module = SimpleNamespace(build_dflash_family_disagg_draft_input=builder) ++ future = SimpleNamespace( ++ publish=lambda *args: calls.append(("publish", args)), ++ stash=lambda *args: calls.append(("stash", args)), ++ ) ++ with patch.dict( ++ sys.modules, {"sglang.srt.speculative.dflash_disaggregation": module} ++ ): ++ for overlap in (False, True): ++ calls.clear() ++ batch = SimpleNamespace( ++ enable_overlap=overlap, ++ seq_lens=torch.tensor([7, 9]), ++ req_pool_indices=torch.tensor([2, 4]), ++ ) ++ tokens = torch.tensor([101, 102]) ++ result = dispatch(Spec(), batch, tokens, future) ++ self.assertIs(result.bonus_tokens, tokens) ++ self.assertIs(result.new_seq_lens, batch.seq_lens) ++ self.assertEqual( ++ [c[0] for c in calls], ["publish", "stash"] if overlap else [] ++ ) ++ if overlap: ++ self.assertIs(calls[-1][1][1].bonus_tokens, tokens) ++ ++ def test_nonfinal_pool_is_minimal_without_mutating_target_config(self): ++ allocate = load_function( ++ "speculative/dflash_worker_v2.py", ++ "alloc_memory_pool", ++ owner="DFlashWorkerV2", ++ globals_dict={"replace": replace}, ++ ) ++ Config = make_dataclass( ++ "Config", ++ [ ++ "max_total_num_tokens", ++ "full_max_total_num_tokens", ++ "swa_max_total_num_tokens", ++ ], ++ ) ++ original = Config(8192, 8192, None) ++ calls = [] ++ worker = SimpleNamespace( ++ page_size=64, ++ _is_context_only_pp_prefill_rank=True, ++ use_compact_draft_cache=False, ++ _draft_worker=SimpleNamespace( ++ alloc_memory_pool=lambda **kw: calls.append(kw) ++ ), ++ ) ++ allocate(worker, original, "req_pool", "allocator") ++ self.assertEqual(original.max_total_num_tokens, 8192) ++ self.assertEqual(calls[0]["memory_pool_config"], Config(64, 64, None)) ++ self.assertEqual(calls[0]["req_to_token_pool"], "req_pool") ++ worker._is_context_only_pp_prefill_rank = False ++ allocate(worker, original, "req_pool", "allocator") ++ self.assertIs(calls[1]["memory_pool_config"], original) ++ ++ def test_only_prefill_skips_draft_graph_initialization(self): ++ init_graphs = load_function( ++ "speculative/dflash_worker_v2.py", ++ "init_cuda_graphs", ++ owner="DFlashWorkerV2", ++ globals_dict={ ++ "get_exec": lambda: SimpleNamespace( ++ graph=SimpleNamespace( ++ cuda_graph_config=SimpleNamespace( ++ decode=SimpleNamespace(backend="disabled") ++ ) ++ ) ++ ), ++ "Backend": SimpleNamespace(DISABLED="disabled"), ++ "is_cuda": lambda: False, ++ }, ++ ) ++ calls = [] ++ worker = SimpleNamespace( ++ _is_pd_prefill=True, ++ _draft_worker=SimpleNamespace( ++ init_cuda_graphs=lambda **kw: calls.append(kw) ++ ), ++ ) ++ init_graphs(worker) ++ self.assertEqual(calls, []) ++ worker._is_pd_prefill = False ++ init_graphs(worker) ++ self.assertEqual(calls, [{"capture_decode_cuda_graph": False}]) ++ ++ def test_pp_proxy_is_consumed_once_even_on_forward_failure(self): ++ forward = load_function( ++ "speculative/dflash_worker_v2.py", ++ "forward_batch_generation", ++ owner="DFlashWorkerV2", ++ ) ++ proxy = object() ++ calls = [] ++ ++ def fail(batch, publish, tensors): ++ calls.append(tensors) ++ raise RuntimeError("target failed") ++ ++ worker = SimpleNamespace( ++ _next_pp_proxy_tensors=proxy, ++ ps=SimpleNamespace(pp_size=8), ++ _validate_phase1_sampling_support=lambda batch: None, ++ _forward_pp_prefill=fail, ++ ) ++ batch = SimpleNamespace(forward_mode=SimpleNamespace(is_extend=lambda: True)) ++ with self.assertRaisesRegex(RuntimeError, "target failed"): ++ forward(worker, batch) ++ self.assertEqual(calls, [proxy]) ++ self.assertIsNone(worker._next_pp_proxy_tensors) ++ ++ ++if __name__ == "__main__": ++ unittest.main(verbosity=2) +diff --git a/test/registered/disaggregation/test_flashinfer_kimi_merge.py b/test/registered/disaggregation/test_flashinfer_kimi_merge.py +new file mode 100644 +--- /dev/null ++++ b/test/registered/disaggregation/test_flashinfer_kimi_merge.py +@@ -0,0 +1,254 @@ ++"""CPU regression for the existing SiTU adapter rebased onto the PP snapshot.""" ++ ++import inspect ++import sys ++from contextlib import nullcontext ++from enum import IntEnum ++from types import SimpleNamespace ++import unittest ++from unittest.mock import patch ++ ++import torch ++from torch.nn import Parameter ++ ++from test_dflash_pp_context import load_function ++ ++ ++class TestFlashInferKimiMerge(unittest.TestCase): ++ def test_activation_selection_accepts_situ_and_legacy_enum(self): ++ class LegacyActivation(IntEnum): ++ Swiglu = 0 ++ Geglu = 1 ++ Relu2 = 2 ++ Identity = 3 ++ ++ class SituActivation(IntEnum): ++ Swiglu = 0 ++ Geglu = 1 ++ Relu2 = 2 ++ Identity = 3 ++ Situ = 10 ++ ++ for enum, kind, expected in ( ++ (LegacyActivation, "silu", LegacyActivation.Swiglu), ++ (SituActivation, "silu", SituActivation.Swiglu), ++ (SituActivation, "situ", SituActivation.Situ), ++ ): ++ with self.subTest(enum=enum.__name__, kind=kind): ++ select = load_function( ++ "layers/moe/moe_runner/flashinfer_cutlass.py", ++ "_activation_type", ++ globals_dict={ ++ "_flashinfer_cutlass_fused_moe": lambda: (None, enum) ++ }, ++ ) ++ with patch.dict( ++ sys.modules, ++ { ++ "sglang.srt.layers.moe.moe_runner.flashinfer_trtllm": SimpleNamespace( ++ get_activation_type=lambda *a, **kw: int(expected) ++ ) ++ }, ++ ): ++ self.assertEqual( ++ select(SimpleNamespace(activation=kind, is_gated=True)), ++ expected, ++ ) ++ ++ def test_weight_layout_and_situ_parameters(self): ++ prepare = load_function( ++ "layers/quantization/mxfp4.py", ++ "_process_weights_for_sm120_cutlass", ++ owner="Mxfp4MoEMethod", ++ globals_dict={"Parameter": Parameter}, ++ ) ++ for interleaved in (False, True): ++ for activation in ("situ", "silu"): ++ with self.subTest(interleaved=interleaved, activation=activation): ++ experts, hidden, intermediate, padded_n = 2, 64, 32, 128 ++ gate = torch.full( ++ (experts, intermediate, hidden // 2), 17, dtype=torch.uint8 ++ ) ++ up = torch.full_like(gate, 29) ++ packed = ( ++ torch.stack((gate, up), dim=2).flatten(1, 2) ++ if interleaved ++ else torch.cat((gate, up), dim=1) ++ ) ++ params = lambda t: Parameter(t, requires_grad=False) ++ layer = SimpleNamespace( ++ num_local_experts=experts, ++ w13_weight=params(packed), ++ w13_weight_scale=params(packed[:, :, : hidden // 32].clone()), ++ w13_weight_bias=params(packed[:, :, 0].to(torch.bfloat16)), ++ w2_weight=params( ++ torch.ones( ++ experts, hidden, intermediate // 2, dtype=torch.uint8 ++ ) ++ ), ++ w2_weight_scale=params( ++ torch.ones( ++ experts, hidden, intermediate // 32, dtype=torch.uint8 ++ ) ++ ), ++ w2_weight_bias=params( ++ torch.zeros(experts, hidden, dtype=torch.bfloat16) ++ ), ++ ) ++ method = SimpleNamespace( ++ _padded_hidden=128, ++ _padded_intermediate=padded_n, ++ moe_runner_config=SimpleNamespace( ++ gate_up_interleaved=interleaved, ++ activation=activation, ++ gemm1_alpha=None, ++ gemm1_beta=None, ++ gemm1_clamp_limit=None, ++ ), ++ ) ++ # Swizzle is a FlashInfer GPU operation; this test checks layout before it. ++ with patch.dict( ++ sys.modules, ++ { ++ "flashinfer": SimpleNamespace( ++ block_scale_interleave=lambda t: t ++ ) ++ }, ++ ): ++ prepare(method, layer) ++ torch.testing.assert_close( ++ layer.w13_weight[:, :intermediate, : hidden // 2], up ++ ) ++ torch.testing.assert_close( ++ layer.w13_weight[ ++ :, padded_n : padded_n + intermediate, : hidden // 2 ++ ], ++ gate, ++ ) ++ self.assertEqual( ++ layer.w13_weight[:, intermediate:padded_n] ++ .count_nonzero() ++ .item(), ++ 0, ++ ) ++ self.assertEqual( ++ layer.w13_weight[:, :, hidden // 2 :].count_nonzero().item(), 0 ++ ) ++ if activation == "situ": ++ torch.testing.assert_close( ++ layer.situ_beta, torch.full((experts,), 4.0) ++ ) ++ torch.testing.assert_close( ++ layer.situ_linear_beta, torch.full((experts,), 25.0) ++ ) ++ self.assertIsNone(layer.swiglu_alpha) ++ else: ++ self.assertIsNone(layer.situ_beta) ++ torch.testing.assert_close( ++ layer.swiglu_alpha, torch.full((experts,), 1.702) ++ ) ++ ++ def test_runner_preserves_swiglu_step_and_old_api(self): ++ calls, inputs = [], [] ++ activation = SimpleNamespace(Swiglu="swiglu", SwigluStep="step", Situ="situ") ++ fused = lambda **kw: calls.append(kw) ++ runner = load_function( ++ "layers/moe/moe_runner/flashinfer_cutlass.py", ++ "fused_experts_none_to_flashinfer_mxfp4", ++ globals_dict={ ++ "register_fused_func": lambda *a: lambda f: f, ++ "FlashInferCutlassMxfp4MoeQuantInfo": SimpleNamespace, ++ "_flashinfer_cutlass_fused_moe": lambda: (fused, activation), ++ "flashinfer_cutlass_supports_situ": lambda: True, ++ "use_symmetric_memory": lambda *a, **k: nullcontext(), ++ "get_tp_group": lambda: None, ++ "is_allocation_symmetric": lambda: False, ++ "next_power_of_2": lambda v: 1 << (v - 1).bit_length(), ++ "envs": SimpleNamespace( ++ SGLANG_FLASHINFER_MOE_FUSED_FINALIZE=SimpleNamespace( ++ get=lambda: True ++ ) ++ ), ++ }, ++ ) ++ ++ def quantize(x, **kw): ++ inputs.append(x.is_contiguous()) ++ return x, torch.ones(1) ++ ++ modules = { ++ "flashinfer": SimpleNamespace(mxfp8_quantize=quantize), ++ "sglang.srt.layers.moe.token_dispatcher.standard": SimpleNamespace( ++ StandardCombineInput=SimpleNamespace ++ ), ++ "sglang.srt.layers.moe.topk": SimpleNamespace( ++ TopKOutputChecker=SimpleNamespace(format_is_bypassed=lambda x: False) ++ ), ++ } ++ quant_info = SimpleNamespace( ++ padded_hidden=None, ++ mxfp4_weight_global_scale=torch.ones(1), ++ w13_weight=torch.zeros(1, 64, 32, dtype=torch.uint8), ++ w2_weight=torch.zeros(1, 64, 32, dtype=torch.uint8), ++ w13_weight_scale=torch.zeros(1, 64, 4, dtype=torch.uint8), ++ w2_weight_scale=torch.zeros(1, 64, 4, dtype=torch.uint8), ++ w13_bias=None, ++ w2_bias=None, ++ swiglu_alpha=None, ++ swiglu_beta=None, ++ swiglu_limit=None, ++ situ_beta=torch.tensor([4.0]), ++ situ_linear_beta=torch.tensor([25.0]), ++ moe_tp_size=4, ++ moe_tp_rank=0, ++ moe_ep_size=4, ++ moe_ep_rank=0, ++ use_swiglu_step=False, ++ ) ++ dispatch = SimpleNamespace( ++ hidden_states=torch.ones(4, 128)[:, ::2], ++ topk_output=SimpleNamespace( ++ topk_ids=torch.zeros(4, 1, dtype=torch.int64), ++ topk_weights=torch.ones(4, 1), ++ ), ++ ) ++ with patch.dict(sys.modules, modules): ++ for kind, step, expected in ( ++ ("situ", False, "situ"), ++ ("silu", True, "step"), ++ ("silu", False, "swiglu"), ++ ): ++ quant_info.use_swiglu_step = step ++ runner(dispatch, quant_info, SimpleNamespace(activation=kind)) ++ self.assertEqual(calls[-1]["activation_type"], expected) ++ self.assertEqual("situ_beta" in calls[-1], kind == "situ") ++ self.assertEqual(inputs, [True, True, True]) ++ ++ def test_capability_requires_public_parameters(self): ++ def legacy(input): ++ pass ++ ++ def public(input, *, situ_beta=None, situ_linear_beta=None): ++ pass ++ ++ for func, has_enum, expected in ( ++ (legacy, True, False), ++ (public, False, False), ++ (public, True, True), ++ ): ++ check = load_function( ++ "layers/moe/moe_runner/flashinfer_cutlass.py", ++ "flashinfer_cutlass_supports_situ", ++ globals_dict={ ++ "inspect": inspect, ++ "_flashinfer_cutlass_fused_moe": lambda: ( ++ func, ++ SimpleNamespace(Situ=10) if has_enum else SimpleNamespace(), ++ ), ++ }, ++ ) ++ self.assertEqual(check(), expected) ++ ++ ++if __name__ == "__main__": ++ unittest.main(verbosity=2) +diff --git a/test/registered/disaggregation/test_mixed_kv_entry_layout.py b/test/registered/disaggregation/test_mixed_kv_entry_layout.py +new file mode 100644 +--- /dev/null ++++ b/test/registered/disaggregation/test_mixed_kv_entry_layout.py +@@ -0,0 +1,302 @@ ++"""CPU regression for Kimi MLA target + GQA draft PD byte transfer plans.""" ++ ++import importlib.util ++import ast ++from collections import deque ++import json ++import os ++from pathlib import Path ++import struct ++import sys ++import unittest ++from types import SimpleNamespace ++ ++ROOT = Path(os.environ.get("SGLANG_SOURCE_ROOT", Path(__file__).resolve().parents[3])) ++MODULE = ROOT / "python/sglang/srt/disaggregation/common/kv_entry_layout.py" ++spec = importlib.util.spec_from_file_location("kv_entry_layout_under_test", MODULE) ++module = importlib.util.module_from_spec(spec) ++sys.modules[spec.name] = module ++spec.loader.exec_module(module) ++KVEntryLayout = module.KVEntryLayout ++plan_kv_entry_transfer = module.plan_kv_entry_transfer ++ ++ ++def upstream_functions(*items): ++ functions = [] ++ for path, name in items: ++ tree = ast.parse((ROOT / path).read_text()) ++ functions.append( ++ next( ++ node ++ for node in ast.walk(tree) ++ if isinstance(node, ast.FunctionDef) and node.name == name ++ ) ++ ) ++ namespace = dict( ++ KVEntryLayout=KVEntryLayout, ++ plan_kv_entry_transfer=plan_kv_entry_transfer, ++ deque=deque, ++ ) ++ tree = ast.Module( ++ body=[ ++ ast.ImportFrom( ++ module="__future__", names=[ast.alias(name="annotations")], level=0 ++ ), ++ *functions, ++ ], ++ type_ignores=[], ++ ) ++ exec( ++ compile(ast.fix_missing_locations(tree), "", "exec"), ++ namespace, ++ ) ++ return namespace ++ ++ ++def layout(tp, heads=8, page=4, dim=2, kind="nhd"): ++ stride = page * max(1, heads // tp) * dim * 2 ++ return KVEntryLayout( ++ kind, page, stride, stride * 6, "torch.bfloat16", 2, heads, dim ++ ) ++ ++ ++def source_ranks(p_tp, d_tp, d_rank): ++ if p_tp <= d_tp: ++ return [d_rank // (d_tp // p_tp)] ++ count = p_tp // d_tp ++ return range(d_rank * count, (d_rank + 1) * count) ++ ++ ++class TestMixedKVEntryLayout(unittest.TestCase): ++ def test_mooncake_mixed_entries_match_layers_before_transfer(self): ++ functions = upstream_functions( ++ ("python/sglang/srt/disaggregation/utils.py", "build_transfer_entry_pairs"), ++ ( ++ "python/sglang/srt/disaggregation/mooncake/conn.py", ++ "_send_mixed_kv_entries", ++ ), ++ ) ++ flat = KVEntryLayout("flat", 64, 73728, 73728 * 6, "torch.bfloat16", 2) ++ src, dst = layout(4, page=64, dim=128), layout(32, page=64, dim=128) ++ records = [flat.to_wire(), src.to_wire(), src.to_wire()] ++ calls = [] ++ manager = SimpleNamespace( ++ attn_tp_size=4, ++ kv_args=SimpleNamespace( ++ kv_entry_layouts=records, ++ kv_data_ptrs=[1000000, 2000000, 3000000], ++ kv_item_lens=[r["item_len"] for r in records], ++ kv_data_lens=[r["buffer_nbytes"] for r in records], ++ kv_layer_ids=[1, 93, 93], ++ engine_rank=0, ++ ), ++ _transfer_data=lambda session, blocks: calls.append((session, blocks)) or 0, ++ ) ++ destination = [ ++ flat.to_wire(), ++ flat.to_wire(), ++ dst.to_wire(), ++ flat.to_wire(), ++ dst.to_wire(), ++ ] ++ args = [ ++ manager, ++ "test", ++ SimpleNamespace(tolist=lambda: [1]), ++ [4000000, 5000000, 6000000, 7000000, 8000000], ++ SimpleNamespace(tolist=lambda: [2]), ++ [7, 1, 93, 17, 93], ++ 32, ++ 0, ++ destination, ++ ] ++ functions["_send_mixed_kv_entries"](*args) ++ self.assertEqual(len(calls), 1) ++ blocks = calls[0][1] ++ self.assertEqual(blocks[0], (1000000 + 73728, 5000000 + 2 * 73728, 73728)) ++ self.assertEqual(blocks[1][1], 6000000 + 2 * 16384) ++ self.assertEqual(blocks[65][1], 8000000 + 2 * 16384) ++ self.assertEqual(sum(block[2] for block in blocks), 73728 + 2 * 16384) ++ calls.clear() ++ destination[-1]["item_len"] += 1 ++ with self.assertRaises(ValueError): ++ functions["_send_mixed_kv_entries"](*args) ++ self.assertEqual(calls, [], "Invalid later entry must not allow earlier writes") ++ args[-1] = None ++ with self.assertRaises(ValueError): ++ functions["_send_mixed_kv_entries"](*args) ++ self.assertEqual(calls, []) ++ ++ def test_all_gqa_shards_and_replicas_byte_exact(self): ++ for heads in (1, 2, 8, 16): ++ for p_tp in (1, 2, 4, 8, 16, 32): ++ for d_tp in (1, 2, 4, 8, 16, 32): ++ with self.subTest(heads=heads, p_tp=p_tp, d_tp=d_tp): ++ src, dst = layout(p_tp, heads), layout(d_tp, heads) ++ for d_rank in range(d_tp): ++ output = bytearray([255]) * dst.buffer_nbytes ++ written = set() ++ dst_head_start = ( ++ d_rank * max(1, heads // d_tp) // max(1, d_tp // heads) ++ ) ++ for p_rank in source_ranks(p_tp, d_tp, d_rank): ++ input_buffer = bytearray(src.buffer_nbytes) ++ src_head_start = ( ++ p_rank ++ * max(1, heads // p_tp) ++ // max(1, p_tp // heads) ++ ) ++ for page in (1, 3): ++ for token in range(src.page_size): ++ for head in range(max(1, heads // p_tp)): ++ value = ( ++ 100 * page ++ + 20 * token ++ + src_head_start ++ + head ++ ) ++ offset = ( ++ page * src.item_len ++ + (token * max(1, heads // p_tp) + head) ++ * 4 ++ ) ++ input_buffer[offset : offset + 4] = ( ++ struct.pack("=13.0 brand=unknown,driver>=535,driver<536 brand=grid,driver>=535,driver<536 brand=tesla,driver>=535,driver<536 brand=nvidia,driver>=535,driver<536 brand=quadro,driver>=535,driver<536 brand=quadrortx,driver>=535,driver<536 brand=nvidiartx,driver>=535,driver<536 brand=vapps,driver>=535,driver<536 brand=vpc,driver>=535,driver<536 brand=vcs,driver>=535,driver<536 brand=vws,driver>=535,driver<536 brand=cloudgaming,driver>=535,driver<536 brand=unknown,driver>=550,driver<551 brand=grid,driver>=550,driver<551 brand=tesla,driver>=550,driver<551 brand=nvidia,driver>=550,driver<551 brand=quadro,driver>=550,driver<551 brand=quadrortx,driver>=550,driver<551 brand=nvidiartx,driver>=550,driver<551 brand=vapps,driver>=550,driver<551 brand=vpc,driver>=550,driver<551 brand=vcs,driver>=550,driver<551 brand=vws,driver>=550,driver<551 brand=cloudgaming,driver>=550,driver<551 brand=unknown,driver>=565,driver<566 brand=grid,driver>=565,driver<566 brand=tesla,driver>=565,driver<566 brand=nvidia,driver>=565,driver<566 brand=quadro,driver>=565,driver<566 brand=quadrortx,driver>=565,driver<566 brand=nvidiartx,driver>=565,driver<566 brand=vapps,driver>=565,driver<566 brand=vpc,driver>=565,driver<566 brand=vcs,driver>=565,driver<566 brand=vws,driver>=565,driver<566 brand=cloudgaming,driver>=565,driver<566 brand=unknown,driver>=570,driver<571 brand=grid,driver>=570,driver<571 brand=tesla,driver>=570,driver<571 brand=nvidia,driver>=570,driver<571 brand=quadro,driver>=570,driver<571 brand=quadrortx,driver>=570,driver<571 brand=nvidiartx,driver>=570,driver<571 brand=vapps,driver>=570,driver<571 brand=vpc,driver>=570,driver<571 brand=vcs,driver>=570,driver<571 brand=vws,driver>=570,driver<571 brand=cloudgaming,driver>=570,driver<571 brand=unknown,driver>=575,driver<576 brand=grid,driver>=575,driver<576 brand=tesla,driver>=575,driver<576 brand=nvidia,driver>=575,driver<576 brand=quadro,driver>=575,driver<576 brand=quadrortx,driver>=575,driver<576 brand=nvidiartx,driver>=575,driver<576 brand=vapps,driver>=575,driver<576 brand=vpc,driver>=575,driver<576 brand=vcs,driver>=575,driver<576 brand=vws,driver>=575,driver<576 brand=cloudgaming,driver>=575,driver<576", + "NV_CUDA_CUDART_VERSION=13.0.96-1", + "CUDA_VERSION=13.0.3", + "LD_LIBRARY_PATH=/usr/local/nvidia/lib:/usr/local/nvidia/lib64:/usr/local/cuda/lib64:/usr/local/nvidia/lib:/usr/local/nvidia/lib64", + "NVIDIA_VISIBLE_DEVICES=all", + "NVIDIA_DRIVER_CAPABILITIES=compute,utility", + "NV_CUDA_LIB_VERSION=13.0.3-1", + "NV_NVTX_VERSION=13.0.85-1", + "NV_LIBNPP_VERSION=13.0.1.2-1", + "NV_LIBNPP_PACKAGE=libnpp-13-0=13.0.1.2-1", + "NV_LIBCUSPARSE_VERSION=12.6.3.3-1", + "NV_LIBCUBLAS_PACKAGE_NAME=libcublas-13-0", + "NV_LIBCUBLAS_VERSION=13.1.1.3-1", + "NV_LIBCUBLAS_PACKAGE=libcublas-13-0=13.1.1.3-1", + "NV_LIBNCCL_PACKAGE_NAME=libnccl2", + "NV_LIBNCCL_PACKAGE_VERSION=2.28.3-1", + "NCCL_VERSION=2.28.3-1", + "NV_LIBNCCL_PACKAGE=libnccl2=2.28.3-1+cuda13.0", + "NVIDIA_PRODUCT_NAME=CUDA", + "NV_CUDA_CUDART_DEV_VERSION=13.0.96-1", + "NV_NVML_DEV_VERSION=13.0.87-1", + "NV_LIBCUSPARSE_DEV_VERSION=12.6.3.3-1", + "NV_LIBNPP_DEV_VERSION=13.0.1.2-1", + "NV_LIBNPP_DEV_PACKAGE=libnpp-dev-13-0=13.0.1.2-1", + "NV_LIBCUBLAS_DEV_VERSION=13.1.1.3-1", + "NV_LIBCUBLAS_DEV_PACKAGE_NAME=libcublas-dev-13-0", + "NV_LIBCUBLAS_DEV_PACKAGE=libcublas-dev-13-0=13.1.1.3-1", + "NV_CUDA_NSIGHT_COMPUTE_VERSION=13.0.3-1", + "NV_CUDA_NSIGHT_COMPUTE_DEV_PACKAGE=cuda-nsight-compute-13-0=13.0.3-1", + "NV_LIBNCCL_DEV_PACKAGE_NAME=libnccl-dev", + "NV_LIBNCCL_DEV_PACKAGE_VERSION=2.28.3-1", + "NV_LIBNCCL_DEV_PACKAGE=libnccl-dev=2.28.3-1+cuda13.0", + "LIBRARY_PATH=/usr/local/cuda/lib64/stubs", + "NV_CUDNN_VERSION=9.14.0.64-1", + "NV_CUDNN_PACKAGE_NAME=libcudnn9-cuda-13", + "NV_CUDNN_PACKAGE=libcudnn9-cuda-13=9.14.0.64-1", + "NV_CUDNN_PACKAGE_DEV=libcudnn9-dev-cuda-13=9.14.0.64-1", + "NV_CUDNN_PACKAGE_DEV_HEADERS=libcudnn9-headers-cuda-13=9.14.0.64-1", + "DEBIAN_FRONTEND=noninteractive", + "CUDA_HOME=/usr/local/cuda", + "GDRCOPY_HOME=/usr/src/gdrdrv-2.5.1/", + "FLASHINFER_VERSION=0.6.17", + "LANG=en_US.UTF-8", + "LANGUAGE=en_US:en", + "LC_ALL=en_US.UTF-8", + "SGLANG_BUILD_COMMIT=daf631719690e18d13f67a20eb513fd48c712327", + "SGLANG_BUILD_URL=https://github.com/sgl-project/sglang/actions/runs/33137164725", + "SGLANG_IMAGE_TAG=lmsysorg/sglang:nightly-dev-20260828-daf63171", + "PYTHONPATH=/opt/kimi-dflash/python", + "PYTHONDONTWRITEBYTECODE=1", + "FLASHINFER_DISABLE_VERSION_CHECK=", + "SGLANG_SOURCE_ROOT=/opt/kimi-dflash" + ], + "Entrypoint": [ + "python3", + "-m", + "sglang.launch_server" + ], + "WorkingDir": "/sgl-workspace/sglang", + "Labels": { + "ai.sglang.build.commit": "daf631719690e18d13f67a20eb513fd48c712327", + "ai.sglang.build.url": "https://github.com/sgl-project/sglang/actions/runs/33137164725", + "ai.sglang.image.tag": "lmsysorg/sglang:nightly-dev-20260828-daf63171", + "com.nvidia.cudnn.version": "9.14.0.64-1", + "maintainer": "NVIDIA CORPORATION ", + "org.opencontainers.image.revision": "daf631719690e18d13f67a20eb513fd48c712327", + "org.opencontainers.image.source": "https://github.com/sgl-project/sglang", + "org.opencontainers.image.url": "https://github.com/sgl-project/sglang/actions/runs/33137164725", + "org.opencontainers.image.version": "lmsysorg/sglang:nightly-dev-20260828-daf63171" + } + }, + "Architecture": "amd64", + "Os": "linux", + "Size": 14967226422, + "RootFS": { + "Type": "layers", + "Layers": [ + "sha256:5e732af9e7c568c9b41ecabc76ac93a58471934d5155070dfd64a9567667fd9d", + "sha256:f5b49b5386a3c18aca8dccc85784c57d64b182f621a632374c44daeb7822e9ca", + "sha256:2950e5ad5b560d8a4561ea596a8589f9be34eb485ad04ae521438c978d98a212", + "sha256:f56f6d7fecabd59c6ae0104489b7e42e0b1c538ba1246f0019b099a12cca8cd6", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:6756a6579040c1541966a48c3e4d3bac9500de08ee027795439141ee0ba0eecd", + "sha256:fdb5c9bcf16b58539209a7e6a469bea3a5b5434d2c4e9d4719968e22fe0ef1d7", + "sha256:c6340b3836ce7fea1a99cadff25d755971a2091ecc7815f7761cce83624961ca", + "sha256:6d6b5c3ac20f833a318cc28f64194bc3d50371b023199e420c641dd5d3264d5e", + "sha256:2dfea3ee7ecb9bf9eaf210231c0d8fc84706edea83c584300f86aef7eb0dfe1d", + "sha256:988d8769daf31ed246461f52c452d8b11248ba47e5b06fe8210915025327d393", + "sha256:cd905001042d6fb31a08ed1272f4b36b9bf1814c0e111355cc8725277cd4bfb7", + "sha256:ae8a69caad45e0cf5977673d59e9b3d444943e89e8eb503a7798f093c65162e8", + "sha256:3f83440160ca573ae240f1525e4e9d6c7ee9d055dfc7be61ac7af84ee1642313", + "sha256:72f31487cece23c9c9f79d1336cf0887726864ba145d967da1fb5d775e533d93", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:3cbf470cb28d2108aba6f0154ecc8a457dc588af7161b915474987bf7973eef2", + "sha256:a679dc7078c9cfa517d01828412d0a591707211483cb74e6313ae603a006f239", + "sha256:8d074d326c9d61118e25e27c432dac303839d3529df5b64a165438848cc2f950", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:7509e379739e5c54072504957c5e22be0cbfb8d9ca12d3b66450c2a725c1f2bb", + "sha256:15785f6df9bee4007ab28a370d83c300ac85ab056869609d1607f6e3c88a6b2c", + "sha256:ada9078922901f5e11dadf528356e717f1619f305c2e32bd8569f92ad76a08b0", + "sha256:a8ee45b1f77587a9ad6bfc8c28d98e2a928f02442251442b5e38bd2a89c480bb", + "sha256:4efe4144ff056f4e35ab86371e2b6905e238d4752e41c0a793821a6469864e3f", + "sha256:10007f71c41cfc3a88da7b29b4dd9edc5d417515b7f63126527bf0b6c1f546fe", + "sha256:c648315cc3944840ae4661ad11397d0acde62110459f7b3b178d1e6f0c3d4079", + "sha256:a1cf5e134ae7bbba0d8002a8deecda067a440170fd9c90318b3be18c25f15c61", + "sha256:b064bfce42d6545b7c7e4c191abc9997e2cd43b96efcf14e56cc342b12201011", + "sha256:f985ef7df34872310bb3c00367c23a55c12def9692698bd2820b9c377e83013b", + "sha256:a1a57562b488e2fdc164abe1045170ac864ddc9417995966034f839058813eac", + "sha256:2258726d23a32a4d5c468c0b1bbe985a8537c47cf34db06ebece8b4d4daa74b1", + "sha256:faa84dbafd41fa49b0b648cb6557e7343849b389d93fdf5556f6e2decaf550a1", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:e8706d16624fb04f6a2dfeb067da4254f54a550decc54267915e26cccb4c0fed", + "sha256:d324fa1e96e663f204caea0ac5e1be0cf5b3a8a67a3bc931a484f134dcfa6a6d", + "sha256:853adfa25cad926899088fd81ace2a3eaecb0b64815ada56c016dc20e17c02e4", + "sha256:81e9c64c1f5b93dfa25081071a7d1744344fbf41ae9233270bf4ea6db59ddd33", + "sha256:c4274c717b4e908fc996bfbee932682d08233f52424ab7954aff5f4014931500", + "sha256:8876e61557f92352e0fa2e79f835d2a2617cd44bdd5cfe72b4ee5deb66645da2", + "sha256:cd60d9eff385b6ff80d2e285a57914c5dc921861ea69a3ce7d7fe3d0bb2c05d0", + "sha256:3cf3641ca559478f4b9316177739b72e159dafd7d0c163d4d08a2b999a034c27", + "sha256:28f200a4f493abe3b5531e85ad3423f72bde68df2b3b61a2417e0c7d9b6177a8", + "sha256:9029cbfb4b4e3bdef2574ec2af6c9240208af42ad47326c038f23acb802f843d", + "sha256:81ce6120dddf6397bdc20e75911d99c8ea891891ecd84a585f28ebe26134061b", + "sha256:f9c2555ebfec8d41e60e4323a8a46a2e7d1b17e8f8ab5a823e9513387b4eac7a", + "sha256:ffabea8dc4f4987995b8a0ff313ac920cd77e534103ad19e800de02d00c210fb", + "sha256:372d209365d371ad6864e432126bec3403add2a274eb5c4d72c09a891fe6c773", + "sha256:74f3ccaf8aa1c45cf376811caaa8cd70da8d349ff8b7c34624570a50c8813c1d", + "sha256:f3cb350e69a5b6c4e176f428401faaadb539c1afc8e891baf1a077fdc6f7df9f", + "sha256:cae7c313be69dccbed4a034062338e1e4e6d92090e9c4c7a38308c4ff261231d", + "sha256:e9ad469189345e9bc779fb888594205c7eb5b997e50d4cc7a33559e39de48a86", + "sha256:817e9f83523ce8b5289b76165dcadf14a21eb45cda7fcd8dc3ac7f7480757bda", + "sha256:4c6d2f48bbc68fe34f554017e18166ad2628d203d435414af7facdba57ba6bc1", + "sha256:82381ac149748e9cd4bfecedd26c6808e4e4e8ff251dfa7c36c6eec115e01408", + "sha256:df5778a62653618e7d7a55efbb2002b8eea2a455da4fab1aa57714fac72e8354", + "sha256:de669107068ddbc9a7ae7bdc4c83e11b7fd17e2a75f93f5384de8d3bcd937073", + "sha256:b185c4deee75530ce0ad0088a587b33c48f6a217d1f13f894e73eea735d76d5a", + "sha256:1fa7159b375ca9c6fbf5b1f61b7d9aaf88583c53bfe52933cc3799a7de2e2b7d", + "sha256:57611659260d363d32f15cee45f01a836a5df931ccd21546e03e965690afd697", + "sha256:e5c5a0f5e39c9e1de42eeea1e198142925fa4a533a956f316fa785c69b6c9a56", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:bd1946dd19d617f0a8a0ebf2a88e7cdab07a77f9fa070a35a605396dbdcaaed2", + "sha256:647babc16cef2877883c9a4276e9b800011130b4a921a1c58dc2d90db912221d", + "sha256:d03dfaade251e784281d5e6cc3d35cecf809e0994e67c9c170d9e0c869df8408", + "sha256:4ba4ee5654a71c370d25e571934ad0d1a03211fcf36a94ae7318d13fcb5b671d", + "sha256:c39e8c43807472b24f161f29d6feb422dfb949fe7100d69f1ef7eadedc14c0fc", + "sha256:bc97787ac190f68c2cc3aa5e08cd2862ea83ddd224796401bfc2f084990c308d", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef", + "sha256:5a8a3022261ca98b6ecb6309d3faa376daa5d819c8a1c3bbfeba389ffce6272b", + "sha256:e6fda71a3f6320ed1295de878c8c1367cef2f616ba1d84cd371f895f0e47a2ba", + "sha256:ba942b61f6a46ad0c35de1c7c63ff5ecd00befac92b56d8cc56146af616f69c4", + "sha256:aa65adc2c641830ccc13c37bb23c7658aeecea9c495d858a52d53141c34a8969", + "sha256:7c448cf35d32b08e4e9c1faba519b3bed1d81043ada7bb1f9834968fc2f06786", + "sha256:311055521e8a3fcc512490f0ff1f602db67dd536a01b62c984a97070fe8f7911", + "sha256:de65b09749fec355f64308704feb09b837cfb4ea9bc6e553c9261c73275b7cb8", + "sha256:5f70bf18a086007016e948b04aed3b82103a36bea41755b6cddfaf10ace3c6ef" + ] + }, + "Metadata": { + "LastTagTime": "2026-08-31T08:44:45.57613455Z" + }, + "Descriptor": { + "mediaType": "application/vnd.oci.image.manifest.v1+json", + "digest": "sha256:281ccb2666a38acd539e6fbb5d55d3682eb2af9ac9089f67e9bf92a8ddd822eb", + "size": 15020, + "platform": { + "architecture": "amd64", + "os": "linux" + } + } + } +] diff --git a/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_delta_metadata.json b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_delta_metadata.json new file mode 100644 index 0000000..0ba5c78 --- /dev/null +++ b/experiments/pro6000/kimi3_pro6000_pd_dflash_validation/results/pwarm1_delta_metadata.json @@ -0,0 +1,93 @@ +{ + "required_base_image": "sha256:6f998f448dd3545423e2bd73000c8bce991b4dbc9be66aff329b978b303c5987", + "image": "sha256:8ab5eee9902556bcba2a2b439a9fa1b14d93e0554503754fb7ea74dad6c3cf79", + "archive": "sglang-kimi-pp-dflash-8ab5eee-pwarm1.delta.oci.tar", + "bytes": 50759680, + "sha256": "8920c7a66a903a0653cd0f1cc087f626d241557bd7b63214db044eca89b768d5", + "included": [ + "blobs", + "blobs/sha256", + "blobs/sha256/108c859dc864e31043ddf1c3fa387c8397d6495939e11227e0a162fbbdba2023", + "blobs/sha256/281ccb2666a38acd539e6fbb5d55d3682eb2af9ac9089f67e9bf92a8ddd822eb", + "blobs/sha256/2ec8cdb54d91b2a0c9bf3aaf8855ab9319c5a07776cca449d6a5b26c9f721d15", + "blobs/sha256/38fca1f6908e11bd0945eff79c3f0fe8175c812bd88d2a3e0d9cfb20d81c1d83", + "blobs/sha256/3d0e896ac992404d479d6a73d4cc65bcb8506bebb982a415d90d15aa46503795", + "blobs/sha256/5e60539501edb70ef0ccd153dc5ba783117d1e916379fd44e116f5e5fd1b9d37", + "blobs/sha256/605208809bbedf1c51355c52df881abfd3ae8dbbfef94567e65b837a876dd319", + "blobs/sha256/8ab5eee9902556bcba2a2b439a9fa1b14d93e0554503754fb7ea74dad6c3cf79", + "blobs/sha256/9446f6f72966bb53dcea092aec81b887153fc8e024da73d3073d390c2e0938a8", + "blobs/sha256/a2d38d45b3b09b273eaa72ebcb7e1c30dc568fac72008492cdf53fd4f4e69f9e", + "blobs/sha256/d30d68757014b973c2edd5d061cc3fd55660a4c68cc3d0af9050fcb7e64c4a78", + "index.json", + "manifest.json", + "oci-layout" + ], + "preexisting_blobs": [ + "blobs/sha256/04ebfffe153ff264a394da45dd64c3e830c5f16c2e285959f72d49982b8330b2", + "blobs/sha256/088360b17cced9ae058a1342b9ab050b50bc7cfb72dfdcf7d04233d0d5a140b6", + "blobs/sha256/2154f8da8e6a860abe32529e59350d6da38d60205e956686560a68878f0f4949", + "blobs/sha256/216bf0c7a62ce24763d96e3f79789c7091108aa1ed05f5311c21dd887b054605", + "blobs/sha256/21e0ed254c8bac8e7e21e80b250a5dc4a92211bbc2f658f9f81ef93c7c623f03", + "blobs/sha256/25367c7aad83c96fda41bf8f039bb2fe57d9b2a475b3d5d7fdb5945278935ea4", + "blobs/sha256/257c5743ada16f22123bd0b28e4a8cd4280a3f72f2c75058a04b7117a9507f31", + "blobs/sha256/2628d54d17a882758e620e3a8b07e9181ff253dce5946bee42d77c8f34065d6d", + "blobs/sha256/2f59ed97b933ab02bc8c208f9bfa5cbd2f2b203c1c972be713196987c61f8e68", + "blobs/sha256/30fb6218a3a4c7ea900128f10ef45e0e33b7df8b19c48f1e91cce984586cdc3e", + "blobs/sha256/3537474cb776983ae719c5a4a141d940777320e985371a4f88887117a64d5108", + "blobs/sha256/3a0d5962c5fcd07ccaf39e4a12985c440536eb2e0768f6f0c08fbc71bf699cfb", + "blobs/sha256/3ff483a4017c9c0d5c5ce8c57df85996144da7a2cdb440d4010914e6e6d37153", + "blobs/sha256/400821fd86ee1a17c460a7addf3892121cfb660a274b37e6a23cf36e7eaa931c", + "blobs/sha256/41aa8951d381527892cc220ff0c1d6e1ddd5527f3b8612066b37be7abcbfa81f", + "blobs/sha256/43f0f6265ae934400a0bd140cca733929de5026ae47360c2608cbc6c99508d89", + "blobs/sha256/44136fa355b3678a1146ad16f7e8649e94fb4fc21fe77e8310c060f61caaff8a", + "blobs/sha256/4477ccaea26652af48319688a16221ebb55a7a0de91c09ac8117cf4a88d203ad", + "blobs/sha256/4f4fb700ef54461cfa02571ae0db9a0dc1e0cdb5577484a6d75e68dc38e8acc1", + "blobs/sha256/5109ac5cb09b4c1eb33e1af8e806e34d0cd2f99375ff9c503de47f85e5611e83", + "blobs/sha256/527294152c9b07684d4d1f02b0afe6609d176282aa61d8b8968867c833481c56", + "blobs/sha256/5347e3e9b03fcfe4c347eb536fb743b5655e2ee0efc286719c930707784dab34", + "blobs/sha256/58d2172ebde0f9fecf755ac29d213f92adaa11806d6bb09d5ba7eee68ca4d341", + "blobs/sha256/5aa2bdef1d0c2fc3cbc86b86d351369901a24126f22e974e914c4c15bb01be0c", + "blobs/sha256/5bec02c16746b1fe15687b8df421f93fa2638f05ca2cd423d5be3c7979979d65", + "blobs/sha256/5bfec0e8bb11ecc202d0ca16cdb1f0e3390ad25705bd6a4b59284f04c85a940c", + "blobs/sha256/5e25946ba40dcbc2410bf1852d135a46cb31cc7f89c05b9ef1ab12a6ef988aeb", + "blobs/sha256/5fe493bec4795a3b68752f0b0f987d20efc48388fce5c5f7fad65e2011d31733", + "blobs/sha256/689b91d88a0f4086057ec826027b128902ecf2b516be510371c115bc55da19a6", + "blobs/sha256/730c8029606bf4b4533d7a7159458cf5e3f348ea956ccfb10ce5da3e20adac8f", + "blobs/sha256/753fe1c8b50013fccd0c156fb7091fe949b926876af6602c0a2e037d049ed14e", + "blobs/sha256/7723c0339d4ff66ed471c2c7b5e900e877a2477efa33b17cba42092ec5ef248b", + "blobs/sha256/79f971bb1a784b1a29ad3ae7e1376f146545dfe9a798ae600feb2eaaadb5fd81", + "blobs/sha256/848e634a564c2ca96ad2ce7a12bc0956e50c8186548a3aceac273afa7387a377", + "blobs/sha256/88d0092c6f752e77b376e7a0e50aa9547d66f8eb8b5e32b6f2e0aa3e1dd51a44", + "blobs/sha256/8b7f3999bbb358a5d5bc445380aba87aec6725bcbb41882d33ef9346ca105881", + "blobs/sha256/8bd080c2e4b22c71c5058be60551a2735c1bead7341fe005313ca704578d9fa5", + "blobs/sha256/92d6eafe91c933626ee8d9d9530de26ded162a5c81e3e5030c89384c98774c99", + "blobs/sha256/95e35692d3cc347d712fd236fd2af3b1c697783518182b1ae716722f694ed616", + "blobs/sha256/9681265a10ff460cc05d28f0bf5a56356b18a7d595f94d66f18ee2fea2ee82a8", + "blobs/sha256/9da825116641fadfcf4f55e849aabf92994ce1ba50760a68a1647f2df9347e86", + "blobs/sha256/a073d6f2f59b8e347d00cf592bb1f62da9f5eee3d25462dabab3f473abbc0204", + "blobs/sha256/a30ad5edf82cd178710efe9eaeac43d1e78bbcca4ab2fd20dac0cc3eb6bf1bba", + "blobs/sha256/a30da298bda6986d68697bc361b9204f189616a3b7cc748c91846ecddbb873aa", + "blobs/sha256/a496798fdd1a6ed0f8df3ca91bbbad569c7a8a5e3ad6f28e06fd237bbfe6f6c5", + "blobs/sha256/adebc1576f3c07100011ce7679aef5e7bb1e4aaca050e6d8e729273d22a2442a", + "blobs/sha256/b2848d797b1413db0b192a69782d4786ce9226cbf6be47154f06f7675f11bab6", + "blobs/sha256/b533a7368b7da8141f98b9a2442f84b5aa5d32f2d286c2611a32f725361f84b6", + "blobs/sha256/b80baf68b599275411a988ae77891a2782b66425a216fa07ef080d2efc6197d8", + "blobs/sha256/be7e2f489592134bebe19820b8233d5ba1da985c43f1fb46133741dbd4b3b272", + "blobs/sha256/cb3db0ad5ae089b6ae0c5b39b73e73cac40a558df052e602f8b0f24972b003b4", + "blobs/sha256/cc689184d8c0977f56d2da63b3a3d225dc9c7fadb9375338768cf4909b016b32", + "blobs/sha256/ccf9c70a85fc6158d617c85f3f719f669efb00c28fd45fa292575d351435bd2f", + "blobs/sha256/d103ce79450cbe8e275382ca45943531c3d08058f5c5fdd1d66f01acefac3a7f", + "blobs/sha256/d2dceafa44338ead1244a6542f284631554823289d552b98df574537fd0cc91a", + "blobs/sha256/d9f0d064333ffda4221fc5ebb606f9c7881ed90a6d8f0edf241f6f03c9cf6c6a", + "blobs/sha256/df898601d93cec8523eff57f3ca1cf20d2cf92793ce98d0c5e9885dd2a9c53d1", + "blobs/sha256/df998f8e25976466ca0e2aef4aff971125a81762b52c5b12446151ad8a24accb", + "blobs/sha256/e1f60ff4d853e5d51dfd06c483771ff1a24cf7e175af54203b4ce3bf831214d4", + "blobs/sha256/e3cd2362e641503f07d2ce339ee172131701b1a127f19d5a32561df9cd03d375", + "blobs/sha256/e4be11be2ac9f8e5786b81aee8b31ed8a210ac8a60c8ef6ec5e16057fb6cd702", + "blobs/sha256/ea972dfc554b3236f553b8b2e34c1c46df152889bfd037bc16d63bd8ceb92bd1", + "blobs/sha256/f07a09afcd4d724fe3068ef61532517a6962269f8dff5675ca221defcb678d34", + "blobs/sha256/f0af354daeda18f18684b23520969ef6bb3de82bb267936f867b12769be9125d", + "blobs/sha256/f3a41fe51ac1d40a35d538a321c4ff1c36911ba9eec5cd39ea0c168491d7385e", + "blobs/sha256/fcb0cc52852695975c6478e1df44c074260fa1004c497e2abe54c9cabe6e5ab1" + ] +}