diff --git a/.claude/skills/add-dynamic-filter/SKILL.md b/.claude/skills/add-dynamic-filter/SKILL.md index 204139439..46e58e161 100644 --- a/.claude/skills/add-dynamic-filter/SKILL.md +++ b/.claude/skills/add-dynamic-filter/SKILL.md @@ -27,7 +27,7 @@ Use this skill when: ### Step 2: Implement the Function Signature -Dynamic sampling filter (called in `slime/rollout/sglang_rollout.py`): +Dynamic sampling filter (called in `slime/rollout/vllm_rollout.py`): ```python def filter_function(args, samples, **kwargs): @@ -97,5 +97,5 @@ Example wiring: - Dynamic filter types: `slime/rollout/filter_hub/base_types.py` - Dynamic filter example: `slime/rollout/filter_hub/dynamic_sampling_filters.py` -- Rollout generation hook points: `slime/rollout/sglang_rollout.py` +- Rollout generation hook points: `slime/rollout/vllm_rollout.py` - Buffer filter hook point: `slime/rollout/data_source.py` diff --git a/.claude/skills/add-eval-dataset-config/SKILL.md b/.claude/skills/add-eval-dataset-config/SKILL.md index f59360687..6ff2463e5 100644 --- a/.claude/skills/add-eval-dataset-config/SKILL.md +++ b/.claude/skills/add-eval-dataset-config/SKILL.md @@ -87,5 +87,5 @@ Use a separate eval function when inference/eval behavior must differ from train - Eval config model: `slime/utils/eval_config.py` - Eval config resolution: `slime/utils/arguments.py` -- Eval rollout path: `slime/rollout/sglang_rollout.py` +- Eval rollout path: `slime/rollout/vllm_rollout.py` - Customization docs: `docs/en/get_started/customization.md` diff --git a/.claude/skills/add-rollout-function/SKILL.md b/.claude/skills/add-rollout-function/SKILL.md index c8623a135..39aa2f8ca 100644 --- a/.claude/skills/add-rollout-function/SKILL.md +++ b/.claude/skills/add-rollout-function/SKILL.md @@ -1,6 +1,6 @@ --- name: add-rollout-function -description: Guide for adding a new rollout function in slime and wiring it through --rollout-function-path. Use when user wants to implement custom rollout data generation logic, custom train/eval rollout outputs, or migrate from the default sglang rollout path. +description: Guide for adding a new rollout function in slime and wiring it through --rollout-function-path. Use when user wants to implement custom rollout data generation logic, custom train/eval rollout outputs, or migrate from the default vLLM rollout path. --- # Add Rollout Function @@ -12,7 +12,7 @@ Implement a custom rollout function and integrate it safely with slime training/ Use this skill when: - User asks to add a new rollout task or rollout generation function -- User asks to replace default `slime.rollout.sglang_rollout.generate_rollout` +- User asks to replace default `slime.rollout.vllm_rollout.generate_rollout` - User asks to customize train/eval data generation behavior ## Step-by-Step Guide @@ -21,10 +21,10 @@ Use this skill when: Start from one of these references: -- Async RL-style rollout: `slime/rollout/sglang_rollout.py` +- Async RL-style rollout: `slime/rollout/vllm_rollout.py` - Simple SFT-style rollout: `slime/rollout/sft_rollout.py` -If the task needs engine-based async generation and rewards, use the sglang path as base. +If the task needs engine-based async generation and rewards, use the vLLM path as base. If the task is file/buffer-driven and simple, use sft path as base. ### Step 2: Create the New Rollout Module @@ -100,7 +100,7 @@ The default and signature expectation are documented in: ## Reference Locations -- Default rollout: `slime/rollout/sglang_rollout.py` +- Default rollout: `slime/rollout/vllm_rollout.py` - Simple custom example: `slime/rollout/sft_rollout.py` - Output dataclasses: `slime/rollout/base_types.py` - Wiring/loading: `slime/ray/rollout.py` diff --git a/.github/workflows/conda-ci.yml b/.github/workflows/conda-ci.yml deleted file mode 100644 index 6cb0f5ce2..000000000 --- a/.github/workflows/conda-ci.yml +++ /dev/null @@ -1,90 +0,0 @@ -name: conda CI - -on: - pull_request: - branches: [main] - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} - cancel-in-progress: true - -jobs: - build-conda: - if: contains(github.event.pull_request.title, '[release]') - runs-on: self-hosted - container: - image: lmsysorg/sglang:v0.5.0rc0-cu126 - options: --privileged --cap-add SYS_NICE --security-opt seccomp=unconfined --gpus all --ipc=host --shm-size=16g --ulimit memlock=-1 --ulimit stack=67108864 --memory=0 --memory-swap=0 -v /mnt/nvme0n1/models:/root/models -v /mnt/nvme0n1/datasets:/root/datasets - - defaults: - run: - working-directory: ${{ github.workspace }} - - steps: - - name: Checkout repository - uses: actions/checkout@v6 - - - name: Construct Conda - run: | - echo "📦 Installing slime..." - cd $GITHUB_WORKSPACE - echo "Current directory: $(pwd)" - - mkdir -p /root/ - BASE_DIR=/root bash build_conda.sh - shell: bash - - - name: Download model and dataset - run: | - echo "🔗 Downloading up model and dataset..." - - # Create cache directories if they don't exist - mkdir -p /root/models /root/datasets - - echo "Downloading Qwen3-30B-A3B..." - hf download Qwen/Qwen3-30B-A3B --local-dir /root/models/Qwen3-30B-A3B - hf download Qwen/Qwen3-30B-A3B-FP8 --local-dir /root/models/Qwen3-30B-A3B-FP8 - - hf download --repo-type dataset zhuzilin/dapo-math-17k --local-dir /root/datasets/dapo-math-17k - - hf download --repo-type dataset zhuzilin/aime-2024 --local-dir /root/datasets/aime-2024 - shell: bash - - - name: Convert checkpoint - run: | - echo "🔄 Converting model checkpoint..." - cd $GITHUB_WORKSPACE - echo "Current directory: $(pwd)" - - source ~/.bashrc - micromamba activate slime - export CUDA_HOME="$CONDA_PREFIX" - - source scripts/models/qwen3-30B-A3B.sh - PYTHONPATH=/root/Megatron-LM torchrun --nproc-per-node 8 tools/convert_hf_to_torch_dist.py \ - ${MODEL_ARGS[@]} \ - --hf-checkpoint /root/models/Qwen3-30B-A3B \ - --save /root/Qwen3-30B-A3B_torch_dist - shell: bash - - - name: Run tests - run: | - echo "🧪 Running tests..." - cd $GITHUB_WORKSPACE - echo "Current directory: $(pwd)" - - source ~/.bashrc - micromamba activate slime - export CUDA_HOME="$CONDA_PREFIX" - - SLIME_TEST_USE_DEEPEP=0 SLIME_TEST_USE_FP8_ROLLOUT=0 python tests/test_qwen3_30B_A3B.py - shell: bash - - - name: Cleanup - if: always() - run: | - echo "🧹 Cleaning up..." - pkill -9 ray || true - ray stop --force || true - pkill -9 python || true - shell: bash diff --git a/README.md b/README.md index 153cea52d..c33021be9 100644 --- a/README.md +++ b/README.md @@ -45,7 +45,7 @@ We also provide examples for some use cases not covered in the quick start guide Arguments in Vime are divided into three categories: 1. **Megatron arguments**: Vime reads all arguments in Megatron. You can configure Megatron by passing arguments like `--tensor-model-parallel-size 2`. -2. **vLLM arguments**: vLLM server and engine options are exposed with a `--vllm-` prefix (for example, `--vllm-gpu-memory-utilization`). Router options live under two prefixes: vllm-router's native options are passed with `--router-` (for example, `--router-policy round_robin`), while Vime-side orchestration knobs that tell Vime *where* the router lives use `--vllm-router-` (`--vllm-router-ip`, `--vllm-router-port`, `--vllm-router-request-timeout-secs`). See [slime/backends/vllm_utils/arguments.py](slime/backends/vllm_utils/arguments.py) for the full surface. +2. **vLLM arguments**: vLLM server and engine options are exposed with a `--vllm-` prefix (for example, `--vllm-gpu-memory-utilization`). Router options live under two prefixes: vllm-router's native options are passed with `--router-` (for example, `--router-policy round_robin`, `--router-request-timeout-secs`), while Vime-side orchestration knobs that tell Vime *where* the router lives use `--vllm-router-` (`--vllm-router-ip`, `--vllm-router-port`). See [slime/backends/vllm_utils/arguments.py](slime/backends/vllm_utils/arguments.py) for the full surface. 3. **Framework-specific arguments**: Shared slime/Vime orchestration flags (rollout GPUs, data paths, RL algorithms, etc.). Please refer to [slime/utils/arguments.py](slime/utils/arguments.py). `--rollout-num-gpus-per-engine` sets the tensor parallel size of each vLLM engine. The default rollout entry is `slime.rollout.vllm_rollout.generate_rollout`. diff --git a/README_zh.md b/README_zh.md index 54b121da9..2804a7728 100644 --- a/README_zh.md +++ b/README_zh.md @@ -45,7 +45,7 @@ Vime 继承 slime 的广泛模型支持,包括: Vime 的参数分为三类: 1. **Megatron 参数**:Vime 会读取 Megatron 中的全部参数,可通过传入如 `--tensor-model-parallel-size 2` 的方式配置 Megatron; -2. **vLLM 参数**:vLLM server 与 engine 相关选项以 `--vllm-` 为前缀(例如 `--vllm-gpu-memory-utilization`)。路由相关选项分两类前缀:vllm-router 自身的选项以 `--router-` 传入(例如 `--router-policy round_robin`),Vime 侧用于告诉 Vime *router 在哪里* 的编排参数则以 `--vllm-router-` 为前缀(`--vllm-router-ip`、`--vllm-router-port`、`--vllm-router-request-timeout-secs`)。完整参数见 [slime/backends/vllm_utils/arguments.py](slime/backends/vllm_utils/arguments.py)。 +2. **vLLM 参数**:vLLM server 与 engine 相关选项以 `--vllm-` 为前缀(例如 `--vllm-gpu-memory-utilization`)。路由相关选项分两类前缀:vllm-router 自身的选项以 `--router-` 传入(例如 `--router-policy round_robin`、`--router-request-timeout-secs`),Vime 侧用于告诉 Vime *router 在哪里* 的编排参数则以 `--vllm-router-` 为前缀(`--vllm-router-ip`、`--vllm-router-port`)。完整参数见 [slime/backends/vllm_utils/arguments.py](slime/backends/vllm_utils/arguments.py)。 3. **框架参数**:与 slime/Vime 编排相关的开关(rollout GPU、数据路径、RL 算法等),见 [slime/utils/arguments.py](slime/utils/arguments.py)。 `--rollout-num-gpus-per-engine` 对应每个 vLLM engine 的 tensor parallel size。默认 rollout 入口为 `slime.rollout.vllm_rollout.generate_rollout`。 diff --git a/build_conda.sh b/build_conda.sh deleted file mode 100644 index 1d8210b23..000000000 --- a/build_conda.sh +++ /dev/null @@ -1,84 +0,0 @@ -#!/bin/bash - -set -ex - -# create conda -yes '' | "${SHELL}" <(curl -L micro.mamba.pm/install.sh) -export PS1=tmp -mkdir -p /root/.cargo/ -touch /root/.cargo/env -source ~/.bashrc - -micromamba create -n slime python=3.12 pip -c conda-forge -y -micromamba activate slime -export CUDA_HOME="$CONDA_PREFIX" -export SGLANG_COMMIT="bbe9c7eeb520b0a67e92d133dfc137a3688dc7f2" -export MEGATRON_COMMIT="3714d81d418c9f1bca4594fc35f9e8289f652862" - -export BASE_DIR=${BASE_DIR:-"/root"} -cd $BASE_DIR - -# install cuda 12.9 as it's the default cuda version for torch -micromamba install -n slime cuda cuda-nvtx cuda-nvtx-dev nccl -c nvidia/label/cuda-12.9.1 -y -micromamba install -n slime -c conda-forge cudnn -y - -pip install cuda-python==12.9 -pip install torch==2.9.1 torchvision==0.24.1 torchaudio==2.9.1 --index-url https://download.pytorch.org/whl/cu129 - -# install sglang -git clone https://github.com/sgl-project/sglang.git -cd sglang -git checkout ${SGLANG_COMMIT} -# Install the python packages -pip install -e "python[all]" - - -pip install cmake ninja - -# flash attn -# the newest version megatron supports is v2.7.4.post1 -MAX_JOBS=64 pip -v install flash-attn==2.7.4.post1 --no-build-isolation - -pip install git+https://github.com/ISEEKYAN/mbridge.git@89eb10887887bc74853f89a4de258c0702932a1c --no-deps -pip install --no-build-isolation "transformer_engine[pytorch]==2.10.0" -pip install flash-linear-attention==0.4.1 -NVCC_APPEND_FLAGS="--threads 4" \ - pip -v install --disable-pip-version-check --no-cache-dir \ - --no-build-isolation \ - --config-settings "--build-option=--cpp_ext --cuda_ext --parallel 8" git+https://github.com/NVIDIA/apex.git@10417aceddd7d5d05d7cbf7b0fc2daad1105f8b4 - -pip install git+https://github.com/fzyzcjy/torch_memory_saver.git@dc6876905830430b5054325fa4211ff302169c6b --no-cache-dir --force-reinstall -pip install git+https://github.com/fzyzcjy/Megatron-Bridge.git@dev_rl --no-build-isolation -pip install nvidia-modelopt[torch]>=0.37.0 --no-build-isolation -pip install https://github.com/zhuzilin/sgl-router/releases/download/v0.3.2-5f8d397/sglang_router-0.3.2-cp38-abi3-manylinux_2_28_x86_64.whl --force-reinstall - -# megatron -cd $BASE_DIR -git clone https://github.com/NVIDIA/Megatron-LM.git --recursive && \ - cd Megatron-LM/ && git checkout ${MEGATRON_COMMIT} && \ - pip install -e . - -# install slime and apply patches - -# if slime does not exist locally, clone it -if [ ! -d "$BASE_DIR/slime" ]; then - cd $BASE_DIR - git clone https://github.com/THUDM/slime.git - cd slime/ - export SLIME_DIR=$BASE_DIR/slime - pip install -e . -else - export SLIME_DIR=$BASE_DIR/slime - cd $SLIME_DIR - pip install -e . -fi - -# https://github.com/pytorch/pytorch/issues/168167 -pip install nvidia-cudnn-cu12==9.16.0.29 -pip install "numpy<2" - -# apply patch -cd $BASE_DIR/sglang -git apply $SLIME_DIR/docker/patch/v0.5.9/sglang.patch -cd $BASE_DIR/Megatron-LM -git apply $SLIME_DIR/docker/patch/v0.5.9/megatron.patch diff --git a/docker/Dockerfile b/docker/Dockerfile index 926b407e1..45dd1bcd9 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -1,11 +1,10 @@ -ARG BASE_IMAGE=vllm/vllm-openai:v0.21.0-cu129-ubuntu2404 +ARG BASE_IMAGE=vllm/vllm-openai:v0.22.0-cu129-ubuntu2404 FROM ${BASE_IMAGE} # ======================================== Arguments ============================================= ARG PATCH_VERSION=latest ARG MEGATRON_COMMIT=1dcf0dafa884ad52ffb243625717a3471643e087 -ARG SGLANG_VERSION=0.5.10.post1 ARG ENABLE_CUDA_13=0 @@ -74,18 +73,9 @@ RUN if [ "$ENABLE_CUDA_13" = "1" ]; then \ (cd /root && git clone -b feat/v350_plus_8045 https://github.com/fzyzcjy/triton.git && cd triton && pip install -r python/requirements.txt && pip install --verbose -e .); \ fi -# Skip sglang lines from requirements (we install --no-deps below; the full -# transitive set conflicts with the vllm base environment). COPY requirements.txt /tmp/requirements.txt -RUN grep -vE '^[[:space:]]*(sglang|sglang-router)([[:space:]]|[<>=!]|$)' /tmp/requirements.txt > /tmp/requirements-vllm.txt && \ - pip install --ignore-installed PyJWT && \ - pip install -r /tmp/requirements-vllm.txt - -# Temporarily install another sgl-kernel version for GB300 without rebuilding the whole image -RUN if [ "$ENABLE_CUDA_13" = "1" ]; then \ - SGL_KERNEL_VERSION=0.3.17.post2 && \ - python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sgl_kernel-${SGL_KERNEL_VERSION}+cu130-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps; \ - fi +RUN pip install --ignore-installed PyJWT && \ + pip install -r /tmp/requirements.txt # https://github.com/pytorch/pytorch/issues/168167 RUN pip install nvidia-cudnn-cu12==9.16.0.29 @@ -93,10 +83,7 @@ RUN pip install nvidia-cudnn-cu12==9.16.0.29 # reinstall numpy 1.x for megatron RUN pip install "numpy<2" -# vime's slime/utils/arguments.py + slime/ray/rollout.py top-level import sglang_router; -# the cu129 vllm base does not ship sglang, so install --no-deps stubs here. -RUN pip install --no-deps "sglang==${SGLANG_VERSION}" sglang-router==0.3.2 && \ - pip install IPython +RUN pip install IPython # Pin vllm-router explicitly so the vllm rollout routing layer is a visible build step # (also in requirements.txt; pulling it here makes the layer cache-able and fail-fast). @@ -135,10 +122,9 @@ RUN cd /root/slime/slime/backends/megatron_utils/kernels/int4_qat && \ # Fail-fast import smoke + flashinfer version pin check. Catches ABI / version # regressions at build time instead of first GPU run. RUN python3 -c "\ -import vllm, sglang, slime, flashinfer; \ +import vllm, slime, flashinfer; \ from slime.backends.vllm_utils.vllm_engine import VLLMEngine; \ print('vllm', vllm.__version__); \ -print('sglang', sglang.__version__); \ print('flashinfer', flashinfer.__version__); \ print('VLLMEngine import ok')" diff --git a/docker/justfile b/docker/justfile index 2acb92552..b02095375 100644 --- a/docker/justfile +++ b/docker/justfile @@ -1,13 +1,17 @@ release-primary: ARG_TAG_POSTFIX="" ARG_BUILD_EXTRA_ARGS="" just _release-raw -# Should be executed on ARM machines +# Should be executed on ARM machines. +# Inherits the Dockerfile's default cu129 BASE_IMAGE; that tag is a multi-arch +# manifest, so docker selects the arm64 image automatically on an ARM host. release-cu129-arm64: - ARG_TAG_POSTFIX="-cu129-arm64" ARG_BUILD_EXTRA_ARGS='--build-arg SGLANG_IMAGE_TAG=v0.5.5.post3-cu129-arm64 --build-arg ENABLE_SGLANG_PATCH=0' just _release-raw + ARG_TAG_POSTFIX="-cu129-arm64" ARG_BUILD_EXTRA_ARGS="" just _release-raw -# Should be executed on ARM machines +# Should be executed on ARM machines. +# The default-CUDA vLLM tag (no cuXXX suffix) already ships CUDA 13.0; +# ENABLE_CUDA_13 then builds the CUDA-13 TransformerEngine/Triton on top. release-cu13-arm64: - ARG_TAG_POSTFIX="-cu13-arm64" ARG_BUILD_EXTRA_ARGS='--build-arg SGLANG_IMAGE_TAG=dev-arm64-cu13-20251122 --build-arg ENABLE_CUDA_13=1 --build-arg ENABLE_SGLANG_PATCH=0' just _release-raw + ARG_TAG_POSTFIX="-cu13-arm64" ARG_BUILD_EXTRA_ARGS='--build-arg BASE_IMAGE=vllm/vllm-openai:v0.22.0-ubuntu2404 --build-arg ENABLE_CUDA_13=1' just _release-raw _release-raw: #!/bin/bash diff --git a/docker/npu_patch/README.md b/docker/npu_patch/README.md deleted file mode 100644 index db93b6763..000000000 --- a/docker/npu_patch/README.md +++ /dev/null @@ -1,192 +0,0 @@ -# Slime NPU Patch Installation Guide - -This guide provides instructions for installing Slime with NPU support, including all required dependencies and patches. - -## Component Version Mapping - -| Component | Version/Commit | Source | -| --------------- | ---------------------------------------- | ------------------------------------------------------------------------------------------------------------------- | -| Slime | v0.2.2 | [GitHub](https://github.com/THUDM/slime/tree/v0.2.2) | -| SGLang | dce8b0606c06d3a191a24c7b8cbe8e238ab316c9 | [GitHub](https://github.com/sgl-project/sglang/tree/sglang-slime) | -| SGL Kernel NPU | 2026.02.01 | [GitHub](https://github.com/sgl-project/sgl-kernel-npu/releases/tag/2026.02.01) | -| Megatron-Bridge | 35b4ebfc486fb15dcc0273ceea804c3606be948a | [GitHub](https://github.com/fzyzcjy/Megatron-Bridge) | -| Megatron-LM | 3714d81d418c9f1bca4594fc35f9e8289f652862 | [GitHub](https://github.com/NVIDIA/Megatron-LM) | -| MindSpeed | fc63de5c48426dd019c3b3f39e65f5bdf56e4086 | [GitCode](https://gitcode.com/Ascend/MindSpeed) | -| HDK | 25.3.RC1 | [Ascend](https://www.hiascend.com/hardware/firmware-drivers/commercial?product=7\&model=33) | -| CANN | 8.5.0 | [Ascend](https://www.hiascend.com/developer/download/community/result?module=cann\&cann=8.5.0\&product=7\&model=33) | - -## Preparing the Running Environment - -### Python Version - -Only `python==3.11` is supported currently. - -```shell -conda create -n slime_release python=3.11 -conda activate slime_release -``` - -### Working Directory Setup - -```shell -mkdir && cd -``` - -### CANN Environment - -Prior to start work with Slime on Ascend you need to install CANN Toolkit, Kernels operator package and NNAL version 8.5.0, check the [installation guide](https://www.hiascend.com/document/detail/zh/CANNCommunityEdition/83RC1/softwareinst/instg/instg_0008.html?Mode=PmIns\&InstallType=local\&OS=openEuler\&Software=cannToolKit) - -```shell -source /ascend-toolkit/set_env.sh -source /nnal/atb/set_env.sh -``` - -### PyTorch and PyTorch NPU - -```shell -pip install torch-npu==2.8.0 -``` - -## Installing Dependencies - -### SGLang - -```shell -cd -git clone https://github.com/sgl-project/sglang.git && cd sglang -git checkout dce8b0606c06d3a191a24c7b8cbe8e238ab316c9 -mv python/pyproject.toml python/pyproject.toml.backup -mv python/pyproject_other.toml python/pyproject.toml -pip install -e "python[srt_npu]" -pip install torch-npu==2.8.0 -``` - -### SGL Kernel NPU and Torch Memory Saver - -Download `sgl-kernel-npu-2026.02.01-torch2.8.0-py311-cann8.5.0-a3-aarch64.zip` from the release link, then install: - -```shell -pip install sgl_kernel_npu-2026.2.1-cp311-cp311-linux_aarch64.whl -pip install torch_memory_saver-0.0.8-cp311-cp311-linux_aarch64.whl -``` - -### Megatron-Bridge - -```shell -pip install git+https://github.com/ISEEKYAN/mbridge.git@89eb10887887bc74853f89a4de258c0702932a1c --no-deps - -cd -git clone https://github.com/fzyzcjy/Megatron-Bridge.git -b dev_rl -pip install nvidia-modelopt[torch]>=0.37.0 --no-build-isolation -``` - -### Megatron-LM - -```shell -cd -git clone https://github.com/NVIDIA/Megatron-LM.git --recursive && \ - cd Megatron-LM/ && git checkout 3714d81d418c9f1bca4594fc35f9e8289f652862 && \ - pip install -e . -``` - -### MindSpeed - -```shell -cd -git clone https://gitcode.com/Ascend/MindSpeed.git && \ - cd MindSpeed/ && git checkout fc63de5c48426dd019c3b3f39e65f5bdf56e4086 && \ - pip install -e . -``` - -### Slime - -```shell -cd -git clone https://github.com/ascend-slime/slime.git && cd slime -cp -r docker/npu_patch ../npu_patch -git checkout v0.2.2 -pip install -e . -``` - -## Applying Patches - -```shell -cd /slime -git apply ../npu_patch/slime.patch - -cd /sglang -git apply ../slime/docker/patch/v0.5.7/sglang.patch -git apply ../npu_patch/sglang.patch - -cd /Megatron-LM -git apply ../slime/docker/patch/v0.5.7/megatron.patch -git apply ../npu_patch/megatron.patch - -cd /Megatron-Bridge -git apply ../npu_patch/megatron-bridge.patch - -cd /MindSpeed -git apply ../npu_patch/mindspeed.patch -``` - -## Additional Dependencies - -```shell -cd /slime -pip install triton-ascend -pip install torch-npu==2.8.0 -pip install torchvision==0.23.0 -pip install numpy==1.26.0 -``` - -## Running the Training - -### Configuration - -Modify the paths in the following files according to your environment (note to use your CANN version): - -**Common (both GRPO and PPO):** - -- `slime/utils/external_utils/command_utils.py` - -**GRPO:** - -- `examples/geo3k_vlm_multi_turn/run_grpo_npu.sh` -- `examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_grpo_npu.py` - -**PPO:** - -- `examples/geo3k_vlm_multi_turn/run_ppo_npu.sh` -- `examples/geo3k_vlm_multi_turn/run_geo3k_vlm_multi_turn_ppo_npu.py` - -### Dataset - -Download the dataset from [HuggingFace](https://huggingface.co/datasets/VeraIsHere/geo3k_imgurl_processed) following the instructions in the script directory. - -### Execute Training - -```shell -cd /slime -# GRPO -bash examples/geo3k_vlm_multi_turn/run_grpo_npu.sh -# PPO -bash examples/geo3k_vlm_multi_turn/run_ppo_npu.sh -``` - -To save logs and display them simultaneously: - -```shell -# GRPO -bash examples/geo3k_vlm_multi_turn/run_grpo_npu.sh 2>&1 | tee -a -# PPO -bash examples/geo3k_vlm_multi_turn/run_ppo_npu.sh 2>&1 | tee -a -``` - -## Placeholders Reference - -| Placeholder | Description | Example | -| ------------- | ------------------------------------ | --------------------- | -| `` | Root directory for all installations | `/root/slime-release` | -| `` | Path to CANN installation directory | `/usr/local/ascend` | -| `` | Path to log file for training output | `training.log` | - diff --git a/docker/npu_patch/megatron-bridge.patch b/docker/npu_patch/megatron-bridge.patch deleted file mode 100644 index 3817097f4..000000000 --- a/docker/npu_patch/megatron-bridge.patch +++ /dev/null @@ -1,93 +0,0 @@ -diff --git a/src/megatron/bridge/models/conversion/param_mapping.py b/src/megatron/bridge/models/conversion/param_mapping.py -index dc7d0be..8156826 100644 ---- a/src/megatron/bridge/models/conversion/param_mapping.py -+++ b/src/megatron/bridge/models/conversion/param_mapping.py -@@ -1088,15 +1088,19 @@ class AutoMapping(MegatronParamMapping[torch.Tensor]): - "ColumnParallelLinear", - "TEColumnParallelLinear", - "TELayerNormColumnParallelLinear", -+ "MindSpeedTELayerNormColumnParallelLinear", - "TEColumnParallelGroupedLinear", -+ "MindSpeedTEColumnParallelGroupedLinear", - "VocabParallelEmbedding", - "DotProductAttention", # for attention sink only - "TEDotProductAttention", # for attention sink only -+ "MindSpeedTEDotProductAttention", - }, - "row": { - "RowParallelLinear", - "TERowParallelLinear", - "TERowParallelGroupedLinear", -+ "MindSpeedTERowParallelGroupedLinear", - }, - "replicated": { - # Normalization layers -@@ -1164,7 +1168,7 @@ class AutoMapping(MegatronParamMapping[torch.Tensor]): - # Handle fused modules like TELayerNormColumnParallelLinear - # These modules have both column-parallel weights (weight, bias) - # and replicated layer norm weights (layer_norm_weight, layer_norm_bias) -- if module_type == "TELayerNormColumnParallelLinear": -+ if module_type == "TELayerNormColumnParallelLinear" or module_type == "MindSpeedTELayerNormColumnParallelLinear": - # Check the actual parameter name to determine the correct parallelism type - if self.megatron_param and ( - self.megatron_param.endswith("layer_norm_weight") or self.megatron_param.endswith("layer_norm_bias") -@@ -1195,7 +1199,7 @@ class AutoMapping(MegatronParamMapping[torch.Tensor]): - return "replicated" - - # Check parallel_mode for TELinear -- if module_type == "TELinear": -+ if module_type == "TELinear" or module_type == "MindSpeedTELinear": - if module.parallel_mode == "column": - return "column" - elif module.parallel_mode == "row": -diff --git a/src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/transformer_block.py b/src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/transformer_block.py -index 3f64d8d..9dadc4b 100644 ---- a/src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/transformer_block.py -+++ b/src/megatron/bridge/models/qwen_vl/modelling_qwen3_vl/transformer_block.py -@@ -71,8 +71,9 @@ class Qwen3VLTransformerBlock(TransformerBlock): - context_mask, - rotary_pos_emb, - visual_pos_masks, -- deepstack_visual_embeds, -+ *deepstack_visual_embeds_args, - ): -+ deepstack_visual_embeds = list(deepstack_visual_embeds_args) if deepstack_visual_embeds_args else None - for index in range(start, end): - layer = self._get_layer(index) - inner_fp8_context = ( -@@ -103,6 +104,8 @@ class Qwen3VLTransformerBlock(TransformerBlock): - return hidden_states, context - - return custom_forward -+ -+ deepstack_visual_embeds_tuple = tuple(deepstack_visual_embeds) if deepstack_visual_embeds else () - - def checkpoint_handler(forward_func): - """Determines whether to use the `te_checkpoint` or `tensor_parallel.checkpoint`""" -@@ -118,7 +121,7 @@ class Qwen3VLTransformerBlock(TransformerBlock): - context_mask, - rotary_pos_emb, - visual_pos_masks, -- deepstack_visual_embeds, -+ *deepstack_visual_embeds_tuple, - ) - else: - return tensor_parallel.checkpoint( -@@ -130,7 +133,7 @@ class Qwen3VLTransformerBlock(TransformerBlock): - context_mask, - rotary_pos_emb, - visual_pos_masks, -- deepstack_visual_embeds, -+ *deepstack_visual_embeds_tuple, - ) - - if self.config.recompute_method == "uniform": -@@ -169,7 +172,7 @@ class Qwen3VLTransformerBlock(TransformerBlock): - context_mask, - rotary_pos_emb, - visual_pos_masks, -- deepstack_visual_embeds, -+ *deepstack_visual_embeds_tuple, - ) - else: - raise ValueError("Invalid activation recompute method.") diff --git a/docker/npu_patch/megatron.patch b/docker/npu_patch/megatron.patch deleted file mode 100644 index 7687827db..000000000 --- a/docker/npu_patch/megatron.patch +++ /dev/null @@ -1,518 +0,0 @@ -diff --git a/megatron/core/activations.py b/megatron/core/activations.py -index 8b422d73a..58fba4667 100644 ---- a/megatron/core/activations.py -+++ b/megatron/core/activations.py -@@ -5,19 +5,19 @@ import torch.nn.functional as F - from megatron.core.jit import jit_fuser - - --@jit_fuser -+ - def squared_relu(x: torch.Tensor) -> torch.Tensor: - """Squared ReLU activation""" - return torch.pow(F.relu(x), 2) - - --@jit_fuser -+ - def quick_gelu(x: torch.Tensor) -> torch.Tensor: - """Quick GELU activation""" - return x * torch.sigmoid(1.702 * x) - - --@jit_fuser -+ - def fast_gelu(x: torch.Tensor) -> torch.Tensor: - """Fast GELU activation""" - return 0.5 * x * (1.0 + torch.tanh(x * 0.7978845608 * (1.0 + 0.044715 * x * x))) -diff --git a/megatron/core/fusions/fused_bias_dropout.py b/megatron/core/fusions/fused_bias_dropout.py -index 336452562..614ee1a48 100644 ---- a/megatron/core/fusions/fused_bias_dropout.py -+++ b/megatron/core/fusions/fused_bias_dropout.py -@@ -64,14 +64,14 @@ def bias_dropout_add_unfused(training): - return _bias_dropout_add - - --@jit_fuser -+ - def bias_dropout_add_fused_train( - x_with_bias: Tuple[torch.Tensor, Optional[torch.Tensor]], residual: torch.Tensor, prob: float - ) -> torch.Tensor: - return _bias_dropout_add_func(x_with_bias, residual, prob, True) - - --@jit_fuser -+ - def bias_dropout_add_fused_inference( - x_with_bias: Tuple[torch.Tensor, Optional[torch.Tensor]], residual: torch.Tensor, prob: float - ) -> torch.Tensor: -diff --git a/megatron/core/fusions/fused_bias_geglu.py b/megatron/core/fusions/fused_bias_geglu.py -index 7a7fbe7f9..f9a438953 100644 ---- a/megatron/core/fusions/fused_bias_geglu.py -+++ b/megatron/core/fusions/fused_bias_geglu.py -@@ -13,7 +13,7 @@ from megatron.core.jit import jit_fuser - # x * 0.5 * (1.0 + torch.erf(x * 0.70710678)) - - --@jit_fuser -+ - def geglu(y): - """Performs GEGLU (GELU-Gated Linear Unit) activation. - -@@ -27,7 +27,7 @@ def geglu(y): - return (y_1 * 0.5 * (1.0 + torch.tanh(0.79788456 * y_1 * (1 + 0.044715 * y_1 * y_1)))) * y_2 - - --@jit_fuser -+ - def bias_geglu(bias, y): - """Performs GEGLU activation with bias addition. - -@@ -45,7 +45,7 @@ def bias_geglu(bias, y): - # gradient of tanh approximation of gelu - # gradient of actual gelu is: - # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x) --@jit_fuser -+ - def geglu_back(g, y): - """Computes the gradient for the GEGLU activation. - -@@ -65,7 +65,7 @@ def geglu_back(g, y): - return torch.cat(((g * y_2) * ff, g * (y_1 * 0.5 * (1.0 + tanh_out))), -1) - - --@jit_fuser -+ - def bias_geglu_back(g, y, bias): - """Computes the gradient for the biased GEGLU activation. - -@@ -181,13 +181,13 @@ def bias_geglu_impl(input, bias): - # ------------------------- QUICK GEGLU FUSION -------------------------- - - --@jit_fuser -+ - def quick_gelu(y: torch.Tensor) -> torch.Tensor: - """Sigmoid approximation of gelu""" - return y * torch.sigmoid(1.702 * y) - - --@jit_fuser -+ - def quick_geglu(y: torch.Tensor, linear_offset: float = 0.0) -> torch.Tensor: - """Performs Quick-GELU-based GEGLU activation : quick_gelu(y1) * (y2 + offset). - -@@ -202,7 +202,7 @@ def quick_geglu(y: torch.Tensor, linear_offset: float = 0.0) -> torch.Tensor: - return quick_gelu(y_1) * (y_2 + linear_offset) - - --@jit_fuser -+ - def weighted_quick_geglu( - y: torch.Tensor, weights: torch.Tensor, linear_offset: float = 0.0 - ) -> torch.Tensor: -@@ -217,7 +217,7 @@ def weighted_quick_geglu( - - - # gradient of sigmoid approximation of gelu --@jit_fuser -+ - def quick_geglu_back(g, y, linear_offset: float = 0.0) -> torch.Tensor: - """Backward helper for Quick-GEGLU. - -@@ -236,7 +236,7 @@ def quick_geglu_back(g, y, linear_offset: float = 0.0) -> torch.Tensor: - return torch.cat((dy_1, dy_2), -1) - - --@jit_fuser -+ - def weighted_quick_geglu_back(g, y, weights, linear_offset: float = 0.0): - """Backward helper for weighted Quick-GEGLU. - Returns gradient w.r.t input `y` and `weights`. -@@ -255,7 +255,7 @@ def weighted_quick_geglu_back(g, y, weights, linear_offset: float = 0.0): - # ---------------- Weighted Bias Quick-GEGLU helpers ----------------- - - --@jit_fuser -+ - def weighted_bias_quick_geglu( - y: torch.Tensor, bias: torch.Tensor, weights: torch.Tensor, linear_offset: float = 0.0 - ) -> torch.Tensor: -@@ -275,7 +275,7 @@ def weighted_bias_quick_geglu( - return res.to(dtype) - - --@jit_fuser -+ - def weighted_bias_quick_geglu_back(g, y, bias, weights, linear_offset: float = 0.0): - """Backward helper for weighted Quick-GEGLU with bias. - -diff --git a/megatron/core/fusions/fused_bias_gelu.py b/megatron/core/fusions/fused_bias_gelu.py -index 8cc90f617..fda8f2f5f 100644 ---- a/megatron/core/fusions/fused_bias_gelu.py -+++ b/megatron/core/fusions/fused_bias_gelu.py -@@ -13,7 +13,7 @@ from megatron.core.jit import jit_fuser - # x * 0.5 * (1.0 + torch.erf(x * 0.70710678)) - - --@jit_fuser -+ - def bias_gelu(bias, y): - x = bias + y - return x * 0.5 * (1.0 + torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))) -@@ -22,7 +22,7 @@ def bias_gelu(bias, y): - # gradient of tanh approximation of gelu - # gradient of actual gelu is: - # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x) --@jit_fuser -+ - def bias_gelu_back(g, bias, y): - x = bias + y - tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x)) -diff --git a/megatron/core/fusions/fused_bias_swiglu.py b/megatron/core/fusions/fused_bias_swiglu.py -index 632470876..105936786 100644 ---- a/megatron/core/fusions/fused_bias_swiglu.py -+++ b/megatron/core/fusions/fused_bias_swiglu.py -@@ -12,7 +12,7 @@ from megatron.core.utils import nvtx_decorator - ###### BIAS SWIGLU FUSION/ NO AUTOGRAD ################ - - --@jit_fuser -+ - def swiglu(y): - """Performs SwiGLU (Swish-Gated Linear Unit) activation function. - -@@ -26,7 +26,7 @@ def swiglu(y): - return F.silu(y_1) * y_2 - - --@jit_fuser -+ - def bias_swiglu(y, bias): - """Performs SwiGLU activation with bias addition. - -@@ -41,7 +41,7 @@ def bias_swiglu(y, bias): - return swiglu(y) - - --@jit_fuser -+ - def weighted_swiglu(y, weights): - dtype = y.dtype - res = swiglu(y) * weights -@@ -51,7 +51,7 @@ def weighted_swiglu(y, weights): - # gradient of tanh approximation of gelu - # gradient of actual gelu is: - # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x) --@jit_fuser -+ - def swiglu_back(g, y): - """Computes the gradient for the SwiGLU activation function. - -@@ -69,7 +69,7 @@ def swiglu_back(g, y): - ) - - --@jit_fuser -+ - def bias_swiglu_back(g, y, bias): - """Computes the gradient for the biased SwiGLU activation function. - -@@ -86,7 +86,7 @@ def bias_swiglu_back(g, y, bias): - return swiglu_back(g, y) - - --@jit_fuser -+ - def weighted_swiglu_back(g, y, weights): - input_dtype = y.dtype - w_dtype = weights.dtype -diff --git a/megatron/core/fusions/fused_cross_entropy.py b/megatron/core/fusions/fused_cross_entropy.py -index 23e4b6031..4d447e0e0 100644 ---- a/megatron/core/fusions/fused_cross_entropy.py -+++ b/megatron/core/fusions/fused_cross_entropy.py -@@ -9,7 +9,7 @@ from megatron.core.tensor_parallel.cross_entropy import VocabParallelCrossEntrop - from megatron.core.tensor_parallel.utils import VocabUtility - - --@jit_fuser -+ - def calculate_logits_max(vocab_parallel_logits: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: - """ - Calculates the maximum logits of the predicted tokens. -@@ -22,7 +22,7 @@ def calculate_logits_max(vocab_parallel_logits: torch.Tensor) -> Tuple[torch.Ten - return vocab_parallel_logits, logits_max - - --@jit_fuser -+ - def calculate_predicted_logits( - vocab_parallel_logits: torch.Tensor, - target: torch.Tensor, -@@ -44,7 +44,7 @@ def calculate_predicted_logits( - return target_mask, masked_target_1d, predicted_logits_sum_exp_logits, exp_logits - - --@jit_fuser -+ - def calculate_cross_entropy_loss( - exp_logits: torch.Tensor, predicted_logits_sum_exp_logits: torch.Tensor - ) -> Tuple[torch.Tensor, torch.Tensor]: -@@ -61,7 +61,7 @@ def calculate_cross_entropy_loss( - return exp_logits, loss - - --@jit_fuser -+ - def calculate_gradients( - softmax: torch.Tensor, - grad_output: torch.Tensor, -diff --git a/megatron/core/fusions/fused_pad_routing_map.py b/megatron/core/fusions/fused_pad_routing_map.py -index c382178b6..563279edd 100644 ---- a/megatron/core/fusions/fused_pad_routing_map.py -+++ b/megatron/core/fusions/fused_pad_routing_map.py -@@ -70,7 +70,7 @@ def _pad_routing_map_kernel( - tl.store(output_row_ptr + token_indices, output_row, mask=token_mask) - - --@jit_fuser -+ - def fused_pad_routing_map(routing_map: torch.Tensor, pad_multiple: int) -> torch.Tensor: - """Fused version of pad_routing_map. - Args: -diff --git a/megatron/core/fusions/fused_weighted_squared_relu.py b/megatron/core/fusions/fused_weighted_squared_relu.py -index 02dabc14c..137022386 100644 ---- a/megatron/core/fusions/fused_weighted_squared_relu.py -+++ b/megatron/core/fusions/fused_weighted_squared_relu.py -@@ -10,7 +10,7 @@ from megatron.core.utils import nvtx_decorator - ###################### WEIGHTED SQUARED ReLU FUSION ###################### - - --@jit_fuser -+ - def weighted_squared_relu(x: torch.Tensor, weights: torch.Tensor) -> torch.Tensor: - """Element-wise weight applied after Squared-ReLU. - -@@ -28,7 +28,7 @@ def weighted_squared_relu(x: torch.Tensor, weights: torch.Tensor) -> torch.Tenso - return res.to(out_dtype) - - --@jit_fuser -+ - def _squared_relu_back(g: torch.Tensor, x: torch.Tensor) -> torch.Tensor: - """Gradient of Squared-ReLU. - -@@ -37,7 +37,7 @@ def _squared_relu_back(g: torch.Tensor, x: torch.Tensor) -> torch.Tensor: - return g * 2 * F.relu(x) - - --@jit_fuser -+ - def weighted_squared_relu_back(g: torch.Tensor, x: torch.Tensor, weights: torch.Tensor): - """Backward for weighted Squared-ReLU. - -diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py -index 712793853..76ea4333b 100755 ---- a/megatron/core/models/gpt/gpt_layer_specs.py -+++ b/megatron/core/models/gpt/gpt_layer_specs.py -@@ -219,7 +219,7 @@ def get_gpt_layer_with_transformer_engine_spec( - 'The fp8 argument in "get_gpt_layer_with_transformer_engine_spec" has been deprecated' - " and will be removed soon. Please update your code accordingly." - ) -- -+ from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider - if use_kitchen: - assert HAVE_KITCHEN - backend: BackendSpecProvider = KitchenSpecProvider( -diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net.py -index dfa6e4c35..8b3621819 100644 ---- a/megatron/core/ssm/gated_delta_net.py -+++ b/megatron/core/ssm/gated_delta_net.py -@@ -413,7 +413,7 @@ class GatedDeltaNet(MegatronModule): - - return out, out_bias - -- @jit_fuser -+ - def _apply_gated_norm(self, x, gate): - # Output Norm - x_dtype = x.dtype -diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py -index 80e9ec6fc..b6bcef6a9 100644 ---- a/megatron/core/transformer/attention.py -+++ b/megatron/core/transformer/attention.py -@@ -1026,7 +1026,7 @@ class Attention(MegatronModule, ABC): - - return output, bias - -- @jit_fuser -+ - def _apply_output_gate(self, x, gate): - x_dtype = x.dtype - gate = gate.contiguous() -diff --git a/megatron/core/transformer/module.py b/megatron/core/transformer/module.py -index 2330df91b..0446c9097 100644 ---- a/megatron/core/transformer/module.py -+++ b/megatron/core/transformer/module.py -@@ -16,9 +16,9 @@ from megatron.core.transformer.utils import ( - sharded_state_dict_default, - ) - --_FLOAT_TYPES = (torch.FloatTensor, torch.cuda.FloatTensor) --_HALF_TYPES = (torch.HalfTensor, torch.cuda.HalfTensor) --_BF16_TYPES = (torch.BFloat16Tensor, torch.cuda.BFloat16Tensor) -+_FLOAT_TYPES = (torch.FloatTensor, torch.cuda.FloatTensor, torch.npu.FloatTensor) -+_HALF_TYPES = (torch.HalfTensor, torch.cuda.HalfTensor, torch.npu.HalfTensor) -+_BF16_TYPES = (torch.BFloat16Tensor, torch.cuda.BFloat16Tensor, torch.npu.BFloat16Tensor) - - - def param_is_not_shared(param): # pylint: disable=missing-function-docstring -diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py -index 5eeafdd8d..a4dce6970 100644 ---- a/megatron/core/transformer/moe/experts.py -+++ b/megatron/core/transformer/moe/experts.py -@@ -91,7 +91,7 @@ class GroupedMLP(MegatronModule): - if self.config.activation_func not in (F.silu, F.gelu): - raise ValueError("Activation function must be silu or gelu when using GroupedMLP.") - -- @jit_fuser -+ - def glu(x): - x = torch.chunk(x, 2, dim=-1) - return self.config.activation_func(x[0]) * x[1] -@@ -108,7 +108,7 @@ class GroupedMLP(MegatronModule): - "moe_act recompute for fp8 or fp4 cannot work with the legacy GroupedMLP." - ) - -- @jit_fuser -+ - def activation_func_with_probs(x, probs): - dtype = x.dtype - res = self.activation_func(x) * probs -diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py -index 517944f25..2c5bc0728 100644 ---- a/megatron/core/transformer/moe/router.py -+++ b/megatron/core/transformer/moe/router.py -@@ -472,7 +472,7 @@ class TopKRouter(Router): - else: - return input - -- @jit_fuser -+ - def _apply_expert_bias(self, routing_map: torch.Tensor): - """ - Update expert bias and tokens_per_expert -diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py -index d0da38d63..6d092c88f 100644 ---- a/megatron/core/transformer/moe/token_dispatcher.py -+++ b/megatron/core/transformer/moe/token_dispatcher.py -@@ -1403,7 +1403,7 @@ class MoEFlexTokenDispatcher(MoETokenDispatcher): - ).contiguous() - return routing_map, probs - -- @jit_fuser -+ - def dispatch_preprocess( - self, hidden_states: torch.Tensor, routing_map: torch.Tensor, probs: torch.Tensor - ): -diff --git a/megatron/core/transformer/torch_norm.py b/megatron/core/transformer/torch_norm.py -index d0ceca7af..f16796680 100644 ---- a/megatron/core/transformer/torch_norm.py -+++ b/megatron/core/transformer/torch_norm.py -@@ -69,7 +69,7 @@ class L2Norm(torch.nn.Module): - self.hidden_size = hidden_size - self.eps = eps - -- @jit_fuser -+ - def _norm(self, x): - """ - Performs the actual L2 normalization. -diff --git a/megatron/core/transformer/utils.py b/megatron/core/transformer/utils.py -index 880c53099..dbc95736e 100644 ---- a/megatron/core/transformer/utils.py -+++ b/megatron/core/transformer/utils.py -@@ -51,7 +51,7 @@ def attention_mask_func(attention_scores, attention_mask): - return attention_scores - - --@jit_fuser -+ - def gelu_impl(x): - """OpenAI's gelu implementation.""" - return 0.5 * x * (1.0 + torch.tanh(0.7978845608028654 * x * (1.0 + 0.044715 * x * x))) -@@ -65,7 +65,7 @@ def openai_gelu(x): - # This is actually Python equivalent of torch.nn.functional.gelu(), also with - # type hints for ONNX exporter - # pylint: disable=missing-function-docstring --@jit_fuser -+ - def erf_gelu(x): - return ( - x * 0.5 * (torch.erf(x / 1.41421).to(dtype=x.dtype) + torch.ones_like(x).to(dtype=x.dtype)) -diff --git a/megatron/legacy/model/fused_bias_gelu.py b/megatron/legacy/model/fused_bias_gelu.py -index e00e63148..ffe4b7ec6 100644 ---- a/megatron/legacy/model/fused_bias_gelu.py -+++ b/megatron/legacy/model/fused_bias_gelu.py -@@ -12,7 +12,7 @@ from megatron.core.jit import jit_fuser - # actual gelu is: - # x * 0.5 * (1.0 + torch.erf(x * 0.70710678)) - --@jit_fuser -+ - def bias_gelu(bias, y): - x = bias + y - return x * 0.5 * (1.0 + torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))) -@@ -20,7 +20,7 @@ def bias_gelu(bias, y): - # gradient of tanh approximation of gelu - # gradient of actual gelu is: - # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x) --@jit_fuser -+ - def bias_gelu_back(g, bias, y): - x = bias + y - tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x)) -diff --git a/megatron/legacy/model/transformer.py b/megatron/legacy/model/transformer.py -index 2a662a55b..2fc3e1bfb 100644 ---- a/megatron/legacy/model/transformer.py -+++ b/megatron/legacy/model/transformer.py -@@ -856,7 +856,7 @@ def get_bias_dropout_add(training): - return _bias_dropout_add - - --@jit_fuser -+ - def bias_dropout_add_fused_train(x: torch.Tensor, - bias: Optional[torch.Tensor], - residual: torch.Tensor, -@@ -864,7 +864,7 @@ def bias_dropout_add_fused_train(x: torch.Tensor, - return bias_dropout_add(x, bias, residual, prob, True) - - --@jit_fuser -+ - def bias_dropout_add_fused_inference(x: torch.Tensor, - bias: Optional[torch.Tensor], - residual: torch.Tensor, -diff --git a/megatron/legacy/model/utils.py b/megatron/legacy/model/utils.py -index 5762000d5..534858df7 100644 ---- a/megatron/legacy/model/utils.py -+++ b/megatron/legacy/model/utils.py -@@ -43,7 +43,7 @@ def get_linear_layer(rows, columns, init_method): - return layer - - --@jit_fuser -+ - def gelu_impl(x): - """OpenAI's gelu implementation.""" - return 0.5 * x * (1.0 + torch.tanh(0.7978845608028654 * x * -@@ -54,7 +54,7 @@ def openai_gelu(x): - - - #This is actually Python equivalent of torch.nn.functional.gelu(), also with type hints for ONNX exporter --@jit_fuser -+ - def erf_gelu(x): - return x * 0.5 * (torch.erf(x / 1.41421).to(dtype=x.dtype)+torch.ones_like(x).to(dtype=x.dtype)) - diff --git a/docker/npu_patch/mindspeed.patch b/docker/npu_patch/mindspeed.patch deleted file mode 100644 index 0cf1c8484..000000000 --- a/docker/npu_patch/mindspeed.patch +++ /dev/null @@ -1,48 +0,0 @@ -diff --git a/mindspeed/core/fusions/fused_rope.py b/mindspeed/core/fusions/fused_rope.py -index a6f02e07..70f7cb08 100644 ---- a/mindspeed/core/fusions/fused_rope.py -+++ b/mindspeed/core/fusions/fused_rope.py -@@ -126,5 +126,6 @@ def apply_rotary_pos_emb( - freqs, - rotary_interleaved=config.rotary_interleaved, - multi_latent_attention=config.multi_latent_attention, -- mscale=mscale -+ mscale=mscale, -+ cp_group=cp_group - ) -diff --git a/mindspeed/megatron_adaptor.py b/mindspeed/megatron_adaptor.py -index f14b231d..4615590e 100644 ---- a/mindspeed/megatron_adaptor.py -+++ b/mindspeed/megatron_adaptor.py -@@ -57,6 +57,7 @@ def delete_lock_file(): - def repatch(args): - MindSpeedFeaturesManager.remove_patches() - full_args = get_full_args() -+ args = vars(args) - for k, v in args.items(): - setattr(full_args, k, v) - MindSpeedFeaturesManager.apply_features_pre_patches(full_args) -diff --git a/mindspeed/te/pytorch/attention/dot_product_attention/dot_product_attention.py b/mindspeed/te/pytorch/attention/dot_product_attention/dot_product_attention.py -index ac4eabe5..78c5866d 100644 ---- a/mindspeed/te/pytorch/attention/dot_product_attention/dot_product_attention.py -+++ b/mindspeed/te/pytorch/attention/dot_product_attention/dot_product_attention.py -@@ -330,6 +330,8 @@ class DotProductAttention(torch.nn.Module): - inference_params: Any = None, - pad_between_seqs: Optional[bool] = None, - fp8_output: Optional[bool] = False, -+ local_cp_size=None, -+ cp_group=None, - ) -> torch.Tensor: - """ - Dot Product Attention Layer. -@@ -659,7 +661,9 @@ class MindSpeedTEDotProductAttention(DotProductAttention): - ) and not getattr(self.config, 'is_llava', False): - self.config.sparse_mode = 2 - attention_mask = get_attention_mask(self.config) -- -+ attention_mask = torch.triu( -+ torch.ones((2048, 2048), -+ device=query.device, dtype=torch.bool), diagonal=1) - packed_seq_kwargs = ( - {key: getattr(packed_seq_params, key) for key in self.kept_packed_seq_params} - if packed_seq_params is not None diff --git a/docker/npu_patch/sglang.patch b/docker/npu_patch/sglang.patch deleted file mode 100644 index b684d3c0e..000000000 --- a/docker/npu_patch/sglang.patch +++ /dev/null @@ -1,610 +0,0 @@ -diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py -index 8a343d43b..3e49dddb0 100644 ---- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py -+++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py -@@ -1050,93 +1050,36 @@ class AscendAttnBackend(AttentionBackend): - ) - - if not self.use_mla: -- num_tokens = q.shape[0] -- """PA will support bs torch.Tensor: -- qkv, _ = self.qkv_proj(hidden_states) -- q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) -- q, k = self.rotary_emb(positions, q, k) -+ if not _is_npu or not hasattr(self.rotary_emb, "get_cos_sin_with_position"): -+ q, k, v = self.forward_prepare_native( -+ positions=positions, -+ hidden_states=hidden_states, -+ ) -+ else: -+ q, k, v = self.forward_prepare_npu( -+ positions=positions, -+ hidden_states=hidden_states, -+ forward_batch=forward_batch, -+ ) -+ - attn_output = self.attn(q, k, v, forward_batch) - output, _ = self.o_proj(attn_output) - return output -diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py -index 47a1a4e4c..2cdbf8906 100644 ---- a/python/sglang/srt/models/qwen3.py -+++ b/python/sglang/srt/models/qwen3.py -@@ -161,12 +161,12 @@ class Qwen3Attention(nn.Module): - qkv, - self.rotary_emb.position_sin, - self.rotary_emb.position_cos, -- self.q_norm.weight, -- self.k_norm.weight, - self.q_size, - self.kv_size, - self.head_dim, -- self.q_norm.variance_epsilon, -+ eps=self.q_norm.variance_epsilon, -+ q_weight=self.q_norm.weight, -+ k_weight=self.k_norm.weight, - q_bias=getattr(self.q_norm, "bias", None), - k_bias=getattr(self.k_norm, "bias", None), - ) -@@ -370,6 +370,7 @@ class Qwen3ForCausalLM(nn.Module): - config.vocab_size, - config.hidden_size, - quant_config=quant_config, -+ use_attn_tp_group=get_global_server_args().enable_dp_lm_head, - prefix=add_prefix("lm_head", prefix), - ) - else: -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index e277d46f2..9e49156da 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -549,12 +549,12 @@ class Qwen3MoeAttention(nn.Module): - qkv, - self.rotary_emb.position_sin, - self.rotary_emb.position_cos, -- self.q_norm.weight, -- self.k_norm.weight, - self.q_size, - self.kv_size, - self.head_dim, -- self.q_norm.variance_epsilon, -+ eps=self.q_norm.variance_epsilon, -+ q_weight=self.q_norm.weight, -+ k_weight=self.k_norm.weight, - q_bias=getattr(self.q_norm, "bias", None), - k_bias=getattr(self.k_norm, "bias", None), - ) -diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py -index 218e32362..9f3141ce2 100644 ---- a/python/sglang/srt/models/qwen3_vl.py -+++ b/python/sglang/srt/models/qwen3_vl.py -@@ -19,6 +19,7 @@ import re - from functools import lru_cache, partial - from typing import Callable, Iterable, List, Optional, Tuple, Union - -+import numpy as np - import torch - import torch.nn as nn - from einops import rearrange -@@ -397,70 +398,89 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): - return cos_combined, sin_combined - - def fast_pos_embed_interpolate(self, grid_thw): -- grid_ts, grid_hs, grid_ws = grid_thw[:, 0], grid_thw[:, 1], grid_thw[:, 2] - num_grid_per_side = int(self.num_position_embeddings**0.5) -- device = self.pos_embed.weight.device - - idx_list = [[] for _ in range(4)] - weight_list = [[] for _ in range(4)] - -- for t, h, w in zip(grid_ts, grid_hs, grid_ws): -- h_idxs = torch.linspace(0, num_grid_per_side - 1, h) -- w_idxs = torch.linspace(0, num_grid_per_side - 1, w) -+ # TODO: use torch instand of np -+ for t, h, w in grid_thw: -+ h_idxs = np.linspace(0, num_grid_per_side - 1, h) -+ w_idxs = np.linspace(0, num_grid_per_side - 1, w) - -- h_idxs_floor = h_idxs.int() -- w_idxs_floor = w_idxs.int() -- h_idxs_ceil = (h_idxs.int() + 1).clip(max=num_grid_per_side - 1) -- w_idxs_ceil = (w_idxs.int() + 1).clip(max=num_grid_per_side - 1) -+ h_idxs_floor = h_idxs.astype(int) -+ w_idxs_floor = w_idxs.astype(int) -+ h_idxs_ceil = (h_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1) -+ w_idxs_ceil = (w_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1) - - dh = h_idxs - h_idxs_floor - dw = w_idxs - w_idxs_floor - -- base_h = h_idxs_floor * num_grid_per_side -- base_h_ceil = h_idxs_ceil * num_grid_per_side -- -- indices = [ -- (base_h[None].T + w_idxs_floor[None]).flatten(), -- (base_h[None].T + w_idxs_ceil[None]).flatten(), -- (base_h_ceil[None].T + w_idxs_floor[None]).flatten(), -- (base_h_ceil[None].T + w_idxs_ceil[None]).flatten(), -- ] -+ idx_list[0].extend( -+ ((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_floor[None]) -+ .flatten() -+ .tolist() -+ * t -+ ) -+ idx_list[1].extend( -+ ((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_ceil[None]) -+ .flatten() -+ .tolist() -+ * t -+ ) -+ idx_list[2].extend( -+ ((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_floor[None]) -+ .flatten() -+ .tolist() -+ * t -+ ) -+ idx_list[3].extend( -+ ((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_ceil[None]) -+ .flatten() -+ .tolist() -+ * t -+ ) - -- weights = [ -- ((1 - dh)[None].T * (1 - dw)[None]).flatten(), -- ((1 - dh)[None].T * dw[None]).flatten(), -- (dh[None].T * (1 - dw)[None]).flatten(), -- (dh[None].T * dw[None]).flatten(), -- ] -+ weight_list[0].extend( -+ ((1 - dh)[None].T * (1 - dw)[None]).flatten().tolist() * t -+ ) -+ weight_list[1].extend(((1 - dh)[None].T * dw[None]).flatten().tolist() * t) -+ weight_list[2].extend((dh[None].T * (1 - dw)[None]).flatten().tolist() * t) -+ weight_list[3].extend((dh[None].T * dw[None]).flatten().tolist() * t) - -- for i in range(4): -- idx_list[i].extend(indices[i].tolist()) -- weight_list[i].extend(weights[i].tolist()) -+ device = self.pos_embed.weight.device -+ dtype = self.pos_embed.weight.dtype - -- idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=device) -- weight_tensor = torch.tensor( -- weight_list, dtype=self.pos_embed.weight.dtype, device=device -+ p0 = ( -+ self.pos_embed(torch.tensor(idx_list[0], dtype=torch.long, device=device)) -+ * torch.tensor(weight_list[0], dtype=dtype, device=device)[:, None] - ) -- pos_embeds = self.pos_embed(idx_tensor).to(device) * weight_tensor[:, :, None] -- patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] -- -- patch_pos_embeds = patch_pos_embeds.split( -- [h * w for h, w in zip(grid_hs, grid_ws)] -+ p1 = ( -+ self.pos_embed(torch.tensor(idx_list[1], dtype=torch.long, device=device)) -+ * torch.tensor(weight_list[1], dtype=dtype, device=device)[:, None] -+ ) -+ p2 = ( -+ self.pos_embed(torch.tensor(idx_list[2], dtype=torch.long, device=device)) -+ * torch.tensor(weight_list[2], dtype=dtype, device=device)[:, None] -+ ) -+ p3 = ( -+ self.pos_embed(torch.tensor(idx_list[3], dtype=torch.long, device=device)) -+ * torch.tensor(weight_list[3], dtype=dtype, device=device)[:, None] - ) - -+ patch_pos_embeds = p0 + p1 + p2 + p3 -+ patch_pos_embeds = patch_pos_embeds.split([t * h * w for t, h, w in grid_thw]) - patch_pos_embeds_permute = [] -- merge_size = self.spatial_merge_size -- for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): -- pos_embed = pos_embed.repeat(t, 1) -+ m_size = self.spatial_merge_size -+ for pos_embed, (t, h, w) in zip(patch_pos_embeds, grid_thw): - pos_embed = ( -- pos_embed.view( -- t, h // merge_size, merge_size, w // merge_size, merge_size, -1 -- ) -+ pos_embed.view(t, h // m_size, m_size, w // m_size, m_size, -1) - .permute(0, 1, 3, 2, 4, 5) - .flatten(0, 4) - ) - patch_pos_embeds_permute.append(pos_embed) -- return torch.cat(patch_pos_embeds_permute) -+ patch_pos_embeds = torch.cat(patch_pos_embeds_permute) -+ return patch_pos_embeds - - def forward( - self, -@@ -650,31 +670,20 @@ class Qwen3LLMModel(Qwen3Model): - hidden_states + residual if residual is not None else hidden_states - ) - -- deepstack_embeds = None -- if input_deepstack_embeds is not None: -- prev_layer_idx = layer_idx - 1 -- if prev_layer_idx in self.deepstack_embed_to_decoder_layer: -- sep = self.hidden_size * prev_layer_idx -- deepstack_embeds = input_deepstack_embeds[ -- :, sep : sep + self.hidden_size -- ] -- -- # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. -- # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 -- # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack -- # The order matters because addition with different tensors is not associative in practice. - hidden_states, residual = layer( - positions, - hidden_states, - forward_batch, - residual, -- post_residual_addition=deepstack_embeds, - ) -+ # process deepstack -+ if ( -+ input_deepstack_embeds is not None -+ and layer_idx in self.deepstack_embed_to_decoder_layer -+ ): - -- # Handle deepstack for the last processed layer if it exists. -- last_deepstack = self.get_deepstack_embeds( -- self.end_layer - 1, input_deepstack_embeds -- ) -+ sep = self.hidden_size * layer_idx -+ hidden_states += input_deepstack_embeds[:, sep : sep + self.hidden_size] - - if not self.pp_group.is_last_rank: - return PPProxyTensors( -@@ -688,9 +697,7 @@ class Qwen3LLMModel(Qwen3Model): - if residual is None: - hidden_states = self.norm(hidden_states) - else: -- hidden_states, _ = self.norm( -- hidden_states, residual, post_residual_addition=last_deepstack -- ) -+ hidden_states, _ = self.norm(hidden_states, residual) - - if len(aux_hidden_states) == 0: - return hidden_states -@@ -805,15 +812,15 @@ class Qwen3VLForConditionalGeneration(nn.Module): - max_images_per_call = get_int_env_var("SGLANG_VLM_MAX_IMAGES_PER_VIT", 0) - - if max_patches_per_call == 0 and max_images_per_call == 0: -- if self.use_data_parallel: -- return run_dp_sharded_mrope_vision_model( -- self.visual, -- pixel_values, -- image_grid_thw.tolist(), -- rope_type="rope_3d", -- ) -- else: -- return self.visual(pixel_values, grid_thw=image_grid_thw) -+ # if self.use_data_parallel: -+ # return run_dp_sharded_mrope_vision_model( -+ # self.visual, -+ # pixel_values, -+ # image_grid_thw.tolist(), -+ # rope_type="rope_3d", -+ # ) -+ # else: -+ return self.visual(pixel_values, grid_thw=image_grid_thw) - - # compute the number of patches per image and the slice positions in pixel_values - grid_thw_list = ( -@@ -995,7 +1002,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): - name = name.replace(r"model.language_model.", r"model.") - layer_id = get_layer_id(name) - -- if self.pp_group.is_last_rank and "model.embed_tokens.weight" in name: -+ if self.pp_group.is_last_rank and "model.embed_tokens.weight" in name and self.config.tie_word_embeddings: - if "lm_head.weight" in params_dict: - lm_head_param = params_dict["lm_head.weight"] - weight_loader = getattr( -diff --git a/scripts/ci/npu_ci_install_dependency.sh b/scripts/ci/npu_ci_install_dependency.sh -index 6172db3b4..8cc8fbf39 100755 ---- a/scripts/ci/npu_ci_install_dependency.sh -+++ b/scripts/ci/npu_ci_install_dependency.sh -@@ -49,7 +49,7 @@ wget -O "${BISHENG_NAME}" "${BISHENG_URL}" && chmod a+x "${BISHENG_NAME}" && "./ - - - ### Install sgl-kernel-npu --SGL_KERNEL_NPU_TAG="20251206" -+SGL_KERNEL_NPU_TAG="2025.12.31" - git clone --depth 1 https://github.com/sgl-project/sgl-kernel-npu.git --branch ${SGL_KERNEL_NPU_TAG} - (cd sgl-kernel-npu && bash ./build.sh && ${PIP_INSTALL} output/deep_ep*.whl output/sgl_kernel_npu*.whl && cd "$(python3 -m pip show deep-ep | grep -E '^Location:' | awk '{print $2}')" && ln -s deep_ep/deep_ep_cpp*.so) - diff --git a/docker/npu_patch/slime.patch b/docker/npu_patch/slime.patch deleted file mode 100644 index 55e7c2d44..000000000 --- a/docker/npu_patch/slime.patch +++ /dev/null @@ -1,941 +0,0 @@ -diff --git a/slime/backends/megatron_utils/__init__.py b/slime/backends/megatron_utils/__init__.py -index a4666fbe..2a07086a 100644 ---- a/slime/backends/megatron_utils/__init__.py -+++ b/slime/backends/megatron_utils/__init__.py -@@ -21,21 +21,35 @@ except ImportError: - logging.warning("deep_ep is not installed, some functionalities may be limited.") - - try: -- from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import ( -- Qwen3VLMoETextRotaryEmbedding, -- Qwen3VLTextRotaryEmbedding, -- ) -- -- def patch_rotary_embedding(cls): -- _original_forward = cls.forward -- -- def _patched_forward(self, *args, packed_seq_params=None, **kwargs): -- return _original_forward(self, *args, **kwargs) -+ from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.text_model import Qwen3VLTextRotaryEmbedding, Qwen3VLMoETextRotaryEmbedding -+ -+ _original_forward = Qwen3VLTextRotaryEmbedding.forward -+ _original_forward_1 = Qwen3VLMoETextRotaryEmbedding.forward -+ -+ def _patched_forward(self, *args, packed_seq_params=None, **kwargs): -+ return _original_forward(self, *args, **kwargs) -+ def _patched_forward_1(self, *args, packed_seq_params=None, **kwargs): -+ return _original_forward_1(self, *args, **kwargs) -+ Qwen3VLTextRotaryEmbedding.forward = _patched_forward -+ Qwen3VLMoETextRotaryEmbedding.forward = _patched_forward_1 -+except ImportError: -+ pass - -- cls.forward = _patched_forward -+try: -+ from mbridge.models.qwen3_vl.model import Qwen3VLModel -+ _original_forward2 = Qwen3VLModel.forward -+ def _patched_forward2(self, *args, loss_mask=None, **kwargs): -+ return _original_forward2(self, *args, **kwargs) -+ Qwen3VLModel.forward = _patched_forward2 -+except ImportError: -+ pass - -- patch_rotary_embedding(Qwen3VLTextRotaryEmbedding) -- patch_rotary_embedding(Qwen3VLMoETextRotaryEmbedding) -+try: -+ from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.model import Qwen3VLModel -+ _original_forward3 = Qwen3VLModel.forward -+ def _patched_forward3(self, *args, loss_mask=None, **kwargs): -+ return _original_forward3(self, *args, **kwargs) -+ Qwen3VLModel.forward = _patched_forward3 - except ImportError: - pass - -diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py -index 7bc4f910..a0b1f63a 100644 ---- a/slime/backends/megatron_utils/actor.py -+++ b/slime/backends/megatron_utils/actor.py -@@ -8,6 +8,10 @@ from contextlib import nullcontext - import ray - import torch - import torch.distributed as dist -+from slime.utils.common import is_npu -+if is_npu(): -+ import mindspeed.megatron_adaptor -+ from mindspeed.megatron_adaptor import repatch - from megatron.core import mpu - from ray.actor import ActorHandle - from torch_memory_saver import torch_memory_saver -@@ -55,6 +59,8 @@ class MegatronTrainRayActor(TrainRayActor): - super().init(args, role, with_ref) - - init(args) -+ if is_npu(): -+ repatch(args) - - if is_megatron_main_rank(): - init_tracking(args, primary=False) -@@ -596,8 +602,12 @@ class MegatronTrainRayActor(TrainRayActor): - - group_name = "actor_critic" - world_size = 2 -+ if is_npu(): -+ backend = "hccl" -+ else: -+ backend = "nccl" - self._actor_critic_groups = init_process_group( -- backend="nccl", -+ backend=backend, - init_method=f"tcp://{master_address}:{master_port}", - world_size=world_size, - rank=0 if self.role == "actor" else 1, -diff --git a/slime/backends/megatron_utils/megatron_to_hf/__init__.py b/slime/backends/megatron_utils/megatron_to_hf/__init__.py -index 84ff899a..a0e2e43c 100644 ---- a/slime/backends/megatron_utils/megatron_to_hf/__init__.py -+++ b/slime/backends/megatron_utils/megatron_to_hf/__init__.py -@@ -7,7 +7,7 @@ from .processors import quantize_params, remove_padding - from .qwen2 import convert_qwen2_to_hf - from .qwen3_next import convert_qwen3_next_to_hf - from .qwen3moe import convert_qwen3moe_to_hf -- -+from .qwen3_vl import convert_qwen3vl_to_hf - - # TODO unify w/ `convert_to_hf` - def postprocess_hf_param(args, megatron_param_name, hf_param_name, param): -@@ -40,7 +40,10 @@ def _convert_to_hf_core(args, model_name, name, param): - elif "qwen3next" in model_name: - converted_named_tensors = convert_qwen3_next_to_hf(args, name, param) - elif "qwen2" in model_name or "qwen3" in model_name: -- converted_named_tensors = convert_qwen2_to_hf(args, name, param) -+ if "qwen3vl" in model_name: -+ converted_named_tensors = convert_qwen3vl_to_hf(args, name, param) -+ else: -+ converted_named_tensors = convert_qwen2_to_hf(args, name, param) - elif "deepseekv3" in model_name: - converted_named_tensors = convert_deepseekv3_to_hf(args, name, param) - -diff --git a/slime/backends/megatron_utils/megatron_to_hf/qwen3_vl.py b/slime/backends/megatron_utils/megatron_to_hf/qwen3_vl.py -new file mode 100644 -index 00000000..5a3aa56d ---- /dev/null -+++ b/slime/backends/megatron_utils/megatron_to_hf/qwen3_vl.py -@@ -0,0 +1,133 @@ -+import re -+import torch -+from megatron.core import parallel_state as mpu -+ -+ -+def convert_qwen3vl_to_hf(args, name, param): -+ """ -+ Convert Megatron-style Qwen3-VL parameter names to HF-style names. -+ Supports both language model and vision model parameters. -+ -+ Args: -+ args: megatron model args (num_attention_heads, kv_channels, etc.) -+ name: str, Megatron parameter name -+ param: torch.Tensor, parameter value -+ -+ Returns: -+ List of tuples [(hf_name, hf_param), ...] -+ """ -+ -+ hf_name_param = None -+ -+ # ---------------------------- -+ # 1. language model & vision model parameters -+ # ---------------------------- -+ -+ try: -+ head_dim = args.kv_channels if args.kv_channels is not None else args.hidden_size // args.num_attention_heads -+ except: -+ head_dim = args.hidden_size // args.num_attention_heads -+ value_num_per_group = args.num_attention_heads // args.num_query_groups -+ language_num_layers = args.num_layers -+ -+ pp_size = args.pipeline_model_parallel_size -+ pp_rank = mpu.get_pipeline_model_parallel_rank() -+ assert language_num_layers % pp_size == 0 -+ -+ num_layers_per_rank = language_num_layers // pp_size -+ offsets = pp_rank * num_layers_per_rank -+ -+ -+ -+ # ---------------------------- -+ # 2. LM Embeddings & output -+ # ---------------------------- -+ if name == "module.module.language_model.embedding.word_embeddings.weight": -+ hf_name_param = [("model.language_model.embed_tokens.weight", param)] -+ elif name == "module.module.language_model.decoder.final_layernorm.weight": -+ hf_name_param = [("model.language_model.norm.weight", param)] -+ elif name == "module.module.language_model.output_layer.weight": -+ if not args.untie_embeddings_and_output_weights: -+ return [("model.language_model.embed_tokens.weight", param)] -+ else: -+ return [("lm_head.weight", param)] -+ -+ else: -+ -+ decoder_layers_pattern = r"module\.module\.language_model.decoder\.layers\.(\d+)\.(.+)" -+ vision_pattern = r"module\.module\.vision_model\.(.+)" -+ -+ if match := re.match(decoder_layers_pattern, name): -+ # ---------------------------- -+ # 3. Attention and MLP layers in language model -+ # ---------------------------- -+ -+ layer_idx, rest = match.groups() -+ layer_idx = str(int(layer_idx) + offsets) -+ # Self-attention projection -+ if rest == "self_attention.linear_proj.weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.self_attn.o_proj.weight", param)] -+ -+ elif rest == "self_attention.linear_qkv.weight": -+ param = param.view(args.num_query_groups, -1, head_dim, args.hidden_size) -+ q_param, k_param, v_param = torch.split( -+ param, split_size_or_sections=[value_num_per_group, 1, 1], dim=1 -+ ) -+ q_param = q_param.reshape(-1, args.hidden_size) -+ k_param = k_param.reshape(-1, args.hidden_size) -+ v_param = v_param.reshape(-1, args.hidden_size) -+ hf_name_param = [ -+ (f"model.language_model.layers.{layer_idx}.self_attn.q_proj.weight", q_param), -+ (f"model.language_model.layers.{layer_idx}.self_attn.k_proj.weight", k_param), -+ (f"model.language_model.layers.{layer_idx}.self_attn.v_proj.weight", v_param), -+ ] -+ -+ # MLP layers -+ elif rest == "mlp.linear_fc1.weight": -+ gate_weight, up_weight = param.chunk(2, dim=0) -+ -+ hf_name_param = [ -+ (f"model.language_model.layers.{layer_idx}.mlp.gate_proj.weight", gate_weight), -+ (f"model.language_model.layers.{layer_idx}.mlp.up_proj.weight", up_weight), -+ ] -+ -+ elif rest == "mlp.linear_fc2.weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.mlp.down_proj.weight", param)] -+ -+ # LayerNorms -+ elif rest == "self_attention.linear_qkv.layer_norm_weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.input_layernorm.weight", param)] -+ elif rest == "mlp.linear_fc1.layer_norm_weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.post_attention_layernorm.weight", param)] -+ elif rest == "self_attention.q_layernorm.weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.self_attn.q_norm.weight", param)] -+ elif rest == "self_attention.k_layernorm.weight": -+ hf_name_param = [(f"model.language_model.layers.{layer_idx}.self_attn.k_norm.weight", param)] -+ -+ elif match_v := re.match(vision_pattern, name): -+ # ---------------------------- -+ # 4. Vision model parameters -+ # ---------------------------- -+ deepstack_merger_pattern = r"deepstack_merger_list\.(\d+)\.(.+)" -+ decoder_layer_pattern = r"blocks\.(\d+)\.(.+)" -+ -+ rest = match_v.groups()[0] -+ -+ if match_layer := re.match(deepstack_merger_pattern, rest): -+ layer_idx, layer_rest = match_layer.groups() -+ layer_idx = str(int(layer_idx) + offsets) -+ hf_name_param = [(f"model.visual.deepstack_merger_list.{layer_idx}.{layer_rest}", param)] -+ -+ # Decoder layers -+ elif match_layer := re.match(decoder_layer_pattern, rest): -+ -+ layer_idx, layer_rest = match_layer.groups() -+ layer_idx = str(int(layer_idx) + offsets) -+ hf_name_param = [(f"model.visual.blocks.{layer_idx}.{layer_rest}", param)] -+ else: -+ hf_name_param = [(f"model.visual.{rest}", param)] -+ -+ if hf_name_param == None: -+ raise ValueError(f"Unknown parameter name: {name}") -+ else: -+ return hf_name_param -diff --git a/slime/backends/megatron_utils/model_provider.py b/slime/backends/megatron_utils/model_provider.py -index 8174c7ac..1f1c0e09 100644 ---- a/slime/backends/megatron_utils/model_provider.py -+++ b/slime/backends/megatron_utils/model_provider.py -@@ -17,7 +17,7 @@ from megatron.core.transformer.transformer_config import TransformerConfig - from megatron.training.arguments import core_transformer_config_from_args - - from slime.utils.misc import load_function -- -+import slime_plugins.patch.mbridge_patch - - # Adapt from https://github.com/volcengine/verl/blob/c3b20575d2bc815fcccd84bddb4c0401fc4b632b/verl/models/llama/megatron/layers/parallel_linear.py#L82 - class LinearForLastLayer(torch.nn.Linear): -@@ -33,7 +33,7 @@ class LinearForLastLayer(torch.nn.Linear): - self.sequence_parallel = config.sequence_parallel - if self.sequence_parallel: - self.weight.sequence_parallel = True -- -+ self.bias.sequence_parallel = True - self.weight.data.normal_(mean=0.0, std=0.02) - if bias: - self.bias.data.zero_() -@@ -51,10 +51,59 @@ class LinearForLastLayer(torch.nn.Linear): - return logits, None - - -+def get_qwen3vl_provide_wrapper(provider, role): -+ -+ def provide_wrapper(pre_process=None, post_process=None, vp_stage=None): -+ """ -+ Provide a Qwen3VL MoE model instance with vision and language components. -+ """ -+ from megatron.bridge.models.qwen_vl.modelling_qwen3_vl.model import Qwen3VLModel -+ language_transformer_config = provider -+ -+ # Create vision transformer config - placeholder for future use -+ hf_config = provider.vision_config -+ -+ language_transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( -+ num_experts=provider.num_moe_experts, -+ moe_grouped_gemm=True, -+ qk_layernorm=provider.qk_layernorm, -+ fp8=False, -+ normalization="RMSNorm", -+ ) -+ -+ # reuse Qwen3VLModel for MoE model but replace the language model with MoE language model -+ model = Qwen3VLModel( -+ language_transformer_config=language_transformer_config, -+ language_transformer_layer_spec=language_transformer_layer_spec, -+ vision_transformer_config=hf_config, -+ pre_process=pre_process, -+ post_process=post_process, -+ ) -+ -+ if role == "critic" and post_process: -+ model.language_model.output_layer = LinearForLastLayer(input_size=provider.hidden_size, output_size=1, config=provider).to( -+ device=model.language_model.output_layer.weight.device, -+ dtype=model.language_model.output_layer.weight.dtype, -+ ) -+ -+ # Apply freeze options if any are enabled for fine-tuning -+ if provider.freeze_language_model or provider.freeze_vision_model or provider.freeze_vision_projection: -+ model.freeze( -+ freeze_language_model=provider.freeze_language_model, -+ freeze_vision_model=provider.freeze_vision_model, -+ freeze_vision_projection=provider.freeze_vision_projection, -+ ) -+ -+ return model -+ return provide_wrapper -+ -+ - def get_model_provider_func( - args: argparse.Namespace, - role: Literal["actor", "critic"] = "actor", - ): -+ from megatron.bridge.models.conversion.param_mapping import AutoMapping -+ AutoMapping.register_module_type('LinearForLastLayer', 'replicated') # 或 'column' / 'replicated' - # Support custom model provider path (similar to --custom-rm-path for reward models) - if getattr(args, "custom_model_provider_path", None): - -@@ -88,6 +137,27 @@ def get_model_provider_func( - provider.expert_model_parallel_size = args.expert_model_parallel_size - provider.expert_tensor_parallel_size = args.expert_tensor_parallel_size - provider.sequence_parallel = args.sequence_parallel -+ provider.gradient_accumulation_fusion = args.gradient_accumulation_fusion -+ provider.recompute_granularity = args.recompute_granularity -+ provider.recompute_method = args.recompute_method -+ provider.recompute_num_layers = args.recompute_num_layers -+ for key, value in vars(args).items(): -+ if hasattr(provider, key): -+ continue -+ setattr(provider, key, value) -+ -+ is_qwen3vl = ( -+ hasattr(bridge.hf_pretrained, 'config') -+ and hasattr(bridge.hf_pretrained.config, 'model_type') -+ and 'qwen3_vl' in bridge.hf_pretrained.config.model_type.lower() -+ ) -+ -+ if role == 'critic' and is_qwen3vl: -+ from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge -+ from slime_plugins.patch.critic_patch import load_weights_hf_to_megatron_wrapper -+ MegatronModelBridge.load_weights_hf_to_megatron = load_weights_hf_to_megatron_wrapper -+ provider.provide = get_qwen3vl_provide_wrapper(provider, role) -+ - provider.finalize() - return provider.provide - -diff --git a/slime/backends/megatron_utils/update_weight/common.py b/slime/backends/megatron_utils/update_weight/common.py -index a2e4e129..ce40b446 100644 ---- a/slime/backends/megatron_utils/update_weight/common.py -+++ b/slime/backends/megatron_utils/update_weight/common.py -@@ -10,7 +10,7 @@ from megatron.core.transformer.transformer_layer import get_transformer_layer_of - - from slime.backends.megatron_utils.misc_utils import strip_param_name_prefix - from slime.utils.types import ParamInfo -- -+from slime.utils.common import is_npu - - def all_gather_param(name: str, param: torch.nn.Parameter) -> torch.Tensor: - """ -@@ -40,6 +40,8 @@ def all_gather_param(name: str, param: torch.nn.Parameter) -> torch.Tensor: - if "linear_fc1.weight" in name: - param_partitions = [p.chunk(2, dim=0) for p in param_partitions] - param_partitions = [p[0] for p in param_partitions] + [p[1] for p in param_partitions] -+ if is_npu(): -+ partition_dim = 0 - # this is bug in megatron's grouped moe. - if "linear_fc2.weight" in name: - if partition_dim == 0: -@@ -102,6 +104,8 @@ def all_gather_params_async( - if "linear_fc1.weight" in info.name: - param_partitions = [p.chunk(2, dim=0) for p in param_partitions] - param_partitions = [p[0] for p in param_partitions] + [p[1] for p in param_partitions] -+ if is_npu(): -+ partition_dim = 0 - # this is bug in megatron's grouped moe. - if "linear_fc2.weight" in info.name: - if partition_dim == 0: -diff --git a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py -index a8e50e0e..b3c6ac24 100644 ---- a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py -+++ b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py -@@ -12,6 +12,7 @@ from ray.actor import ActorHandle - from tqdm import tqdm - - from slime.utils.distributed_utils import get_gloo_group, init_process_group -+from slime.utils.common import is_npu - - from ..megatron_to_hf import convert_to_hf - from .common import all_gather_param, named_params_and_buffers -@@ -253,6 +254,7 @@ def connect_rollout_engines_from_distributed( - master_port = sock.getsockname()[1] - world_size = len(rollout_engines) * args.rollout_num_gpus_per_engine + 1 - -+ backend = "hccl" if is_npu() else "nccl" - refs = [ - engine.init_weights_update_group.remote( - master_address, -@@ -260,12 +262,12 @@ def connect_rollout_engines_from_distributed( - i * args.rollout_num_gpus_per_engine + 1, - world_size, - group_name, -- backend="nccl", -+ backend=backend, - ) - for i, engine in enumerate(rollout_engines) - ] - model_update_groups = init_process_group( -- backend="nccl", -+ backend=backend, - init_method=f"tcp://{master_address}:{master_port}", - world_size=world_size, - rank=0, -diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py -index af0f5620..03ded2c8 100644 ---- a/slime/backends/sglang_utils/sglang_engine.py -+++ b/slime/backends/sglang_utils/sglang_engine.py -@@ -15,6 +15,7 @@ from urllib3.exceptions import NewConnectionError - - from slime.ray.ray_actor import RayActor - from slime.utils.http_utils import get_host_info -+from slime.utils.common import is_npu - - logger = logging.getLogger(__name__) - -@@ -33,7 +34,10 @@ def get_base_gpu_id(args, rank): - - - def _to_local_gpu_id(physical_gpu_id: int) -> int: -- cvd = os.environ.get("CUDA_VISIBLE_DEVICES") -+ if is_npu(): -+ cvd = os.environ.get("ASCEND_RT_VISIBLE_DEVICES") -+ else: -+ cvd = os.environ.get("CUDA_VISIBLE_DEVICES") - if not cvd: - return physical_gpu_id # no remapping - # CUDA_VISIBLE_DEVICES can be like "4,5,6,7" -diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py -index 4cac07a2..098bdf92 100644 ---- a/slime/ray/actor_group.py -+++ b/slime/ray/actor_group.py -@@ -5,6 +5,7 @@ from ray.util.placement_group import PlacementGroup - from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy - - from slime.ray.utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST -+from slime.utils.common import is_npu - - - class RayTrainGroup: -@@ -87,19 +88,19 @@ class RayTrainGroup: - - actor_impl = FSDPTrainRayActor - -- TrainRayActor = ray.remote(num_gpus=1, runtime_env={"env_vars": env_vars})(actor_impl) -- -+ TrainRayActor = ray.remote(runtime_env={"env_vars": env_vars})(actor_impl) -+ device_name = "NPU" if is_npu() else "GPU" - # Create worker actors - self._actor_handlers = [] - master_addr, master_port = None, None - for rank in range(world_size): - actor = TrainRayActor.options( - num_cpus=num_gpus_per_actor, -- num_gpus=num_gpus_per_actor, - scheduling_strategy=PlacementGroupSchedulingStrategy( - placement_group=pg, - placement_group_bundle_index=reordered_bundle_indices[rank], - ), -+ resources={device_name:num_gpus_per_actor} - ).remote(world_size, rank, master_addr, master_port) - if rank == 0: - master_addr, master_port = ray.get(actor.get_master_addr_and_port.remote()) -diff --git a/slime/ray/placement_group.py b/slime/ray/placement_group.py -index eb232b16..963b4071 100644 ---- a/slime/ray/placement_group.py -+++ b/slime/ray/placement_group.py -@@ -4,6 +4,7 @@ import socket - import ray - from ray.util.placement_group import placement_group - from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy -+from slime.utils.common import is_npu - - from .actor_group import RayTrainGroup - from .rollout import RolloutManager -@@ -11,10 +12,13 @@ from .rollout import RolloutManager - logger = logging.getLogger(__name__) - - --@ray.remote(num_gpus=1) -+@ray.remote - class InfoActor: - def get_ip_and_gpu_id(self): -- return ray.util.get_node_ip_address(), ray.get_gpu_ids()[0] -+ if is_npu(): -+ return ray.util.get_node_ip_address(), ray.get_runtime_context().get_accelerator_ids()["NPU"][0] -+ else: -+ return ray.util.get_node_ip_address(), ray.get_gpu_ids()[0] - - - def sort_key(x): -@@ -35,12 +39,13 @@ def sort_key(x): - # representation that allows for sorting. - node_ip_parts = [ord(c) for c in node_identifier] - -- return (node_ip_parts, gpu_id) -+ return (node_ip_parts, int(gpu_id)) - - - def _create_placement_group(num_gpus): - """Create a placement group with the specified number of GPUs.""" -- bundles = [{"GPU": 1, "CPU": 1} for _ in range(num_gpus)] -+ device_name = "NPU" if is_npu() else "GPU" -+ bundles = [{device_name: 1, "CPU": 1} for _ in range(num_gpus)] - pg = placement_group(bundles, strategy="PACK") - num_bundles = len(bundles) - -@@ -53,7 +58,8 @@ def _create_placement_group(num_gpus): - scheduling_strategy=PlacementGroupSchedulingStrategy( - placement_group=pg, - placement_group_bundle_index=i, -- ) -+ ), -+ resources={device_name:1} - ).remote() - ) - gpu_ids = ray.get([actor.get_ip_and_gpu_id.remote() for actor in info_actors]) -@@ -167,9 +173,11 @@ def create_training_models(args, pgs, rollout_manager): - - - def create_rollout_manager(args, pg): -+ device_name = "NPU" if is_npu() else "GPU" - rollout_manager = RolloutManager.options( - num_cpus=1, - num_gpus=0, -+ resources={device_name:0} - ).remote(args, pg) - - # calculate num_rollout from num_epoch -diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py -index 75cb053c..c54a3854 100644 ---- a/slime/ray/rollout.py -+++ b/slime/ray/rollout.py -@@ -28,6 +28,7 @@ from slime.utils.metric_utils import ( - from slime.utils.misc import Box, group_by, load_function - from slime.utils.seqlen_balancing import get_seqlen_balanced_partitions - from slime.utils.types import Sample -+from slime.utils.common import is_npu - - from ..utils.metric_utils import has_repetition - from .utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST, Lock -@@ -76,7 +77,8 @@ class RolloutManager: - self.all_rollout_engines = [None] * num_engines - self.num_new_engines = init_rollout_engines(args, pg, self.all_rollout_engines) - self.nodes_per_engine = max(1, args.rollout_num_gpus_per_engine // args.num_gpus_per_node) -- self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0).remote() -+ device_name = "NPU" if is_npu() else "GPU" -+ self.rollout_engine_lock = Lock.options(num_cpus=1, num_gpus=0, resources={device_name:0}).remote() - self.rollout_id = -1 - - self._metric_checker = MetricChecker.maybe_create(args) -@@ -467,6 +469,7 @@ def init_rollout_engines(args, pg, all_rollout_engines): - RolloutRayActor = ray.remote(SGLangEngine) - - rollout_engines = [] -+ device_name = "NPU" if is_npu() else "GPU" - for i in range(num_engines): - if all_rollout_engines[i] is not None: - continue -@@ -503,11 +506,11 @@ def init_rollout_engines(args, pg, all_rollout_engines): - - rollout_engine = RolloutRayActor.options( - num_cpus=num_cpus, -- num_gpus=num_gpus, - scheduling_strategy=scheduling_strategy, - runtime_env={ - "env_vars": env_vars, - }, -+ resources={device_name:num_gpus} - ).remote(args, rank=i, worker_type=worker_type, base_gpu_id=base_gpu_id) - - rollout_engines.append((i, rollout_engine)) -diff --git a/slime/ray/train_actor.py b/slime/ray/train_actor.py -index 2e900ca5..d0a25583 100644 ---- a/slime/ray/train_actor.py -+++ b/slime/ray/train_actor.py -@@ -13,16 +13,23 @@ from slime.ray.ray_actor import RayActor - from slime.utils.distributed_utils import init_gloo_group - from slime.utils.logging_utils import configure_logger - from slime.utils.memory_utils import clear_memory, print_memory -+from slime.utils.common import is_npu - - logger = logging.getLogger(__name__) - - - def get_local_gpu_id(): -- cvd = os.environ.get("CUDA_VISIBLE_DEVICES", None) -+ if is_npu(): -+ env_var = "ASCEND_RT_VISIBLE_DEVICES" -+ device_ids = ray.get_runtime_context().get_accelerator_ids()["NPU"] -+ else: -+ env_var = "CUDA_VISIBLE_DEVICES" -+ device_ids = ray.get_gpu_ids() -+ cvd = os.environ.get(env_var, None) - if cvd is None: -- return ray.get_gpu_ids()[0] -+ return device_ids[0] - else: -- return cvd.split(",").index(str(ray.get_gpu_ids()[0])) -+ return cvd.split(",").index(str(device_ids[0])) - - - class TrainRayActor(RayActor): -diff --git a/slime/utils/common.py b/slime/utils/common.py -new file mode 100644 -index 00000000..60ee5b5c ---- /dev/null -+++ b/slime/utils/common.py -@@ -0,0 +1,12 @@ -+import torch -+ -+def is_npu() -> bool: -+ if not hasattr(torch, "npu"): -+ return False -+ -+ if not torch.npu.is_available(): -+ raise RuntimeError( -+ "torch_npu detected, but NPU device is not available or visible." -+ ) -+ -+ return True -diff --git a/slime/utils/external_utils/command_utils.py b/slime/utils/external_utils/command_utils.py -index 9f51ecdf..d4b47eca 100644 ---- a/slime/utils/external_utils/command_utils.py -+++ b/slime/utils/external_utils/command_utils.py -@@ -184,6 +184,108 @@ def execute_train( - f"{train_args}" - ) - -+def execute_train_npu( -+ train_args: str, -+ megatron_model_type: str | None, -+ train_script: str = "train.py", -+ before_ray_job_submit=None, -+ extra_env_vars=None, -+ config: ExecuteTrainConfig | None = None, -+): -+ if extra_env_vars is None: -+ extra_env_vars = {} -+ if config is None: -+ config = ExecuteTrainConfig() -+ external_ray = get_bool_env_var("SLIME_SCRIPT_EXTERNAL_RAY") -+ master_addr = os.environ.get("MASTER_ADDR", "127.0.0.1") -+ -+ train_backend_fsdp = "--train-backend fsdp" in train_args -+ assert train_backend_fsdp == (megatron_model_type is None) -+ -+ exec_command( -+ "pkill -9 sglang; " -+ "sleep 3; " -+ f"{'' if external_ray else 'ray stop --force; '}" -+ f"{'' if external_ray else 'pkill -9 ray; '}" -+ # cannot be run in CI, o/w kill the parent script -+ # TODO: do we really need this kill? (or can we instead kill slime) -+ # "pkill -9 python; " -+ "pkill -9 slime; " -+ "sleep 3; " -+ f"{'' if external_ray else 'pkill -9 ray; '}" -+ # "pkill -9 python; " -+ "pkill -9 slime; " -+ "pkill -9 redis; " -+ "true; " -+ ) -+ -+ if not external_ray: -+ exec_command( -+ # will prevent ray from buffering stdout/stderr -+ f"export PYTHONBUFFERED=16 && " -+ f"ray start --head --node-ip-address {master_addr} --disable-usage-stats" -+ ) -+ -+ if (f := before_ray_job_submit) is not None: -+ f() -+ -+ runtime_env_json = json.dumps( -+ { -+ "env_vars": { -+ # "PYTHONPATH": "/root/Megatron-LM/", -+ "CUDA_DEVICE_MAX_CONNECTIONS": "1", -+ "RAY_EXPERIMENTAL_NOSET_ASCEND_RT_VISIBLE_DEVICES": "1", -+ "ASCEND_TOOLKIT_HOME": "/path/to/ascend/Ascend/ascend-toolkit/latest/", -+ "ASCEND_OPP_PATH": "/path/to/ascend/Ascend/ascend-toolkit/latest/opp/", -+ "ASCEND_AICPU_PATH": "/path/to/ascend/Ascend/ascend-toolkit/latest/", -+ "ASCEND_HOME_PATH": "/path/to/ascend/Ascend/ascend-toolkit/latest/", -+ "set_env_path": "/path/to/ascend/Ascend/nnal/atb/set_env.sh", -+ "HYDRA_FULL_ERROR": "1", -+ "HCCL_HOST_SOCKET_PORT_RANGE": "60000-60050", -+ "HCCL_NPU_SOCKET_PORT_RANGE": "61000-61050", -+ # If setting this in FSDP, the computation communication overlapping may have issues -+ **( -+ {} -+ if train_backend_fsdp -+ else { -+ "CUDA_DEVICE_MAX_CONNECTIONS": "1", -+ } -+ ), -+ "NCCL_NVLS_ENABLE": str(int(check_has_nvlink())), -+ "no_proxy": f"127.0.0.1,{master_addr}", -+ # This is needed by megatron / torch distributed in multi-node setup -+ "MASTER_ADDR": master_addr, -+ **( -+ { -+ "CUDA_ENABLE_COREDUMP_ON_EXCEPTION": "1", -+ "CUDA_COREDUMP_SHOW_PROGRESS": "1", -+ "CUDA_COREDUMP_GENERATION_FLAGS": "skip_nonrelocated_elf_images,skip_global_memory,skip_shared_memory,skip_local_memory,skip_constbank_memory", -+ "CUDA_COREDUMP_FILE": "/root/shared_data/cuda_coredump_%h.%p.%t", -+ } -+ if config.cuda_core_dump -+ else {} -+ ), -+ **extra_env_vars, -+ **_parse_extra_env_vars(config.extra_env_vars), -+ } -+ } -+ ) -+ -+ if get_bool_env_var("SLIME_SCRIPT_ENABLE_RAY_SUBMIT", "1"): -+ cmd_megatron_model_source = ( -+ f'source "{repo_base_dir}/scripts/models/{megatron_model_type}.sh" && ' -+ if megatron_model_type is not None -+ else "" -+ ) -+ exec_command( -+ f"export no_proxy=127.0.0.1 && export PYTHONBUFFERED=16 && " -+ f"{cmd_megatron_model_source}" -+ f'ray job submit --address="http://127.0.0.1:8265" ' -+ f"--runtime-env-json='{runtime_env_json}' " -+ f"-- python3 {train_script} " -+ f"{'${MODEL_ARGS[@]}' if megatron_model_type is not None else ''} " -+ f"{train_args}" -+ ) - - def _parse_extra_env_vars(text: str): - try: -diff --git a/slime/utils/memory_utils.py b/slime/utils/memory_utils.py -index c12f3cd0..89078266 100644 ---- a/slime/utils/memory_utils.py -+++ b/slime/utils/memory_utils.py -@@ -3,6 +3,7 @@ import logging - - import torch - import torch.distributed as dist -+from slime.utils.common import is_npu - - logger = logging.getLogger(__name__) - -@@ -12,12 +13,19 @@ def clear_memory(clear_host_memory: bool = False): - gc.collect() - torch.cuda.empty_cache() - if clear_host_memory: -- torch._C._host_emptyCache() -+ if is_npu(): -+ torch.npu.empty_cache() -+ else: -+ torch._C._host_emptyCache() - - - def available_memory(): -- device = torch.cuda.current_device() -- free, total = torch.cuda.mem_get_info(device) -+ if is_npu(): -+ device = torch.npu.current_device() -+ free, total = torch.npu.mem_get_info(device) -+ else: -+ device = torch.cuda.current_device() -+ free, total = torch.cuda.mem_get_info(device) - return { - "gpu": str(device), - "total_GB": _byte_to_gb(total), -diff --git a/slime_plugins/patch/critic_patch.py b/slime_plugins/patch/critic_patch.py -new file mode 100644 -index 00000000..1f4e1bc6 ---- /dev/null -+++ b/slime_plugins/patch/critic_patch.py -@@ -0,0 +1,94 @@ -+import torch -+from typing import ( -+ List, -+ Mapping, -+ TypeVar, -+ Union, -+) -+from megatron.core.transformer.module import MegatronModule -+HFPreTrained = TypeVar("HFPreTrained") -+MegatronModel = TypeVar("MegatronModel", bound=MegatronModule) -+ -+ -+def load_weights_hf_to_megatron_wrapper( -+ self, hf_pretrained: HFPreTrained, megatron_model: Union[MegatronModel, List[MegatronModel]] -+ ) -> List[MegatronModel]: -+ """Load HuggingFace weights into Megatron models. -+ -+ This method orchestrates the complete weight loading process from HuggingFace -+ format to Megatron's distributed format. It builds a conversion task and -+ executes it with proper progress tracking and error handling. -+ -+ The actual weight transformations and distribution are delegated to the -+ appropriate MegatronParamMapping instances based on the state mappings. -+ -+ Args: -+ hf_pretrained (HFPreTrained): HuggingFace model or state source containing the -+ weights to load. -+ megatron_model (Union[MegatronModel, List[MegatronModel]]): Megatron model instance -+ or list of model instances (one per virtual pipeline stage). -+ -+ Returns: -+ List[MegatronModel]: The input megatron_model as a list with loaded weights. -+ -+ Process: -+ 1. Build a task mapping each Megatron parameter to its source -+ 2. For each parameter in the task: -+ - Fetch source weights from HuggingFace state -+ - Apply format transformation via the param mapping -+ - Distribute to appropriate TP/PP ranks -+ - Copy into the Megatron parameter -+ -+ Example: -+ .. code-block:: python -+ -+ hf_model = PreTrainedCausalLM.from_pretrained("gpt2") -+ megatron_model = create_megatron_model() # Single model or list -+ bridge.load_weights_hf_to_megatron(hf_model, megatron_model) -+ -+ Note: -+ Progress is shown only on rank 0 to avoid cluttered output in -+ distributed environments. -+ -+ Raises: -+ ValueError: If hf_pretrained doesn't have state attribute or if weight shapes don't match. -+ AttributeError: If required HF weights are missing. -+ """ -+ if not isinstance(megatron_model, list): -+ megatron_model = [megatron_model] -+ -+ hf_to_megatron_tasks = self.build_conversion_tasks(hf_pretrained, megatron_model) -+ hf_state_dict: Mapping[str, torch.Tensor] = hf_pretrained.state if hasattr(hf_pretrained, "state") else {} -+ -+ description = f"Loading from {hf_pretrained.model_name_or_path}" -+ for task in self._with_progress_tracking(hf_to_megatron_tasks, description): -+ # None means megatron module not on current rank, skip if this task is not going to happen -+ if task.megatron_module is None: -+ continue -+ # 1) Fetch source tensor(s) from HF state dict -+ hf_weights = self.maybe_modify_loaded_hf_weight(task.mapping.hf_param, hf_state_dict) -+ -+ # 2) Delegate conversion & distribution to the bridge -+ converted_weights = task.mapping.hf_to_megatron(hf_weights, task.megatron_module) -+ -+ # 3) Copy into Megatron param if this rank received a shard -+ if converted_weights is not None: -+ # Assert that param_weight is not None for HF->Megatron tasks -+ assert task.param_weight is not None, "param_weight is required for HF->Megatron conversion" -+ -+ if converted_weights.shape != task.param_weight.shape and task.param_name == 'language_model.output_layer.weight': -+ continue -+ -+ # Check shape compatibility before copying -+ if converted_weights.shape != task.param_weight.shape: -+ raise ValueError( -+ f"Shape mismatch for megatron param {task.mapping.megatron_param}:\n" -+ f" Expected shape: {task.param_weight.shape}\n" -+ f" Got shape: {converted_weights.shape}\n" -+ f" Bridge type: {type(task.mapping).__name__}\n" -+ f" HF mapping: {task.mapping.hf_param}" -+ ) -+ task.param_weight.data.copy_(converted_weights) -+ -+ self._broadcast_shared_embeddings(megatron_model) -+ return megatron_model -\ No newline at end of file -diff --git a/slime_plugins/patch/mbridge_patch.py b/slime_plugins/patch/mbridge_patch.py -new file mode 100644 -index 00000000..606fff9b ---- /dev/null -+++ b/slime_plugins/patch/mbridge_patch.py -@@ -0,0 +1,25 @@ -+try: -+ from megatron.bridge.models.conversion.model_bridge import MegatronModelBridge -+ _original_build_conversion_tasks = MegatronModelBridge.build_conversion_tasks -+ -+ def _patched_build_conversion_tasks(self, hf_pretrained, megatron_model): -+ """ -+ Invoke the original build_conversion_tasks and filter out any actual None tasks. -+ -+ The original implementation might return List[None | WeightConversionTask]. -+ We consolidate it here into List[WeightConversionTask] to avoid errors -+ when accessing None.task.xxx later. -+ """ -+ tasks = _original_build_conversion_tasks(self, hf_pretrained, megatron_model) -+ -+ if tasks is None: -+ return [] -+ -+ filtered = [t for t in tasks if t is not None] -+ -+ return filtered -+ -+ MegatronModelBridge.build_conversion_tasks = _patched_build_conversion_tasks -+ -+except ImportError: -+ pass -diff --git a/train.py b/train.py -index 01883c47..18faa6d2 100644 ---- a/train.py -+++ b/train.py -@@ -1,5 +1,7 @@ - import ray -- -+from slime.utils.common import is_npu -+if is_npu(): -+ import mindspeed.megatron_adaptor - from slime.ray.placement_group import create_placement_groups, create_rollout_manager, create_training_models - from slime.utils.arguments import parse_args - from slime.utils.logging_utils import configure_logger, init_tracking diff --git a/docker/patch/latest/sglang.patch b/docker/patch/latest/sglang.patch deleted file mode 100644 index 4a13e2f9b..000000000 --- a/docker/patch/latest/sglang.patch +++ /dev/null @@ -1,2329 +0,0 @@ -diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py -index 691f06411d..671ac81c48 100644 ---- a/python/sglang/srt/configs/model_config.py -+++ b/python/sglang/srt/configs/model_config.py -@@ -294,6 +294,7 @@ class ModelConfig: - - if is_draft_model and self.hf_config.architectures[0] in [ - "DeepseekV3ForCausalLM", -+ "DeepseekV32ForCausalLM", - "GlmMoeDsaForCausalLM", - ]: - self.hf_config.architectures[0] = "DeepseekV3ForCausalLMNextN" -diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py -index f7d4092d85..3aae51c849 100644 ---- a/python/sglang/srt/disaggregation/base/conn.py -+++ b/python/sglang/srt/disaggregation/base/conn.py -@@ -17,6 +17,7 @@ class KVArgs: - kv_data_ptrs: List[int] - kv_data_lens: List[int] - kv_item_lens: List[int] -+ aux_buffer_names: List[str] - aux_data_ptrs: List[int] - aux_data_lens: List[int] - aux_item_lens: List[int] -diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py -index f54c882cc2..03832002f0 100644 ---- a/python/sglang/srt/disaggregation/decode.py -+++ b/python/sglang/srt/disaggregation/decode.py -@@ -21,6 +21,7 @@ Life cycle of a request in the decode server - from __future__ import annotations - - import logging -+import os - import time - from collections import deque - from dataclasses import dataclass -@@ -42,8 +43,10 @@ from sglang.srt.disaggregation.utils import ( - MetadataBuffers, - ReqToMetadataIdxAllocator, - TransferBackend, -+ apply_prefill_timing_payload, - get_kv_class, - is_mla_backend, -+ is_slime_profiling_enabled, - kv_to_page_indices, - poll_and_all_reduce, - poll_and_all_reduce_with_staging, -@@ -344,6 +347,7 @@ class DecodePreallocQueue: - kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = ( - self.metadata_buffers.get_buf_infos() - ) -+ kv_args.aux_buffer_names = self.metadata_buffers.get_aux_buffer_names() - - if hasattr(self.token_to_kv_pool, "get_state_buf_infos"): - state_data_ptrs, state_data_lens, state_item_lens = ( -@@ -398,6 +402,16 @@ class DecodePreallocQueue: - ) - return kv_manager - -+ def release_memory_occupation(self): -+ self.queue.clear() -+ self.retracted_queue.clear() -+ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): -+ self.kv_manager.deregister_buffer_to_engine() -+ -+ def resume_memory_occupation(self): -+ if hasattr(self.kv_manager, "register_buffer_to_engine"): -+ self.kv_manager.register_buffer_to_engine() -+ - def add(self, req: Req, is_retracted: bool = False) -> None: - """Add a request to the pending queue.""" - if self._check_if_req_exceed_kv_capacity(req): -@@ -525,12 +539,37 @@ class DecodePreallocQueue: - [decode_req.kv_receiver for decode_req in self.queue], self.gloo_group - ) - -+ # Bootstrap timeout: if a request has been stuck in Bootstrapping for too long, treat it as failed. -+ bootstrap_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): - if rids_to_check is not None and decode_req.req.rid not in rids_to_check: - continue - - if poll == KVPoll.Bootstrapping: -- pass -+ # Check for bootstrap timeout -+ entry_time = getattr( -+ decode_req.req.time_stats, -+ "decode_prealloc_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > bootstrap_timeout: -+ error_message = ( -+ f"Decode bootstrap timed out after {now - entry_time:.1f}s " -+ f"for request rank={self.tp_rank} " -+ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ prepare_abort( -+ decode_req.req, -+ error_message, -+ status_code=HTTPStatus.GATEWAY_TIMEOUT, -+ ) -+ if self.scheduler.enable_metrics: -+ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() - elif poll == KVPoll.WaitingForInput: - decode_req.waiting_for_input = True - decode_req.req.time_stats.set_bootstrap_done_time() -@@ -770,6 +809,7 @@ class DecodePreallocQueue: - self.req_to_metadata_buffer_idx_allocator.alloc() - ) - assert decode_req.metadata_buffer_index is not None -+ self.metadata_buffers.clear_profiling_buf(decode_req.metadata_buffer_index) - page_indices = kv_to_page_indices(kv_indices, page_size) - decode_req.kv_receiver.send_metadata( - page_indices, decode_req.metadata_buffer_index, state_indices -@@ -964,6 +1004,7 @@ class DecodeTransferQueue: - output_topk_index, - output_hidden_states, - output_bootstrap_room, -+ output_prefill_timing, - ) = self.metadata_buffers.get_buf(idx) - - # Validate bootstrap_room to detect context corruption -@@ -1025,6 +1066,14 @@ class DecodeTransferQueue: - output_top_logprobs_idx[: decode_req.req.top_logprobs_num].tolist() - ) - -+ # Inject prefill-side PD timing forwarded from the P instance. -+ # Layout: [bootstrap_queue, forward, transfer_queue, bootstrap, -+ # alloc_waiting, transfer_speed, transfer_mb, retry_count] -+ if is_slime_profiling_enabled(): -+ apply_prefill_timing_payload( -+ decode_req.req.time_stats, output_prefill_timing -+ ) -+ - decode_req.kv_receiver.clear() - decode_req.kv_receiver = None - decode_req.req.time_stats.set_wait_queue_entry_time() -@@ -1057,6 +1106,13 @@ class DecodeTransferQueue: - [dr.kv_receiver for dr in self.queue], self.gloo_group - ) - -+ # Transfer timeout: if a request has been in the transfer queue for too long -+ # (e.g., stuck in Bootstrapping/WaitingForInput/Transferring), treat it as failed. -+ transfer_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - transferred_reqs = [] - indices_to_remove = set() - for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): -@@ -1111,7 +1167,20 @@ class DecodeTransferQueue: - KVPoll.WaitingForInput, - KVPoll.Transferring, - ]: -- pass -+ # Check for transfer timeout -+ entry_time = getattr( -+ decode_req.req.time_stats, -+ "decode_transfer_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > transfer_timeout: -+ error_message = ( -+ f"Decode transfer timed out after {now - entry_time:.1f}s " -+ f"(state={poll}) for request rank={self.tp_rank} " -+ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ decode_req.kv_receiver.abort() - else: - raise ValueError(f"Unexpected poll case: {poll}") - -@@ -1132,6 +1201,14 @@ class DecodeTransferQueue: - - return transferred_reqs - -+ def release_memory_occupation(self): -+ """Clean up all in-flight transfers before releasing GPU memory.""" -+ self.queue.clear() -+ -+ def resume_memory_occupation(self): -+ """Resume after GPU memory re-allocation. Queue was already cleared on release.""" -+ pass -+ - - class SchedulerDisaggregationDecodeMixin: - -@@ -1301,7 +1378,15 @@ class SchedulerDisaggregationDecodeMixin: - resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs() - self.waiting_queue.extend(resumed_reqs) - if len(self.disagg_decode_prealloc_queue.retracted_queue) > 0: -- # if there are still retracted requests, we do not allocate new requests -+ # Still have retracted requests that couldn't resume (not enough memory). -+ # Don't accept new requests (pop_preallocated) — they would consume memory -+ # that retracted requests need. -+ # But DO drain completed transfers: their KV is already committed, and -+ # moving them to waiting_queue frees the reserved-decode-token budget -+ # in _allocatable_tokens(), which may unblock resume on the next iteration. -+ # Without this, completed transfers hold memory indefinitely → deadlock. -+ alloc_reqs = self.disagg_decode_transfer_queue.pop_transferred() -+ self.waiting_queue.extend(alloc_reqs) - return - - if not hasattr(self, "polling_count"): -diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py -index 64d97f5c69..4ef08446aa 100644 ---- a/python/sglang/srt/disaggregation/mooncake/conn.py -+++ b/python/sglang/srt/disaggregation/mooncake/conn.py -@@ -31,6 +31,7 @@ from sglang.srt.disaggregation.mooncake.utils import ( - from sglang.srt.disaggregation.utils import ( - DisaggregationMode, - filter_kv_indices_for_cp_rank, -+ iter_aux_transfer_specs, - ) - from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine - from sglang.srt.environ import envs -@@ -276,6 +277,16 @@ class MooncakeKVManager(CommonKVManager): - self.kv_args.state_data_ptrs, self.kv_args.state_data_lens - ) - -+ def deregister_buffer_to_engine(self): -+ if self.kv_args.kv_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.kv_data_ptrs) -+ -+ if self.kv_args.aux_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.aux_data_ptrs) -+ -+ if self.kv_args.state_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.state_data_ptrs) -+ - # ------------------------------------------------------------------ - # Staging buffer methods (all delegate to staging_handler.py) - # ------------------------------------------------------------------ -@@ -884,10 +895,14 @@ class MooncakeKVManager(CommonKVManager): - prefill_aux_ptrs = self.kv_args.aux_data_ptrs - prefill_aux_item_lens = self.kv_args.aux_item_lens - -- for i, dst_aux_ptr in enumerate(dst_aux_ptrs): -- length = prefill_aux_item_lens[i] -- src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -- dst_addr = dst_aux_ptrs[i] + length * req.dst_aux_index -+ for _, src_addr, dst_addr, length in iter_aux_transfer_specs( -+ self.kv_args.aux_buffer_names, -+ prefill_aux_ptrs, -+ prefill_aux_item_lens, -+ dst_aux_ptrs, -+ prefill_aux_index, -+ req.dst_aux_index, -+ ): - transfer_blocks.append((src_addr, dst_addr, length)) - - return self._transfer_data(req.mooncake_session_id, transfer_blocks) -@@ -901,9 +916,14 @@ class MooncakeKVManager(CommonKVManager): - prefill_aux_ptrs = self.kv_args.aux_data_ptrs - prefill_aux_item_lens = self.kv_args.aux_item_lens - -- for i in range(len(prefill_aux_ptrs)): -- length = prefill_aux_item_lens[i] -- src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -+ for i, src_addr, _, length in iter_aux_transfer_specs( -+ self.kv_args.aux_buffer_names, -+ prefill_aux_ptrs, -+ prefill_aux_item_lens, -+ dst_aux_ptrs, -+ prefill_aux_index, -+ req.dst_aux_index, -+ ): - data = AuxDataCodec.serialize_data_from_buffer(src_addr, length) - - self.send_aux_data_to_endpoint( -@@ -1002,13 +1022,13 @@ class MooncakeKVManager(CommonKVManager): - raise RuntimeError( - f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." - ) -- if len(prefill_state_indices) < len(req.dst_state_indices): -- logger.warning( -- f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(req.dst_state_indices)}" -+ if len(prefill_state_indices) != len(req.dst_state_indices): -+ logger.error( -+ "PD extra-state index mismatch, reject transfer to avoid corrupted outputs: " -+ f"len(prefill_state_indices)={len(prefill_state_indices)}, " -+ f"len(dst_state_indices)={len(req.dst_state_indices)}" - ) -- prefill_state_indices = prefill_state_indices[ -- : len(req.dst_state_indices) -- ] -+ return -1 - # Reuse _send_kvcache_generic interface to send extra pool data - prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) - dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32) -@@ -1266,12 +1286,6 @@ class MooncakeKVManager(CommonKVManager): - if ret != 0: - with self.session_lock: - self.session_failures[req.mooncake_session_id] += 1 -- # Failures should never happen if the session is not dead, if the session fails once, mark it as failed -- if self.session_failures[req.mooncake_session_id] >= 1: -- self.failed_sessions.add(req.mooncake_session_id) -- logger.error( -- f"Session {req.mooncake_session_id} failed." -- ) - self.record_failure( - kv_chunk.room, - f"Failed to send kv chunk of {kv_chunk.room} to " -@@ -1289,13 +1303,31 @@ class MooncakeKVManager(CommonKVManager): - - if kv_chunk.is_last_chunk: - if kv_chunk.state_indices is not None: -- self.maybe_send_extra( -+ ret = self.maybe_send_extra( - req, - kv_chunk.state_indices, - target_rank_registration_info.dst_state_data_ptrs, - executor, - target_rank_registration_info, - ) -+ if ret != 0: -+ remote_addr = NetworkAddress( -+ req.endpoint, req.dst_port -+ ).to_host_port_str() -+ self.record_failure( -+ kv_chunk.room, -+ f"Failed to send extra state chunk of {kv_chunk.room} to " -+ f"{remote_addr}", -+ ) -+ self.update_status(kv_chunk.room, KVPoll.Failed) -+ self.sync_status_to_decode_endpoint( -+ req.endpoint, -+ req.dst_port, -+ req.room, -+ KVPoll.Failed, -+ prefill_unique_rank, -+ ) -+ break - - # Only the last chunk we need to send the aux data - ret = self.send_aux( -diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py -index 8eadf81954..c180ce79f3 100644 ---- a/python/sglang/srt/disaggregation/prefill.py -+++ b/python/sglang/srt/disaggregation/prefill.py -@@ -20,6 +20,8 @@ Life cycle of a request in the prefill server - from __future__ import annotations - - import logging -+import os -+import time - from collections import deque - from http import HTTPStatus - from typing import TYPE_CHECKING, List, Optional -@@ -165,6 +167,7 @@ class PrefillBootstrapQueue: - kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = ( - self.metadata_buffers.get_buf_infos() - ) -+ kv_args.aux_buffer_names = self.metadata_buffers.get_aux_buffer_names() - kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device - kv_args.gpu_id = self.scheduler.gpu_id - -@@ -290,6 +293,11 @@ class PrefillBootstrapQueue: - self.scheduler.attn_tp_cpu_group, - ) - -+ bootstrap_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - for i, (req, poll) in enumerate(zip(self.queue, polls)): - if rids_to_check is not None: - # if req not in reqs_info_to_check, skip -@@ -297,6 +305,26 @@ class PrefillBootstrapQueue: - continue - - if poll == KVPoll.Bootstrapping: -+ entry_time = getattr( -+ req.time_stats, -+ "prefill_bootstrap_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > bootstrap_timeout: -+ error_message = ( -+ f"Prefill bootstrap timed out after {now - entry_time:.1f}s " -+ f"for request rank={self.tp_rank} " -+ f"{req.rid=} {req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ prepare_abort( -+ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT -+ ) -+ self.scheduler.stream_output([req], req.return_logprob) -+ indices_to_remove.add(i) -+ failed_reqs.append(req) -+ if self.scheduler.enable_metrics: -+ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() - continue - elif poll == KVPoll.Failed: - error_message = f"Prefill bootstrap failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}" -@@ -346,6 +374,15 @@ class PrefillBootstrapQueue: - else: - return bootstrapped_reqs, failed_reqs - -+ def release_memory_occupation(self): -+ self.queue.clear() -+ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): -+ self.kv_manager.deregister_buffer_to_engine() -+ -+ def resume_memory_occupation(self): -+ if hasattr(self.kv_manager, "register_buffer_to_engine"): -+ self.kv_manager.register_buffer_to_engine() -+ - - class SchedulerDisaggregationPrefillMixin: - """ -@@ -568,12 +605,17 @@ class SchedulerDisaggregationPrefillMixin: - self.send_kv_chunk(req, last_chunk=False, end_idx=req.tmp_end_idx) - req.time_stats.set_last_chunked_prefill_finish_time() - -- can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) -- self.report_prefill_stats( -- prefill_stats=batch.prefill_stats, -- can_run_cuda_graph=can_run_cuda_graph, -- dp_cooperation_info=batch.dp_cooperation_info, -- ) -+ if ( -+ self.current_scheduler_metrics_enabled -+ and hasattr(batch, "prefill_stats") -+ and batch.prefill_stats is not None -+ ): -+ can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) -+ self.report_prefill_stats( -+ prefill_stats=batch.prefill_stats, -+ can_run_cuda_graph=can_run_cuda_graph, -+ dp_cooperation_info=getattr(batch, "dp_cooperation_info", None), -+ ) - - def process_disagg_prefill_inflight_queue( - self: Scheduler, rids_to_check: Optional[List[str]] = None -@@ -593,6 +635,11 @@ class SchedulerDisaggregationPrefillMixin: - self.attn_tp_cpu_group, - ) - -+ transfer_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - undone_reqs: List[Req] = [] - # Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue - for req, poll in zip(self.disagg_prefill_inflight_queue, polls): -@@ -618,7 +665,29 @@ class SchedulerDisaggregationPrefillMixin: - continue - - if poll in [KVPoll.WaitingForInput, KVPoll.Transferring]: -- undone_reqs.append(req) -+ entry_time = getattr( -+ req.time_stats, -+ "prefill_transfer_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > transfer_timeout: -+ error_message = ( -+ f"Prefill transfer timed out after {now - entry_time:.1f}s " -+ f"(state={poll}) for request rank={self.tp_rank} " -+ f"{req.rid=} {req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ release_kv_cache(req, self.tree_cache) -+ prepare_abort( -+ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT -+ ) -+ if hasattr(req.disagg_kv_sender, "clear"): -+ req.disagg_kv_sender.clear() -+ done_reqs.append(req) -+ if self.enable_metrics: -+ self.metrics_collector.increment_transfer_failed_reqs() -+ else: -+ undone_reqs.append(req) - elif poll == KVPoll.Success: # transfer done - release_kv_cache(req, self.tree_cache) # unlock the tree - req.finished_reason = FINISH_LENGTH(length=0) -diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py -index d7956a6048..0ced278713 100644 ---- a/python/sglang/srt/disaggregation/utils.py -+++ b/python/sglang/srt/disaggregation/utils.py -@@ -28,6 +28,17 @@ if TYPE_CHECKING: - # Constants & Enums - ######################### - FAKE_BOOTSTRAP_HOST = "2.2.2.2" -+PREFILL_TIMING_AUX_BUFFER_NAME = "prefill_timing" -+PREFILL_TIMING_DEST_ATTRS = ( -+ ("fwd_prefill_bootstrap_queue_duration", float), -+ ("fwd_prefill_forward_duration", float), -+ ("fwd_prefill_transfer_queue_duration", float), -+ ("fwd_bootstrap_duration", float), -+ ("fwd_alloc_waiting_duration", float), -+ ("fwd_transfer_speed_gb_s", float), -+ ("fwd_transfer_total_mb", float), -+ ("fwd_prefill_retry_count", int), -+) - - - class DisaggregationMode(Enum): -@@ -193,46 +204,35 @@ class MetadataBuffers: - self.bootstrap_room = torch.zeros( - (size, 8), dtype=bootstrap_room_dtype, device=device - ) -+ # Prefill-side PD timing (8 floats, padded to 16 for RDMA alignment). -+ # Layout: [bootstrap_queue, forward, transfer_queue, bootstrap, -+ # alloc_waiting, transfer_speed, transfer_mb, retry_count] -+ self.prefill_timing = torch.zeros( -+ (size, 16), dtype=torch.float32, device=device -+ ) -+ self.aux_buffers = [ -+ ("output_ids", self.output_ids), -+ ("cached_tokens", self.cached_tokens), -+ ("output_token_logprobs_val", self.output_token_logprobs_val), -+ ("output_token_logprobs_idx", self.output_token_logprobs_idx), -+ ("output_top_logprobs_val", self.output_top_logprobs_val), -+ ("output_top_logprobs_idx", self.output_top_logprobs_idx), -+ ("output_topk_p", self.output_topk_p), -+ ("output_topk_index", self.output_topk_index), -+ ("output_hidden_states", self.output_hidden_states), -+ ("bootstrap_room", self.bootstrap_room), -+ (PREFILL_TIMING_AUX_BUFFER_NAME, self.prefill_timing), -+ ] - - def get_buf_infos(self): -- ptrs = [ -- self.output_ids.data_ptr(), -- self.cached_tokens.data_ptr(), -- self.output_token_logprobs_val.data_ptr(), -- self.output_token_logprobs_idx.data_ptr(), -- self.output_top_logprobs_val.data_ptr(), -- self.output_top_logprobs_idx.data_ptr(), -- self.output_topk_p.data_ptr(), -- self.output_topk_index.data_ptr(), -- self.output_hidden_states.data_ptr(), -- self.bootstrap_room.data_ptr(), -- ] -- data_lens = [ -- self.output_ids.nbytes, -- self.cached_tokens.nbytes, -- self.output_token_logprobs_val.nbytes, -- self.output_token_logprobs_idx.nbytes, -- self.output_top_logprobs_val.nbytes, -- self.output_top_logprobs_idx.nbytes, -- self.output_topk_p.nbytes, -- self.output_topk_index.nbytes, -- self.output_hidden_states.nbytes, -- self.bootstrap_room.nbytes, -- ] -- item_lens = [ -- self.output_ids[0].nbytes, -- self.cached_tokens[0].nbytes, -- self.output_token_logprobs_val[0].nbytes, -- self.output_token_logprobs_idx[0].nbytes, -- self.output_top_logprobs_val[0].nbytes, -- self.output_top_logprobs_idx[0].nbytes, -- self.output_topk_p[0].nbytes, -- self.output_topk_index[0].nbytes, -- self.output_hidden_states[0].nbytes, -- self.bootstrap_room[0].nbytes, -- ] -+ ptrs = [buffer.data_ptr() for _, buffer in self.aux_buffers] -+ data_lens = [buffer.nbytes for _, buffer in self.aux_buffers] -+ item_lens = [buffer[0].nbytes for _, buffer in self.aux_buffers] - return ptrs, data_lens, item_lens - -+ def get_aux_buffer_names(self): -+ return [name for name, _ in self.aux_buffers] -+ - def get_buf(self, idx: int): - return ( - self.output_ids[idx], -@@ -245,8 +245,12 @@ class MetadataBuffers: - self.output_topk_index[idx], - self.output_hidden_states[idx], - self.bootstrap_room[idx], -+ self.prefill_timing[idx], - ) - -+ def clear_profiling_buf(self, idx: int): -+ self.prefill_timing[idx].zero_() -+ - def set_buf(self, req: Req): - - self.output_ids[req.metadata_buffer_index][0] = req.output_ids[0] -@@ -294,6 +298,99 @@ class MetadataBuffers: - self.bootstrap_room[req.metadata_buffer_index, 0] = ( - req.bootstrap_room if req.bootstrap_room is not None else 0 - ) -+ # Pack prefill-side PD timing durations for transfer to decode instance. -+ # Note: set_buf is called at the START of the last KV chunk send, so -+ # completion_time and prefill_transfer_queue_entry_time are not yet set. -+ # We use time.perf_counter() as the "forward just completed" timestamp. -+ import time -+ -+ ts = req.time_stats -+ timing = self.prefill_timing[req.metadata_buffer_index] -+ self.clear_profiling_buf(req.metadata_buffer_index) -+ if not is_slime_profiling_enabled(): -+ return -+ for idx, value in enumerate( -+ build_prefill_timing_payload(ts, now=time.perf_counter()) -+ ): -+ if value > 0: -+ timing[idx] = value -+ -+ -+def is_slime_profiling_enabled() -> bool: -+ return envs.SLIME_ENABLE_PROFILING.get() -+ -+ -+def build_prefill_timing_payload(time_stats, now: float) -> tuple[float, ...]: -+ bootstrap_queue_duration = 0.0 -+ if ( -+ time_stats.prefill_bootstrap_queue_entry_time > 0 -+ and time_stats.wait_queue_entry_time > 0 -+ ): -+ bootstrap_queue_duration = ( -+ time_stats.wait_queue_entry_time -+ - time_stats.prefill_bootstrap_queue_entry_time -+ ) -+ -+ prefill_forward_duration = ( -+ now - time_stats.forward_entry_time -+ if time_stats.forward_entry_time > 0 -+ else 0.0 -+ ) -+ -+ bootstrap_duration = 0.0 -+ alloc_waiting_duration = 0.0 -+ if ( -+ time_stats.prefill_bootstrap_queue_entry_time > 0 -+ and time_stats.bootstrap_done_time > 0 -+ ): -+ bootstrap_duration = ( -+ time_stats.bootstrap_done_time -+ - time_stats.prefill_bootstrap_queue_entry_time -+ ) -+ if time_stats.bootstrap_done_time > 0 and time_stats.wait_queue_entry_time > 0: -+ alloc_waiting_duration = ( -+ time_stats.wait_queue_entry_time - time_stats.bootstrap_done_time -+ ) -+ -+ return ( -+ bootstrap_queue_duration, -+ prefill_forward_duration, -+ 0.0, -+ max(0.0, bootstrap_duration), -+ max(0.0, alloc_waiting_duration), -+ max(0.0, time_stats.transfer_speed_gb_s), -+ max(0.0, time_stats.transfer_total_mb), -+ float(max(0, time_stats.prefill_retry_count)), -+ ) -+ -+ -+def apply_prefill_timing_payload(time_stats, timing) -> None: -+ for value, (attr_name, caster) in zip( -+ timing[: len(PREFILL_TIMING_DEST_ATTRS)].tolist(), -+ PREFILL_TIMING_DEST_ATTRS, -+ ): -+ if value > 0: -+ setattr(time_stats, attr_name, caster(value)) -+ -+ -+def iter_aux_transfer_specs( -+ aux_buffer_names: list[str], -+ prefill_aux_ptrs: list[int], -+ prefill_aux_item_lens: list[int], -+ dst_aux_ptrs: list[int], -+ prefill_aux_index: int, -+ dst_aux_index: int, -+): -+ profiling_enabled = is_slime_profiling_enabled() -+ for i, (buffer_name, dst_aux_ptr) in enumerate(zip(aux_buffer_names, dst_aux_ptrs)): -+ if not profiling_enabled and buffer_name == PREFILL_TIMING_AUX_BUFFER_NAME: -+ continue -+ length = prefill_aux_item_lens[i] -+ if length <= 0: -+ continue -+ src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -+ dst_addr = dst_aux_ptr + length * dst_aux_index -+ yield i, src_addr, dst_addr, length - - - ######################### -diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py -index d864e4abaa..3a000a80f2 100644 ---- a/python/sglang/srt/entrypoints/engine.py -+++ b/python/sglang/srt/entrypoints/engine.py -@@ -69,6 +69,7 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqInput, - MultimodalDataInputFormat, - OpenSessionReqInput, -+ PostProcessWeightsReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, - RpcReqInput, -@@ -957,6 +958,24 @@ class Engine(EngineScoreMixin, EngineBase): - self.tokenizer_manager.update_weights_from_ipc(obj, None) - ) - -+ def post_process_weights( -+ self, -+ restore_weights_before_load: bool = False, -+ post_process_quantization: bool = False, -+ ): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ Should be called after weight update is finished. -+ """ -+ obj = PostProcessWeightsReqInput( -+ restore_weights_before_load=restore_weights_before_load, -+ post_process_quantization=post_process_quantization, -+ ) -+ -+ return self.loop.run_until_complete( -+ self.tokenizer_manager.post_process_weights(obj, None) -+ ) -+ - def get_weights_by_name(self, name: str, truncate_size: int = 100): - """Get weights by parameter name.""" - obj = GetWeightsByNameReqInput(name=name, truncate_size=truncate_size) -diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py -index 6978e0c062..80dc159e8f 100644 ---- a/python/sglang/srt/entrypoints/http_server.py -+++ b/python/sglang/srt/entrypoints/http_server.py -@@ -127,6 +127,7 @@ from sglang.srt.managers.io_struct import ( - OpenSessionReqInput, - ParseFunctionCallReq, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -582,10 +583,8 @@ async def model_info(): - @app.get("/weight_version") - async def weight_version(): - """Get the current weight version.""" -- raise HTTPException( -- status_code=404, -- detail="Endpoint '/get_weight_version' or '/weight_version' is deprecated. Please use '/model_info' instead.", -- ) -+ result = await model_info() -+ return {"weight_version": result.get("weight_version", None)} - - - @app.get("/get_server_info") -@@ -602,9 +601,19 @@ async def get_server_info(): - async def server_info(): - """Get the server information.""" - # Returns internal states per DP. -- internal_states: List[Dict[Any, Any]] = ( -- await _global_state.tokenizer_manager.get_internal_state() -- ) -+ # In large/disaggregated deployments this can occasionally block; keep endpoint responsive. -+ server_info_timeout = float(os.environ.get("SGLANG_SERVER_INFO_TIMEOUT", "2")) -+ try: -+ internal_states: List[Dict[Any, Any]] = await asyncio.wait_for( -+ _global_state.tokenizer_manager.get_internal_state(), -+ timeout=server_info_timeout, -+ ) -+ except asyncio.TimeoutError: -+ logger.warning( -+ "Timed out getting internal state for /server_info after %.1fs; returning empty internal_states", -+ server_info_timeout, -+ ) -+ internal_states = [] - - # This field is not serializable. - if hasattr(_global_state.tokenizer_manager.server_args, "model_config"): -@@ -1121,6 +1130,23 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re - return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) - - -+@app.post("/post_process_weights") -+@auth_level(AuthLevel.ADMIN_OPTIONAL) -+async def post_process_weights(req: PostProcessWeightsReqInput, request: Request): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ This should be called selectively after `update_weights_from_distributed/update_weights_from_tensor`. -+ """ -+ success, message = await _global_state.tokenizer_manager.post_process_weights( -+ req, request -+ ) -+ -+ content = {"success": success, "message": message} -+ return ORJSONResponse( -+ content, status_code=200 if success else HTTPStatus.BAD_REQUEST -+ ) -+ -+ - @app.post("/update_weight_version") - @auth_level(AuthLevel.ADMIN_OPTIONAL) - async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request): -diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py -index dfc5507de0..be9501b05a 100644 ---- a/python/sglang/srt/environ.py -+++ b/python/sglang/srt/environ.py -@@ -242,6 +242,7 @@ class Envs: - SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2) - SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300) - SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX") -+ SLIME_ENABLE_PROFILING = EnvBool(False) - SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False) - # Extra slots in req_to_token_pool for decode workers (only effective when - # max_num_reqs > 32). Increases pool capacity so more KV cache transfers -diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -index 02ef4e2440..fd5a43cce8 100644 ---- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -+++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -@@ -1,6 +1,7 @@ - from __future__ import annotations - - import contextlib -+import os - from abc import ABC, abstractmethod - from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple - -@@ -213,14 +214,31 @@ class Indexer(MultiPlatformOp): - prefix=add_prefix("weights_proj", prefix), - ) - self.k_norm = LayerNorm(self.head_dim, dtype=torch.float32) -+ server_args = get_global_server_args() -+ disable_flag = server_args.disable_indexer_rope_neox_style -+ env_raw = os.environ.get("INDEXER_ROPE_NEOX_STYLE", None) -+ if env_raw is not None: -+ env_value = env_raw == "1" -+ if disable_flag and env_value: -+ raise ValueError( -+ "Conflict: --disable-indexer-rope-neox-style is set but " -+ "INDEXER_ROPE_NEOX_STYLE='1'. " -+ "Please remove one or make them consistent." -+ ) -+ resolved_neox_style = env_value -+ elif disable_flag: -+ resolved_neox_style = False -+ else: -+ resolved_neox_style = is_neox_style -+ - self.rotary_emb = get_rope_wrapper( - rope_head_dim, - rotary_dim=rope_head_dim, - max_position=max_position_embeddings, - base=rope_theta, # type: ignore - rope_scaling=rope_scaling, -- is_neox_style=is_neox_style, -- device=get_global_server_args().device, -+ is_neox_style=resolved_neox_style, -+ device=server_args.device, - ) - self.block_size = block_size - self.scale_fmt = scale_fmt -@@ -266,6 +284,11 @@ class Indexer(MultiPlatformOp): - @torch.compile(dynamic=True) if not _is_hip else lambda f: f - def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor): - weights = self._weights_proj_bf16_in_fp32_out(x) -+ if weights.shape[1] < q_scale.shape[1]: -+ assert q_scale.shape[1] % weights.shape[1] == 0 -+ weights = weights.repeat_interleave( -+ q_scale.shape[1] // weights.shape[1], dim=1 -+ ) - weights = weights * self.n_heads**-0.5 - weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale - return weights -@@ -1078,6 +1101,9 @@ class Indexer(MultiPlatformOp): - query, key = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch - ) -+ if query.shape[1] < 32: -+ assert 32 % query.shape[1] == 0 -+ query = query.repeat_interleave(32 // query.shape[1], dim=1) - q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) - with torch.cuda.stream(self.alt_stream): - self._store_index_k_cache( -@@ -1092,6 +1118,9 @@ class Indexer(MultiPlatformOp): - query, key = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch - ) -+ if query.shape[1] < 32: -+ assert 32 % query.shape[1] == 0 -+ query = query.repeat_interleave(32 // query.shape[1], dim=1) - - if enable_dual_stream: - current_stream = torch.cuda.current_stream() -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -index 72483f4ea6..2e1148d189 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -@@ -702,6 +702,7 @@ class FusedMoE(torch.nn.Module): - "CompressedTensorsWNA16TritonMoE", - ] - ) -+ and "zero" not in weight_name - else loaded_weight - ) - -@@ -921,6 +922,7 @@ class FusedMoE(torch.nn.Module): - "CompressedTensorsWNA16TritonMoE", - ] - ) -+ and "zero" not in weight_name - else loaded_weight - ) - -diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py -index 00bd687555..12d5577af2 100644 ---- a/python/sglang/srt/layers/moe/routed_experts_capturer.py -+++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py -@@ -8,10 +8,15 @@ import torch - - from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.layers.dp_attention import ( -+ attn_tp_all_gather_into_tensor, - get_attention_dp_rank, -+ get_attention_tp_size, - get_dp_local_info, - is_dp_attention_enabled, - ) -+from sglang.srt.layers.moe import ( -+ get_moe_a2a_backend, -+) - from sglang.srt.mem_cache.memory_pool import ReqToTokenPool - from sglang.srt.model_executor.forward_batch_info import ForwardBatch - from sglang.srt.server_args import get_global_server_args -@@ -181,13 +186,26 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - device=device, - ) - -+ if get_moe_a2a_backend().is_deepep(): -+ attn_tp_size = get_attention_tp_size() if is_dp_attention_enabled() else 1 -+ self.gather_buffer = torch.empty( -+ ( -+ self.device_cache.buffer.shape[0] * attn_tp_size, -+ self.device_cache.buffer.shape[2], -+ ), -+ dtype=torch.int32, -+ device=device, -+ ) -+ - def _sync_fwd_experts_buffer_DtoH( - self, - forward_batch: ForwardBatch, - can_run_graph: bool, - cuda_graph_batch: int, - ): -- if is_dp_attention_enabled(): -+ # When DeepEP is enabled, capture() already does all_gather, so device_cache.buffer -+ # contains data from all DP ranks. We should not slice by DP rank in this case. -+ if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep(): - local_start_pos, local_num_tokens = get_dp_local_info(forward_batch) - # handle with cuda graph padding - if can_run_graph: -@@ -206,6 +224,12 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - ].cpu() - - def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ if get_moe_a2a_backend().is_deepep(): -+ local_topk_ids = topk_ids -+ topk_ids = self.gather_buffer[ -+ : local_topk_ids.size(0) * get_attention_tp_size() -+ ] -+ attn_tp_all_gather_into_tensor(topk_ids, local_topk_ids) - self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) - - def get_routed_experts( -diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -index a13c53af4d..1d80d06b13 100644 ---- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -+++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -@@ -500,7 +500,7 @@ class CompressedTensorsConfig(QuantizationConfig): - ) - is_static = not weight_quant.dynamic - -- return is_channel_group and input_quant_none and is_symmetric and is_static -+ return is_channel_group and input_quant_none and is_static - - def _is_mxint4a16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool: - input_quant_none = input_quant is None -@@ -969,6 +969,9 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase): - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - layer.scheme.process_weights_after_loading(layer) - -+ def restore_weights_before_loading(self, layer: torch.nn.Module) -> None: -+ layer.scheme.restore_weights_before_loading(layer) -+ - def create_weights( - self, - layer: torch.nn.Module, -diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -index 7a8fb65421..f1c85899cd 100644 ---- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -+++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -@@ -17,7 +17,10 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import ( - CompressedTensorsMoEScheme, - ) - from sglang.srt.layers.quantization.gptq import gptq_marlin_moe_repack --from sglang.srt.layers.quantization.marlin_utils import marlin_moe_permute_scales -+from sglang.srt.layers.quantization.marlin_utils import ( -+ marlin_moe_permute_scales, -+ moe_awq_to_marlin_zero_points, -+) - from sglang.srt.layers.quantization.utils import replace_parameter - from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs - -@@ -64,7 +67,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - self.strategy = config.strategy - self.group_size = config.group_size - self.actorder = config.actorder -- assert config.symmetric, "Only symmetric quantization is supported for MoE" -+ self.sym = config.symmetric - - if not ( - self.quant_config.quant_format == CompressionFormat.pack_quantized.value -@@ -124,7 +127,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - - # In the case where we have actorder/g_idx, - # we do not partition the w2 scales -- load_full_w2 = self.actorder and self.group_size != -1 -+ load_full_w2 = (self.actorder != "static") and self.group_size != -1 - - if load_full_w2: - w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size -@@ -172,6 +175,32 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - layer.register_parameter("w13_weight_shape", w13_weight_shape) - set_weight_attrs(w13_weight_shape, extra_weight_attrs) - -+ # add zero param -+ if not self.sym: -+ w13_qzeros = torch.nn.Parameter( -+ torch.empty( -+ num_experts, -+ num_groups_w13, -+ 2 * intermediate_size_per_partition // self.packed_factor, -+ dtype=torch.int32, -+ ), -+ requires_grad=False, -+ ) -+ layer.register_parameter("w13_weight_zero_point", w13_qzeros) -+ set_weight_attrs(w13_qzeros, extra_weight_attrs) -+ -+ w2_qzeros = torch.nn.Parameter( -+ torch.empty( -+ num_experts, -+ num_groups_w2, -+ hidden_size // self.packed_factor, -+ dtype=torch.int32, -+ ), -+ requires_grad=False, -+ ) -+ layer.register_parameter("w2_weight_zero_point", w2_qzeros) -+ set_weight_attrs(w2_qzeros, extra_weight_attrs) -+ - w13_g_idx = torch.nn.Parameter( - torch.empty( - num_experts, -@@ -231,6 +260,10 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) - layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) - -+ if not self.sym: -+ layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape -+ layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape) -+ - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - - # Skip if the layer is already converted to Marlin format to prevent double-packing. -@@ -334,6 +367,24 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - ) - replace_tensor("w2_weight_scale", marlin_w2_scales) - -+ # Repack zero -+ if not self.sym: -+ marlin_w13_zp = moe_awq_to_marlin_zero_points( -+ layer.w13_weight_zero_point, -+ size_k=layer.w13_weight_zero_point.shape[1], -+ size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor, -+ num_bits=self.num_bits, -+ ) -+ replace_tensor("w13_weight_zero_point", marlin_w13_zp) -+ -+ marlin_w2_zp = moe_awq_to_marlin_zero_points( -+ layer.w2_weight_zero_point, -+ size_k=layer.w2_weight_zero_point.shape[1], -+ size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor, -+ num_bits=self.num_bits, -+ ) -+ replace_tensor("w2_weight_zero_point", marlin_w2_zp) -+ - layer.is_marlin_converted = True - - def restore_weights_before_loading(self, layer: torch.nn.Module): -@@ -399,6 +450,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - g_idx2=layer.w2_weight_g_idx, - sort_indices1=layer.w13_g_idx_sort_indices, - sort_indices2=layer.w2_g_idx_sort_indices, -+ w1_zeros=layer.w13_weight_zero_point if not self.sym else None, -+ w2_zeros=layer.w2_weight_zero_point if not self.sym else None, - num_bits=self.num_bits, - is_k_full=self.is_k_full, - routed_scaling_factor=self.moe_runner_config.routed_scaling_factor, -diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index bd97965345..e6a147c1b4 100644 ---- a/python/sglang/srt/managers/io_struct.py -+++ b/python/sglang/srt/managers/io_struct.py -@@ -1449,6 +1449,18 @@ class ResumeMemoryOccupationReqOutput(BaseReq): - pass - - -+@dataclass -+class PostProcessWeightsReqInput(BaseReq): -+ restore_weights_before_load: bool = False -+ post_process_quantization: bool = False -+ -+ -+@dataclass -+class PostProcessWeightsReqOutput(BaseReq): -+ success: bool -+ message: str -+ -+ - @dataclass - class CheckWeightsReqInput(BaseReq): - action: str -@@ -1753,6 +1765,8 @@ class GetLoadReqOutput(BaseReq): - num_waiting_reqs: int - num_tokens: int - ts_tic: float -+ queue_details: Optional[List[Dict[str, Any]]] = None -+ running_details: Optional[Dict[str, Any]] = None - - - @dataclass -diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py -index e0a1669fb3..fbbb6bb12b 100644 ---- a/python/sglang/srt/managers/multi_tokenizer_mixin.py -+++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py -@@ -496,6 +496,35 @@ def monkey_patch_uvicorn_multiprocessing(timeout: float = 10): - "uvicorn.supervisors.multiprocess not found, skipping monkey patch" - ) - -+ try: -+ import uvicorn._subprocess as uvicorn_subprocess -+ import uvicorn.supervisors.multiprocess as uvicorn_multiprocess -+ -+ def _safe_get_stdin_fileno(): -+ try: -+ fileno = sys.stdin.fileno() -+ dup_fd = os.dup(fileno) -+ os.close(dup_fd) -+ return fileno -+ except (AttributeError, OSError): -+ return None -+ -+ def _patched_get_subprocess(config, target, sockets): -+ kwargs = { -+ "config": config, -+ "target": target, -+ "sockets": sockets, -+ "stdin_fileno": _safe_get_stdin_fileno(), -+ } -+ return uvicorn_subprocess.spawn.Process( -+ target=uvicorn_subprocess.subprocess_started, kwargs=kwargs -+ ) -+ -+ uvicorn_subprocess.get_subprocess = _patched_get_subprocess -+ uvicorn_multiprocess.get_subprocess = _patched_get_subprocess -+ except Exception: -+ pass -+ - - class SenderWrapper: - def __init__(self, port_args: PortArgs, send_to_scheduler: zmq.Socket): -diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index 0b26be6c6d..2ea1042cf9 100644 ---- a/python/sglang/srt/managers/schedule_batch.py -+++ b/python/sglang/srt/managers/schedule_batch.py -@@ -1972,7 +1972,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - while first_iter or ( - not self.check_decode_mem(selected_indices=sorted_indices) - ): -- if len(sorted_indices) == 1: -+ # We should allow all requests to be retracted in decode disaggregation mode -+ # because there call be prealloc prefill requests. -+ num_minimum_reqs = 0 if server_args.disaggregation_mode == "decode" else 1 -+ if len(sorted_indices) == num_minimum_reqs: - # Always keep at least one request - break - -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index 67af2d0de9..122ddb3874 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -120,6 +120,7 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqOutput, - OpenSessionReqInput, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReq, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -1232,6 +1233,7 @@ class Scheduler( - ), - (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor), - (UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc), -+ (PostProcessWeightsReqInput, self.post_process_weights), - (GetWeightsByNameReqInput, self.get_weights_by_name), - (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), - (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), -diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -index 496cd96656..cf2d43015a 100644 ---- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py -+++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -@@ -1154,7 +1154,7 @@ class SchedulerOutputProcessorMixin: - dp_ranks = [self.dp_rank] * len(rids) if rids else None - - # Send to detokenizer -- if reqs or is_idle_batch: -+ if rids or is_idle_batch: - self.send_to_detokenizer.send_output( - BatchTokenIDOutput( - rids=rids, -diff --git a/python/sglang/srt/managers/scheduler_profiler_mixin.py b/python/sglang/srt/managers/scheduler_profiler_mixin.py -index c02ed7997d..61733c4127 100644 ---- a/python/sglang/srt/managers/scheduler_profiler_mixin.py -+++ b/python/sglang/srt/managers/scheduler_profiler_mixin.py -@@ -349,7 +349,7 @@ class SchedulerProfilerMixin: - if self.profiler_prefill_ct > self.profiler_target_prefill_ct: - if self.profile_in_progress: - self.stop_profile(stage=ForwardMode.EXTEND) -- elif batch.forward_mode.is_decode(): -+ elif batch.forward_mode.is_decode() or batch.forward_mode.is_prebuilt(): - if self.profiler_decode_ct == 0: - if self.profile_in_progress: - # force trace flush -diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -index abcda67946..a53848b79d 100644 ---- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py -+++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -@@ -12,6 +12,7 @@ from sglang.srt.constants import ( - GPU_MEMORY_TYPE_KV_CACHE, - GPU_MEMORY_TYPE_WEIGHTS, - ) -+from sglang.srt.disaggregation.utils import DisaggregationMode - from sglang.srt.managers.io_struct import ( - CheckWeightsReqInput, - CheckWeightsReqOutput, -@@ -21,6 +22,8 @@ from sglang.srt.managers.io_struct import ( - GetWeightsByNameReqOutput, - InitWeightsUpdateGroupReqInput, - InitWeightsUpdateGroupReqOutput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ReleaseMemoryOccupationReqInput, - ReleaseMemoryOccupationReqOutput, - ResumeMemoryOccupationReqInput, -@@ -117,6 +120,11 @@ class SchedulerUpdateWeightsMixin: - torch.distributed.barrier(group=self.tp_cpu_group) - return UpdateWeightsFromIPCReqOutput(success, message) - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Optional post-processing for updated weights (e.g., Marlin conversion).""" -+ success, message = self.tp_worker.post_process_weights(recv_req) -+ return PostProcessWeightsReqOutput(success, message) -+ - def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): - parameter = self.tp_worker.get_weights_by_name(recv_req) - return GetWeightsByNameReqOutput(parameter) -@@ -140,6 +148,15 @@ class SchedulerUpdateWeightsMixin: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) - self.flush_cache() - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_transfer_queue"): -+ self.disagg_decode_transfer_queue.release_memory_occupation() -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.release_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.release_memory_occupation() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.stashed_model_static_state = _export_static_state( - self.tp_worker.model_runner.model -@@ -180,6 +197,15 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_KV_CACHE in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_transfer_queue"): -+ self.disagg_decode_transfer_queue.resume_memory_occupation() -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.resume_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.resume_memory_occupation() -+ - return ResumeMemoryOccupationReqOutput() - - def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): -diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -index 544c609401..841658c30e 100644 ---- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py -+++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -@@ -59,6 +59,8 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqOutput, - LoRAUpdateOutput, - OpenSessionReqInput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ProfileReq, - ProfileReqOutput, - ProfileReqType, -@@ -187,6 +189,9 @@ class TokenizerCommunicatorMixin: - self.update_weights_from_ipc_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -+ self.post_process_weights_communicator = _Communicator( -+ self.send_to_scheduler, server_args.dp_size -+ ) - self.get_weights_by_name_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -@@ -272,6 +277,10 @@ class TokenizerCommunicatorMixin: - UpdateWeightsFromIPCReqOutput, - self.update_weights_from_ipc_communicator.handle_recv, - ), -+ ( -+ PostProcessWeightsReqOutput, -+ self.post_process_weights_communicator.handle_recv, -+ ), - ( - GetWeightsByNameReqOutput, - self.get_weights_by_name_communicator.handle_recv, -@@ -530,6 +539,17 @@ class TokenizerCommunicatorMixin: - - return success, message - -+ async def post_process_weights( -+ self: TokenizerManager, -+ obj: PostProcessWeightsReqInput, -+ request: Optional[fastapi.Request] = None, -+ ) -> Tuple[bool, str]: -+ """Trigger post-processing hooks for weights after loading (e.g., Marlin conversion).""" -+ self.auto_create_handle_loop() -+ async with self.model_update_lock.writer_lock: -+ results = await self.post_process_weights_communicator(obj) -+ return _Communicator.merge_results(results) -+ - async def init_weights_send_group_for_remote_instance( - self: TokenizerManager, - obj: InitWeightsSendGroupForRemoteInstanceReqInput, -diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index 81424329a0..2c132be63d 100644 ---- a/python/sglang/srt/managers/tokenizer_manager.py -+++ b/python/sglang/srt/managers/tokenizer_manager.py -@@ -1383,7 +1383,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): - async with self.is_pause_cond: - self.is_pause = True - if obj.mode != "abort": -- await self.send_to_scheduler.send_pyobj(obj) -+ self.send_to_scheduler.send_pyobj(obj) - else: - # we are using the model_update_lock to check if there is still on-going requests. - while True: -@@ -1397,7 +1397,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): - async def continue_generation(self, obj: ContinueGenerationReqInput): - async with self.is_pause_cond: - self.is_pause = False -- await self.send_to_scheduler.send_pyobj(obj) -+ self.send_to_scheduler.send_pyobj(obj) - self.is_pause_cond.notify_all() - - async def update_weights_from_disk( -@@ -1965,25 +1965,23 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin): - priority = getattr(state.obj, "priority", None) - if priority is not None: - labels["priority"] = str(priority) -- if ( -- not state.ttft_observed -- and self.disaggregation_mode != DisaggregationMode.PREFILL -- ): -+ if not state.ttft_observed: - state.ttft_observed = True - state.last_completion_tokens = completion_tokens -- self.metrics_collector.observe_time_to_first_token( -- labels, state.time_stats.get_first_token_latency() -- ) -+ if self.disaggregation_mode != DisaggregationMode.PREFILL: -+ self.metrics_collector.observe_time_to_first_token( -+ labels, state.time_stats.get_first_token_latency() -+ ) - else: - num_new_tokens = completion_tokens - state.last_completion_tokens -- if num_new_tokens: -+ if num_new_tokens > 0: - self.metrics_collector.observe_inter_token_latency( - labels, - state.time_stats.get_interval(), - num_new_tokens, - ) - state.time_stats.set_last_time() -- state.last_completion_tokens = completion_tokens -+ state.last_completion_tokens = completion_tokens - - if state.finished: - retraction_count = ( -diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py -index 7f63610da8..fb56de1583 100644 ---- a/python/sglang/srt/managers/tp_worker.py -+++ b/python/sglang/srt/managers/tp_worker.py -@@ -29,6 +29,7 @@ from sglang.srt.managers.io_struct import ( - InitWeightsUpdateGroupReqInput, - LoadLoRAAdapterFromTensorsReqInput, - LoadLoRAAdapterReqInput, -+ PostProcessWeightsReqInput, - SendWeightsToRemoteInstanceReqInput, - UnloadLoRAAdapterReqInput, - UpdateWeightFromDiskReqInput, -@@ -170,6 +171,11 @@ class BaseTpWorker(ABC): - success, message = self.model_runner.update_weights_from_ipc(recv_req) - return success, message - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Perform optional post-processing on the updated model weights (e.g., Marlin conversion).""" -+ success, message = self.model_runner.post_process_weights(recv_req) -+ return success, message -+ - def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): - parameter = self.model_runner.get_weights_by_name( - recv_req.name, recv_req.truncate_size -diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py -index 3c1e97daab..e5128e5ee2 100644 ---- a/python/sglang/srt/mem_cache/hiradix_cache.py -+++ b/python/sglang/srt/mem_cache/hiradix_cache.py -@@ -755,9 +755,8 @@ class HiRadixCache(RadixCache): - self._update_leaf_status(node) - self._update_host_leaf_status(node) - if node.parent is None: -- assert ( -- node is self.root_node -- ), f"This request holds the node from another tree" -+ # Node belongs to a stale (flushed) tree — stop traversal gracefully. -+ break - node = node.parent - return DecLockRefResult(delta=delta) - -@@ -832,6 +831,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(node) - # update leaf status for the parent because the node is evicted - self._update_leaf_status(node.parent) -+ self._update_host_leaf_status(node.parent) - return num_evicted - - def _evict_regular(self, node: TreeNode): -@@ -1354,6 +1354,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(node) - # update parent status as a new leaf is added into device - self._update_leaf_status(node.parent) -+ self._update_host_leaf_status(node.parent) - else: - self._inc_hit_count(node, chunked) - total_prefix_length += prefix_len -@@ -1369,6 +1370,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(new_node) - # update parent status as a new leaf is added into device - self._update_leaf_status(new_node.parent) -+ self._update_host_leaf_status(new_node.parent) - else: - self._inc_hit_count(new_node, chunked) - total_prefix_length += prefix_len -diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py -index e4c158cda9..cf7333235f 100644 ---- a/python/sglang/srt/mem_cache/memory_pool.py -+++ b/python/sglang/srt/mem_cache/memory_pool.py -@@ -1854,9 +1854,12 @@ class NSATokenToKVPool(MLATokenToKVPool): - else: - assert self.page_size == 64 - with ( -- torch.cuda.use_mem_pool(self.custom_mem_pool) -- if self.custom_mem_pool -- else nullcontext() -+ ( -+ torch.cuda.use_mem_pool(self.custom_mem_pool) -+ if self.custom_mem_pool -+ else nullcontext() -+ ), -+ self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE), - ): - self.index_k_with_scale_buffer = [ - torch.zeros( -@@ -1878,6 +1881,11 @@ class NSATokenToKVPool(MLATokenToKVPool): - ) - for _ in range(layer_num) - ] -+ self.index_k_with_scale_buffer_ptrs = torch.tensor( -+ [x.data_ptr() for x in self.index_k_with_scale_buffer], -+ dtype=torch.uint64, -+ device=self.device, -+ ) - self._finalize_allocation_log(size) - - def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: -@@ -1960,6 +1968,50 @@ class NSATokenToKVPool(MLATokenToKVPool): - ] - return data_ptrs, data_lens, item_lens - -+ def get_cpu_copy(self, indices): -+ # First, save the kv_buffer (inherited from MLATokenToKVPool) -+ kv_cache_cpu = super().get_cpu_copy(indices) -+ -+ # Additionally, save the index_k_with_scale_buffer (page-indexed) -+ page_indices = indices[:: self.page_size] // self.page_size -+ torch.cuda.synchronize() -+ index_k_cpu = [] -+ chunk_size = self.cpu_offloading_chunk_size -+ # Convert chunk_size from token-level to page-level -+ page_chunk_size = max(1, chunk_size // self.page_size) -+ for layer_id in range(self.layer_num): -+ index_k_cpu.append([]) -+ for i in range(0, len(page_indices), page_chunk_size): -+ chunk_page_indices = page_indices[i : i + page_chunk_size] -+ idx_cpu = self.index_k_with_scale_buffer[layer_id][ -+ chunk_page_indices -+ ].to("cpu", non_blocking=True) -+ index_k_cpu[-1].append(idx_cpu) -+ torch.cuda.synchronize() -+ -+ return {"kv": kv_cache_cpu, "index_k": index_k_cpu} -+ -+ def load_cpu_copy(self, kv_cache_cpu_dict, indices): -+ # Restore the kv_buffer (inherited from MLATokenToKVPool) -+ super().load_cpu_copy(kv_cache_cpu_dict["kv"], indices) -+ -+ # Restore the index_k_with_scale_buffer (page-indexed) -+ page_indices = indices[:: self.page_size] // self.page_size -+ index_k_cpu = kv_cache_cpu_dict["index_k"] -+ torch.cuda.synchronize() -+ chunk_size = self.cpu_offloading_chunk_size -+ page_chunk_size = max(1, chunk_size // self.page_size) -+ for layer_id in range(self.layer_num): -+ for i in range(0, len(page_indices), page_chunk_size): -+ chunk_page_indices = page_indices[i : i + page_chunk_size] -+ idx_cpu = index_k_cpu[layer_id][i // page_chunk_size] -+ assert idx_cpu.shape[0] == len(chunk_page_indices) -+ idx_chunk = idx_cpu.to( -+ self.index_k_with_scale_buffer[0].device, non_blocking=True -+ ) -+ self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk -+ torch.cuda.synchronize() -+ - def get_kv_size_bytes(self): - kv_size_bytes = super().get_kv_size_bytes() - for index_k_cache in self.index_k_with_scale_buffer: -diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py -index 7d16160372..70fbdc702f 100644 ---- a/python/sglang/srt/mem_cache/radix_cache.py -+++ b/python/sglang/srt/mem_cache/radix_cache.py -@@ -512,7 +512,17 @@ class RadixCache(BasePrefixCache): - if self.disable: - return - -- token_ids = req.fill_ids -+ # Limit to kv_committed_len to avoid including tokens (e.g., the just-generated -+ # token in disagg prefill) that don't have computed KV yet. If fill_ids is longer -+ # than kv_committed_len, the extra tokens would produce stale values (0 from -+ # req_to_token_pool initialization), leading to spurious tree nodes and memory -+ # leak when page-aligned token counts happen to cross a page boundary. -+ kv_committed_len = req.kv_committed_len -+ token_ids = ( -+ req.fill_ids[:kv_committed_len] -+ if kv_committed_len < len(req.fill_ids) -+ else req.fill_ids -+ ) - kv_indices = self.req_to_token_pool.req_to_token[ - req.req_pool_idx, : len(token_ids) - ] -@@ -638,9 +648,8 @@ class RadixCache(BasePrefixCache): - node.lock_ref -= 1 - self._update_leaf_status(node) - if node.parent is None: -- assert ( -- node is self.root_node -- ), f"This request holds the node from another tree" -+ # Node belongs to a stale (flushed) tree — stop traversal gracefully. -+ break - node = node.parent - return DecLockRefResult(delta=delta) - -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index a59742b943..a7347c15b8 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -406,7 +406,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): - self.forward_stream = torch.get_device_module(self.device).Stream() - - # CPU offload -- set_offloader(create_offloader_from_server_args(server_args, dp_rank=dp_rank)) -+ # For draft worker (e.g., MTP), do not set offloader to avoid overriding -+ # the main model's offloader. Draft worker uses NoopOffloader instead. -+ if not is_draft_worker: -+ set_offloader( -+ create_offloader_from_server_args(server_args, dp_rank=dp_rank) -+ ) - - self._weight_checker = WeightChecker(model_runner=self) - -@@ -646,7 +651,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): - ) - - # Init routed experts capturer -- self.init_routed_experts_capturer() -+ if not self.is_draft_worker: -+ self.init_routed_experts_capturer() - - if self.device == "cuda" or self.device == "musa": - self.init_cublas() -@@ -2767,11 +2773,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): - output.expert_distribution_metrics = recorder_outputs.get("metrics") - - # Copy cached routing experts' buffers back to CPU cache -- get_global_experts_capturer().on_forward_end( -- forward_batch=forward_batch, -- can_run_graph=output.can_run_graph, -- cuda_graph_batch=getattr(self.graph_runner, "bs", None), -- ) -+ if not self.is_draft_worker: -+ # In speculative decoding, num_tokens_per_bs > 1, so we need to pass -+ # the actual number of tokens per dp rank in cuda graph, not batch size. -+ cuda_graph_num_tokens = None -+ if getattr(self.graph_runner, "bs", None): -+ cuda_graph_num_tokens = ( -+ self.graph_runner.bs * self.graph_runner.num_tokens_per_bs -+ ) -+ get_global_experts_capturer().on_forward_end( -+ forward_batch=forward_batch, -+ can_run_graph=output.can_run_graph, -+ cuda_graph_batch=cuda_graph_num_tokens, -+ ) - - if self.eplb_manager is not None: - self.eplb_manager.on_forward_pass_end() -@@ -3021,6 +3035,42 @@ class ModelRunner(ModelRunnerKVCacheMixin): - device=self.device, - ) - -+ def post_process_weights(self, recv_req): -+ """ -+ Execute post-processing logic for model weights, such as Marlin quantization format conversion. -+ """ -+ from sglang.srt.model_loader.loader import device_loading_context -+ -+ target_device = torch.device("cuda", torch.cuda.current_device()) -+ -+ if recv_req.restore_weights_before_load: -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports restoring weights -+ if quant_method is not None and hasattr( -+ quant_method, "restore_weights_before_loading" -+ ): -+ -+ with device_loading_context(module, target_device): -+ quant_method.restore_weights_before_loading(module) -+ -+ if recv_req.post_process_quantization: -+ # Iterate through all modules to apply specific post-loading processing -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports quantization post-processing -+ if quant_method is not None and hasattr( -+ quant_method, "process_weights_after_loading" -+ ): -+ -+ # Apply the post-processing (e.g., repacking weights for Marlin kernel) -+ with device_loading_context(module, target_device): -+ quant_method.process_weights_after_loading(module) -+ -+ return True, "Success" -+ - - def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]): - params_dict = dict(model.named_parameters()) -diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py -index 2f0074924d..1f991932c6 100644 ---- a/python/sglang/srt/models/glm4v_moe.py -+++ b/python/sglang/srt/models/glm4v_moe.py -@@ -52,11 +52,31 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - self.num_fused_shared_experts = 0 - self.determine_num_fused_shared_experts() - -- self.model = Glm4MoeModel( -- config, -- quant_config, -- prefix=add_prefix("language_model", prefix), -- ) -+ if not self.config.encoder_only: -+ self.model = Glm4MoeModel( -+ config, -+ quant_config, -+ prefix=add_prefix("language_model", prefix), -+ ) -+ -+ if self.pp_group.is_last_rank: -+ if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: -+ self.lm_head = self.model.embed_tokens -+ else: -+ self.lm_head = ParallelLMHead( -+ config.vocab_size, -+ config.hidden_size, -+ quant_config=quant_config, -+ prefix=add_prefix("lm_head", prefix), -+ use_attn_tp_group=get_global_server_args().enable_dp_lm_head, -+ ) -+ else: -+ # ranks other than the last rank will have a placeholder layer -+ self.lm_head = PPMissingLayer() -+ else: -+ # encoder_only mode: no language model, so no lm_head needed -+ self.lm_head = None -+ - self.visual = Glm4vVisionModel( - config.vision_config, - quant_config=quant_config, -@@ -64,24 +84,14 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - use_data_parallel=self.use_data_parallel, - ) - -- if self.pp_group.is_last_rank: -- if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: -- self.lm_head = self.model.embed_tokens -- else: -- self.lm_head = ParallelLMHead( -- config.vocab_size, -- config.hidden_size, -- quant_config=quant_config, -- prefix=add_prefix("lm_head", prefix), -- use_attn_tp_group=get_global_server_args().enable_dp_lm_head, -- ) -- else: -- # ranks other than the last rank will have a placeholder layer -- self.lm_head = PPMissingLayer() -- - self.logits_processor = LogitsProcessor(config) - self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) -- self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling -+ _rope_cfg = ( -+ getattr(self.config, "rope_scaling", None) -+ or getattr(self.config, "rope_parameters", None) -+ or {} -+ ) -+ self.is_mrope_enabled = "mrope_section" in _rope_cfg - - # For EAGLE3 support - self.capture_aux_hidden_states = False -@@ -219,6 +229,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue -+ # Skip loading visual/language model weights -+ if ( -+ self.config.encoder_only or self.config.language_only -+ ) and name not in params_dict: -+ continue - if name not in params_dict: - continue - -@@ -234,6 +249,8 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in name: - continue -+ if "visual" in name or self.config.encoder_only: -+ continue - - # Mark as expert weight regardless of whether we can process it - is_expert_weight = True -@@ -265,6 +282,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue -+ # Skip loading mm/language parameters -+ if ( -+ self.config.encoder_only or self.config.language_only -+ ) and name not in params_dict: -+ continue - if name not in params_dict: - continue - -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index 912891b6a7..fd67a7b580 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -325,7 +325,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - topk_output = self.topk(hidden_states, router_logits) - final_hidden_states = self.experts(hidden_states, topk_output) - -- if self.ep_size > 1 and not should_allreduce_fusion: -+ if self.ep_size > 1 and not should_allreduce_fusion and not use_reduce_scatter: - final_hidden_states = moe_expert_parallel_all_reduce(final_hidden_states) - - if ( -diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py -index 7746b24459..57b65fe06f 100644 ---- a/python/sglang/srt/models/qwen3_vl.py -+++ b/python/sglang/srt/models/qwen3_vl.py -@@ -1005,14 +1005,19 @@ class Qwen3LLMModel(Qwen3Model): - hidden_states + residual if residual is not None else hidden_states - ) - -+ deepstack_embeds = None -+ if input_deepstack_embeds is not None: -+ prev_layer_idx = layer_idx - 1 -+ if prev_layer_idx in self.deepstack_embed_to_decoder_layer: -+ sep = self.hidden_size * prev_layer_idx -+ deepstack_embeds = input_deepstack_embeds[ -+ :, sep : sep + self.hidden_size -+ ] -+ - # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. - # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 - # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack - # The order matters because addition with different tensors is not associative in practice. -- # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. -- deepstack_embeds = self.get_deepstack_embeds( -- layer_idx - 1, input_deepstack_embeds -- ) - hidden_states, residual = layer( - positions, - hidden_states, -diff --git a/python/sglang/srt/multimodal/processors/glm4v.py b/python/sglang/srt/multimodal/processors/glm4v.py -index a44f14b6ca..6d6c65ea49 100644 ---- a/python/sglang/srt/multimodal/processors/glm4v.py -+++ b/python/sglang/srt/multimodal/processors/glm4v.py -@@ -1,7 +1,13 @@ - from typing import List, Union - -+import torch -+ - from sglang.srt.layers.rotary_embedding import MRotaryEmbedding --from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput -+from sglang.srt.managers.schedule_batch import ( -+ Modality, -+ MultimodalDataItem, -+ MultimodalProcessorOutput, -+) - from sglang.srt.models.glm4v import Glm4vForConditionalGeneration - from sglang.srt.models.glm4v_moe import Glm4vMoeForConditionalGeneration - from sglang.srt.multimodal.processors.base_processor import ( -@@ -46,6 +52,8 @@ class Glm4vImageProcessor(SGLangBaseProcessor): - self.IMAGE_END_TOKEN_ID = hf_config.image_end_token_id - self.VIDEO_START_TOKEN_ID = hf_config.video_start_token_id - self.VIDEO_END_TOKEN_ID = hf_config.video_end_token_id -+ self.IM_START_TOKEN_ID = self.IMAGE_START_TOKEN_ID -+ self.IM_END_TOKEN_ID = self.IMAGE_END_TOKEN_ID - - # Vision config - self.IMAGE_FACTOR = 28 -@@ -60,6 +68,39 @@ class Glm4vImageProcessor(SGLangBaseProcessor): - video_token_id=self.IM_TOKEN_ID, - ).build(_processor) - -+ def get_mm_data(self, prompt, embeddings, img_grid_thw): -+ input_ids, offsets, _ = self.build_input_ids(prompt, img_grid_thw=img_grid_thw) -+ image_embeddings = ( -+ embeddings.get(Modality.IMAGE, embeddings) -+ if isinstance(embeddings, dict) -+ else embeddings -+ ) -+ mm_items = [ -+ MultimodalDataItem( -+ modality=Modality.IMAGE, -+ offsets=offsets, -+ precomputed_embeddings=image_embeddings, -+ ) -+ ] -+ -+ mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index_glm4v( -+ input_ids=torch.tensor(input_ids, dtype=torch.long).unsqueeze(0), -+ hf_config=self.hf_config, -+ image_grid_thw=img_grid_thw, -+ video_grid_thw=None, -+ attention_mask=None, -+ ) -+ -+ return MultimodalProcessorOutput( -+ input_ids=input_ids, -+ mm_items=mm_items, -+ im_start_id=self.IM_START_TOKEN_ID, -+ im_end_id=self.IM_END_TOKEN_ID, -+ im_token_id=self.IM_TOKEN_ID, -+ mrope_positions=mrope_positions.squeeze(1), -+ mrope_position_delta=mrope_position_delta, -+ ) -+ - def compute_mrope_positions(self, input_ids, mm_items): - image_grid_thw = None - video_grid_thw = None -diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py -index 3f102567d0..6fb3899021 100644 ---- a/python/sglang/srt/multimodal/processors/qwen_vl.py -+++ b/python/sglang/srt/multimodal/processors/qwen_vl.py -@@ -499,7 +499,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): - **kwargs, - ): - entry_time = time.perf_counter() -- base_output = self.load_mm_data( -+ base_output = self.legacy_load_mm_data( - prompt=input_text, - image_data=image_data, - video_data=request_obj.video_data, -diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py -index 8caf21c320..51d1edc584 100644 ---- a/python/sglang/srt/observability/req_time_stats.py -+++ b/python/sglang/srt/observability/req_time_stats.py -@@ -21,7 +21,10 @@ import uuid - from dataclasses import dataclass, field - from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union - --from sglang.srt.disaggregation.utils import DisaggregationMode -+from sglang.srt.disaggregation.utils import ( -+ DisaggregationMode, -+ is_slime_profiling_enabled, -+) - from sglang.srt.model_executor.forward_batch_info import ForwardMode - from sglang.srt.observability.metrics_collector import ( - SchedulerMetricsCollector, -@@ -553,6 +556,14 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): - transfer_total_mb: float = 0.0 - # Number of prefill retries for this request - prefill_retry_count: int = 0 -+ fwd_prefill_bootstrap_queue_duration: Optional[float] = None -+ fwd_prefill_forward_duration: Optional[float] = None -+ fwd_prefill_transfer_queue_duration: Optional[float] = None -+ fwd_bootstrap_duration: Optional[float] = None -+ fwd_alloc_waiting_duration: Optional[float] = None -+ fwd_transfer_speed_gb_s: Optional[float] = None -+ fwd_transfer_total_mb: Optional[float] = None -+ fwd_prefill_retry_count: Optional[int] = None - - def __getstate__(self) -> object: - # send to detokenizer/tokenizer -@@ -560,11 +571,33 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): - return {} - - state = { -+ "disagg_mode": self.disagg_mode, - "wait_queue_entry_time": self.wait_queue_entry_time, - "forward_entry_time": self.forward_entry_time, - "prefill_run_batch_start_time": self.prefill_run_batch_start_time, - "prefill_run_batch_end_time": self.prefill_run_batch_end_time, - "prefill_finished_time": self.prefill_finished_time, -+ "completion_time": self.completion_time, -+ "prefill_bootstrap_queue_entry_time": ( -+ self.prefill_bootstrap_queue_entry_time -+ ), -+ "prefill_transfer_queue_entry_time": self.prefill_transfer_queue_entry_time, -+ "decode_prealloc_queue_entry_time": self.decode_prealloc_queue_entry_time, -+ "decode_transfer_queue_entry_time": self.decode_transfer_queue_entry_time, -+ "bootstrap_done_time": self.bootstrap_done_time, -+ "transfer_speed_gb_s": self.transfer_speed_gb_s, -+ "transfer_total_mb": self.transfer_total_mb, -+ "prefill_retry_count": self.prefill_retry_count, -+ "fwd_prefill_bootstrap_queue_duration": ( -+ self.fwd_prefill_bootstrap_queue_duration -+ ), -+ "fwd_prefill_forward_duration": self.fwd_prefill_forward_duration, -+ "fwd_prefill_transfer_queue_duration": self.fwd_prefill_transfer_queue_duration, -+ "fwd_bootstrap_duration": self.fwd_bootstrap_duration, -+ "fwd_alloc_waiting_duration": self.fwd_alloc_waiting_duration, -+ "fwd_transfer_speed_gb_s": self.fwd_transfer_speed_gb_s, -+ "fwd_transfer_total_mb": self.fwd_transfer_total_mb, -+ "fwd_prefill_retry_count": self.fwd_prefill_retry_count, - "diff_realtime_monotonic": global_diff_realtime_monotonic, - } - return state -@@ -916,6 +949,149 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): - return self.prefill_run_batch_end_time - self.prefill_run_batch_start_time - return None - -+ def get_pd_prefill_bootstrap_queue_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_bootstrap_queue_duration is not None: -+ return self.fwd_prefill_bootstrap_queue_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.prefill_bootstrap_queue_entry_time > 0.0 -+ and self.wait_queue_entry_time > 0.0 -+ ): -+ return self.wait_queue_entry_time - self.prefill_bootstrap_queue_entry_time -+ return None -+ -+ def get_pd_prefill_forward_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_forward_duration is not None: -+ return self.fwd_prefill_forward_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.forward_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.forward_entry_time -+ return None -+ -+ def get_pd_prefill_transfer_queue_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_transfer_queue_duration is not None: -+ return self.fwd_prefill_transfer_queue_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.prefill_transfer_queue_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.prefill_transfer_queue_entry_time -+ return None -+ -+ def get_pd_decode_prealloc_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_prealloc_queue_entry_time > 0.0 -+ and self.decode_transfer_queue_entry_time > 0.0 -+ ): -+ return ( -+ self.decode_transfer_queue_entry_time -+ - self.decode_prealloc_queue_entry_time -+ ) -+ return None -+ -+ def get_pd_decode_transfer_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_transfer_queue_entry_time > 0.0 -+ and self.wait_queue_entry_time > 0.0 -+ ): -+ return self.wait_queue_entry_time - self.decode_transfer_queue_entry_time -+ return None -+ -+ def get_pd_decode_forward_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.forward_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.forward_entry_time -+ return None -+ -+ def get_pd_bootstrap_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_bootstrap_duration is not None: -+ return self.fwd_bootstrap_duration -+ if self.bootstrap_done_time <= 0.0: -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.prefill_bootstrap_queue_entry_time > 0.0 -+ ): -+ return self.bootstrap_done_time - self.prefill_bootstrap_queue_entry_time -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_prealloc_queue_entry_time > 0.0 -+ ): -+ return self.bootstrap_done_time - self.decode_prealloc_queue_entry_time -+ return None -+ -+ def get_pd_alloc_waiting_duration(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_alloc_waiting_duration is not None: -+ return self.fwd_alloc_waiting_duration -+ if self.bootstrap_done_time <= 0.0: -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.wait_queue_entry_time > 0.0 -+ ): -+ return self.wait_queue_entry_time - self.bootstrap_done_time -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_transfer_queue_entry_time > 0.0 -+ ): -+ return self.decode_transfer_queue_entry_time - self.bootstrap_done_time -+ return None -+ -+ def get_pd_transfer_speed_gb_s(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_transfer_speed_gb_s is not None: -+ return self.fwd_transfer_speed_gb_s -+ if ( -+ self.disagg_mode != DisaggregationMode.NULL -+ and self.transfer_speed_gb_s > 0.0 -+ ): -+ return self.transfer_speed_gb_s -+ return None -+ -+ def get_pd_transfer_total_mb(self) -> Optional[float]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_transfer_total_mb is not None: -+ return self.fwd_transfer_total_mb -+ if self.disagg_mode != DisaggregationMode.NULL and self.transfer_total_mb > 0.0: -+ return self.transfer_total_mb -+ return None -+ -+ def get_pd_prefill_retry_count(self) -> Optional[int]: -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_retry_count is not None: -+ return self.fwd_prefill_retry_count -+ if self.disagg_mode == DisaggregationMode.PREFILL: -+ return self.prefill_retry_count -+ return None -+ - def convert_to_duration(self) -> str: - if self.disagg_mode == DisaggregationMode.NULL: - queue_duration = self.forward_entry_time - self.wait_queue_entry_time -@@ -1038,6 +1214,24 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): - "prefill_launch_latency": self.get_prefill_launch_latency(), - } - ) -+ if is_slime_profiling_enabled(): -+ for key, value in { -+ "pd_prefill_bootstrap_queue_duration": ( -+ self.get_pd_prefill_bootstrap_queue_duration() -+ ), -+ "pd_prefill_forward_duration": self.get_pd_prefill_forward_duration(), -+ "pd_prefill_transfer_queue_duration": self.get_pd_prefill_transfer_queue_duration(), -+ "pd_decode_prealloc_duration": self.get_pd_decode_prealloc_duration(), -+ "pd_decode_transfer_duration": self.get_pd_decode_transfer_duration(), -+ "pd_decode_forward_duration": self.get_pd_decode_forward_duration(), -+ "pd_bootstrap_duration": self.get_pd_bootstrap_duration(), -+ "pd_alloc_waiting_duration": self.get_pd_alloc_waiting_duration(), -+ "pd_transfer_speed_gb_s": self.get_pd_transfer_speed_gb_s(), -+ "pd_transfer_total_mb": self.get_pd_transfer_total_mb(), -+ "pd_prefill_retry_count": self.get_pd_prefill_retry_count(), -+ }.items(): -+ if value is not None: -+ meta_data[key] = value - return meta_data - - def format_duration(self, duration: float) -> str: -diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py -index ff5695ce2e..588379a85d 100644 ---- a/python/sglang/srt/observability/scheduler_metrics_mixin.py -+++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py -@@ -883,12 +883,42 @@ class SchedulerMetricsMixin: - num_tokens += sum(req.seqlen for queue in waiting_queues for req in queue) - num_waiting_reqs = sum(len(queue) for queue in waiting_queues) - -+ queue_names = ["waiting_queue"] -+ if self.disaggregation_mode == DisaggregationMode.PREFILL: -+ queue_names.append("bootstrap_queue") -+ elif self.disaggregation_mode == DisaggregationMode.DECODE: -+ queue_names.append("prealloc_queue") -+ queue_names.append("transfer_queue") -+ queue_names.append("retracted_queue") -+ -+ queue_details = [] -+ for name, queue in zip(queue_names, waiting_queues): -+ reqs_info = [{"seqlen": req.seqlen} for req in queue] -+ queue_details.append( -+ { -+ "name": name, -+ "num_reqs": len(queue), -+ "num_tokens": sum(req_info["seqlen"] for req_info in reqs_info), -+ "reqs": reqs_info, -+ } -+ ) -+ -+ running_reqs_info = [ -+ {"seqlen": req.seqlen} for req in self.running_batch.reqs -+ ] -+ running_details = { -+ "num_reqs": len(self.running_batch.reqs), -+ "reqs": running_reqs_info, -+ } -+ - return GetLoadReqOutput( - dp_rank=self.dp_rank, - num_reqs=len(self.running_batch.reqs) + num_waiting_reqs, - num_waiting_reqs=num_waiting_reqs, - num_tokens=num_tokens, - ts_tic=time.perf_counter(), -+ queue_details=queue_details, -+ running_details=running_details, - ) - - def get_loads(self: Scheduler, req: GetLoadsReqInput = None) -> GetLoadsReqOutput: -diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index d91ced805f..4c8774bb64 100644 ---- a/python/sglang/srt/server_args.py -+++ b/python/sglang/srt/server_args.py -@@ -670,6 +670,7 @@ class ServerArgs: - # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 - enable_nsa_prefill_context_parallel: bool = False - nsa_prefill_cp_mode: str = "round-robin-split" -+ disable_indexer_rope_neox_style: bool = False - enable_fused_qk_norm_rope: bool = False - enable_precise_embedding_interpolation: bool = False - enable_fused_moe_sum_all_reduce: bool = False -@@ -5659,6 +5660,12 @@ class ServerArgs: - help="Token splitting mode for the prefill phase of DeepSeek v3.2 under context parallelism. Optional values: 'round-robin-split'(default), 'in-seq-split' " - "'round-robin-split' distributes tokens across ranks based on token_idx %% cp_size. It supports multi-batch prefill, fused MoE, and FP8 KV cache.", - ) -+ parser.add_argument( -+ "--disable-indexer-rope-neox-style", -+ action="store_true", -+ help="Disable NSA indexer RoPE neox style (equivalent to INDEXER_ROPE_NEOX_STYLE=0). " -+ "If the environment variable INDEXER_ROPE_NEOX_STYLE is also set and conflicts, an error is raised.", -+ ) - parser.add_argument( - "--enable-prefill-context-parallel", - action="store_true", -diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -index 40e859b2d6..2604ae037c 100644 ---- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -+++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -@@ -377,6 +377,10 @@ class EAGLEDraftCudaGraphRunner: - buffers.seq_lens.fill_(self.seq_len_fill_value) - buffers.out_cache_loc.zero_() - buffers.positions.zero_() -+ buffers.topk_p.zero_() -+ buffers.topk_index.zero_() -+ buffers.hidden_states.zero_() -+ buffers.req_pool_indices.zero_() - - num_tokens = bs * self.num_tokens_per_bs - -@@ -386,8 +390,12 @@ class EAGLEDraftCudaGraphRunner: - forward_batch.out_cache_loc - ) - buffers.positions[:raw_num_token].copy_(forward_batch.positions) -- buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) -- buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) -+ buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p.clamp(0, 1)) -+ buffers.topk_index[:raw_bs].copy_( -+ forward_batch.spec_info.topk_index.clamp( -+ 0, self.model_runner.model_config.vocab_size - 1 -+ ) -+ ) - buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) - buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) - -diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index dbb91f555e..a04caefc34 100644 ---- a/python/sglang/srt/speculative/eagle_info.py -+++ b/python/sglang/srt/speculative/eagle_info.py -@@ -776,6 +776,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] - self.verified_id = self.verified_id[: len(new_indices)] -+ if self.accept_length is not None: -+ self.accept_length = self.accept_length[: len(new_indices)] -+ if self.accept_length_cpu is not None: -+ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] - else: - # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` - self.topk_p = self.topk_p[new_indices] -@@ -807,6 +811,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) - self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) - self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) -+ if self.accept_length is not None and spec_info.accept_length is not None: -+ self.accept_length = torch.cat( -+ [self.accept_length, spec_info.accept_length] -+ ) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif self.accept_length is not None: -+ zeros = torch.zeros( -+ [spec_info.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([self.accept_length, zeros]) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif spec_info.accept_length is not None: -+ zeros = torch.zeros( -+ [self.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([zeros, spec_info.accept_length]) -+ self.accept_length_cpu = self.accept_length.tolist() - - - @dataclass -diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py -index b0be70d751..44a78d684e 100644 ---- a/python/sglang/srt/utils/common.py -+++ b/python/sglang/srt/utils/common.py -@@ -2157,6 +2157,7 @@ class SafeUnpickler(pickle.Unpickler): - "sglang.srt.layers.", - "sglang.srt.utils.", - "torch_npu.", -+ "slime.", - } - - DENY_CLASSES = { -diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py -index 3be16446e0..1b2371c839 100644 ---- a/python/sglang/srt/utils/weight_checker.py -+++ b/python/sglang/srt/utils/weight_checker.py -@@ -69,6 +69,9 @@ def _check_tensors( - actual_should_compare, - actual, - ) in zip(expect_tensors, actual_tensors, strict=True): -+ if ".cos_sin_cache" in expect_name: -+ # skip cos/sin cache which is deterministic from shape and dtype and may have different shapes due to different implementations. -+ continue - assert expect_name == actual_name, f"{expect_name=} {actual_name=}" - assert ( - expect_should_compare == actual_should_compare diff --git a/docker/patch/v0.5.0rc0-cu126/sglang.patch b/docker/patch/v0.5.0rc0-cu126/sglang.patch deleted file mode 100644 index 990c2e628..000000000 --- a/docker/patch/v0.5.0rc0-cu126/sglang.patch +++ /dev/null @@ -1,203 +0,0 @@ -diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py -index bdb124e51..3edf30ab1 100644 ---- a/python/sglang/srt/configs/model_config.py -+++ b/python/sglang/srt/configs/model_config.py -@@ -454,14 +454,14 @@ class ModelConfig: - ).lower() - - # Detect which checkpoint is it -- for _, method in QUANTIZATION_METHODS.items(): -- quantization_override = method.override_quantization_method( -- quant_cfg, self.quantization -- ) -- if quantization_override: -- quant_method = quantization_override -- self.quantization = quantization_override -- break -+ # for _, method in QUANTIZATION_METHODS.items(): -+ # quantization_override = method.override_quantization_method( -+ # quant_cfg, self.quantization -+ # ) -+ # if quantization_override: -+ # quant_method = quantization_override -+ # self.quantization = quantization_override -+ # break - - # Verify quantization configurations. - if self.quantization is None: -diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py -index 2dd2c75f1..f2adb18f8 100644 ---- a/python/sglang/srt/entrypoints/http_server.py -+++ b/python/sglang/srt/entrypoints/http_server.py -@@ -264,6 +264,10 @@ async def validate_json_request(raw_request: Request): - - - @app.get("/health") -+async def health(request: Request) -> Response: -+ return Response(status_code=200) -+ -+ - @app.get("/health_generate") - async def health_generate(request: Request) -> Response: - """ -diff --git a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py -index 372717bf9..40665cc90 100644 ---- a/python/sglang/srt/layers/moe/token_dispatcher/deepep.py -+++ b/python/sglang/srt/layers/moe/token_dispatcher/deepep.py -@@ -190,6 +190,7 @@ class DeepEPBuffer: - f"Consider using --deepep-config to change the behavior." - ) - -+ num_qps_per_rank = 20 - cls._buffer = Buffer( - group, - num_nvl_bytes, -diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py -index 956264fc9..69f729336 100644 ---- a/python/sglang/srt/layers/quantization/fp8.py -+++ b/python/sglang/srt/layers/quantization/fp8.py -@@ -351,10 +351,10 @@ class Fp8LinearMethod(LinearMethodBase): - return - else: - weight, weight_scale = layer.weight.data, layer.weight_scale_inv.data -- layer.weight = torch.nn.Parameter(weight, requires_grad=False) -- layer.weight_scale_inv = torch.nn.Parameter( -- weight_scale, requires_grad=False -- ) -+ # layer.weight = torch.nn.Parameter(weight, requires_grad=False) -+ # layer.weight_scale_inv = torch.nn.Parameter( -+ # weight_scale, requires_grad=False -+ # ) - return - - layer.weight = torch.nn.Parameter(layer.weight.data, requires_grad=False) -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index 95a529c89..758fbfd5f 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -1359,7 +1359,7 @@ class Scheduler( - - if memory_leak: - msg = "token_to_kv_pool_allocator memory leak detected! " f"{token_msg}" -- raise ValueError(msg) -+ # raise ValueError(msg) - - if self.disaggregation_mode == DisaggregationMode.DECODE: - req_total_size = ( -@@ -1374,7 +1374,7 @@ class Scheduler( - f"available_size={len(self.req_to_token_pool.free_slots)}, " - f"total_size={self.req_to_token_pool.size}\n" - ) -- raise ValueError(msg) -+ # raise ValueError(msg) - - if ( - self.enable_metrics -@@ -1830,6 +1830,7 @@ class Scheduler( - deepep_mode=DeepEPMode(self.server_args.deepep_mode), - require_mlp_tp_gather=require_mlp_tp_gather(self.server_args), - disable_overlap_schedule=self.server_args.disable_overlap_schedule, -+ offload_tags=self.offload_tags, - ) - - def handle_dp_balance_data(self, local_batch: ScheduleBatch): -@@ -1927,6 +1928,7 @@ class Scheduler( - deepep_mode: DeepEPMode, - require_mlp_tp_gather: bool, - disable_overlap_schedule: bool, -+ offload_tags: set[str], - ): - # Check if other DP workers have running batches - if local_batch is None: -@@ -1957,7 +1959,7 @@ class Scheduler( - ) - - tbo_preparer = TboDPAttentionPreparer() -- if disable_overlap_schedule: -+ if len(offload_tags) == 0 and disable_overlap_schedule: - group = tp_group.device_group - device = tp_group.device - else: -diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index 58220b1d6..3c3d081a8 100644 ---- a/python/sglang/srt/managers/tokenizer_manager.py -+++ b/python/sglang/srt/managers/tokenizer_manager.py -@@ -1044,10 +1044,15 @@ class TokenizerManager: - request: Optional[fastapi.Request] = None, - ) -> Tuple[bool, str]: - self.auto_create_handle_loop() -- assert ( -- self.server_args.dp_size == 1 -- ), "dp_size must be 1 for init parameter update group" -- result = (await self.init_weights_update_group_communicator(obj))[0] -+ results = await self.init_weights_update_group_communicator(obj) -+ if self.server_args.dp_size == 1: -+ result = results[0] -+ return result.success, result.message -+ else: -+ all_success = all([r.success for r in results]) -+ all_message = [r.message for r in results] -+ all_message = " | ".join(all_message) -+ return all_success, all_message - return result.success, result.message - - async def update_weights_from_distributed( -@@ -1056,9 +1061,6 @@ class TokenizerManager: - request: Optional[fastapi.Request] = None, - ) -> Tuple[bool, str]: - self.auto_create_handle_loop() -- assert ( -- self.server_args.dp_size == 1 or self.server_args.enable_dp_attention -- ), "dp_size must be 1 or dp attention must be enabled for update weights from distributed" - - if obj.abort_all_requests: - self.abort_request(abort_all=True) -@@ -1066,8 +1068,15 @@ class TokenizerManager: - # This means that weight sync - # cannot run while requests are in progress. - async with self.model_update_lock.writer_lock: -- result = (await self.update_weights_from_distributed_communicator(obj))[0] -- return result.success, result.message -+ results = await self.update_weights_from_distributed_communicator(obj) -+ if self.server_args.dp_size == 1: -+ result = results[0] -+ return result.success, result.message -+ else: -+ all_success = all([r.success for r in results]) -+ all_message = [r.message for r in results] -+ all_message = " | ".join(all_message) -+ return all_success, all_message - - async def update_weights_from_tensor( - self, -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 5222bff0a..ff0bbc62a 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -22,6 +22,7 @@ import os - import time - from dataclasses import dataclass - from typing import List, Optional, Tuple, Union -+from contextlib import nullcontext - - import torch - import torch.distributed as dist -@@ -675,7 +676,7 @@ class ModelRunner: - monkey_patch_vllm_parallel_state() - monkey_patch_isinstance_for_vllm_base_layer() - -- with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_WEIGHTS): -+ with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_WEIGHTS) if not self.is_draft_worker else nullcontext(): - self.model = get_model( - model_config=self.model_config, - load_config=self.load_config, -diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py -index e0f0b373d..a18ac10f1 100644 ---- a/python/sglang/srt/models/glm4_moe.py -+++ b/python/sglang/srt/models/glm4_moe.py -@@ -1108,5 +1108,4 @@ class Glm4MoeForCausalLM(DeepseekV2ForCausalLM): - ) - weight_loader(param, loaded_weight) - -- - EntryClass = [Glm4MoeForCausalLM] diff --git a/docker/patch/v0.5.5.post1/sglang.patch b/docker/patch/v0.5.5.post1/sglang.patch deleted file mode 100644 index 910af9fc1..000000000 --- a/docker/patch/v0.5.5.post1/sglang.patch +++ /dev/null @@ -1,1492 +0,0 @@ -diff --git a/python/sglang/srt/distributed/device_communicators/pynccl.py b/python/sglang/srt/distributed/device_communicators/pynccl.py -index f485c24c2..901010610 100644 ---- a/python/sglang/srt/distributed/device_communicators/pynccl.py -+++ b/python/sglang/srt/distributed/device_communicators/pynccl.py -@@ -352,3 +352,9 @@ class PyNcclCommunicator: - - self.disabled = old_disable - self.stream = old_stream -+ -+ def nccl_pause(self): -+ self.nccl.ncclPause(self.comm) -+ -+ def nccl_resume(self): -+ self.nccl.ncclResume(self.comm) -diff --git a/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py b/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -index 579811777..3c0854550 100644 ---- a/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -+++ b/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -@@ -304,6 +304,15 @@ class NCCLLibrary: - Function("ncclGroupEnd", ncclResult_t, []), - ] - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ exported_functions.extend([ -+ # ncclResult_t ncclPause(ncclComm_t comm); -+ Function("ncclPause", ncclResult_t, [ncclComm_t]), -+ # ncclResult_t ncclResume(ncclComm_t comm); -+ Function("ncclResume", ncclResult_t, [ncclComm_t]), -+ Function("ncclSetGroupID", ncclResult_t, [ctypes.c_int]), -+ ]) -+ - exported_functions_symm_mem = [ - # ncclResult_t ncclCommWindowRegister(ncclComm_t comm, void* buff, size_t size, ncclWindow_t* win, int winFlags); - Function( -@@ -551,6 +560,12 @@ class NCCLLibrary: - def ncclGroupEnd(self) -> None: - self.NCCL_CHECK(self._funcs["ncclGroupEnd"]()) - -+ def ncclPause(self, comm: ncclComm_t) -> None: -+ self.NCCL_CHECK(self._funcs["ncclPause"](comm)) -+ -+ def ncclResume(self, comm: ncclComm_t) -> None: -+ self.NCCL_CHECK(self._funcs["ncclResume"](comm)) -+ - - __all__ = [ - "NCCLLibrary", -diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py -index c954d1e52..c5d2067b2 100644 ---- a/python/sglang/srt/distributed/parallel_state.py -+++ b/python/sglang/srt/distributed/parallel_state.py -@@ -1758,7 +1758,10 @@ def get_tensor_model_parallel_world_size(): - - def get_tensor_model_parallel_rank(): - """Return my rank for the tensor model parallel group.""" -- return get_tp_group().rank_in_group -+ try: -+ return get_tp_group().rank_in_group -+ except Exception: -+ return 0 - - - def get_pipeline_model_parallel_world_size(): -diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py -index ebab42a8f..05b2cb466 100644 ---- a/python/sglang/srt/entrypoints/engine.py -+++ b/python/sglang/srt/entrypoints/engine.py -@@ -179,6 +179,7 @@ class Engine(EngineBase): - lora_path: Optional[List[Optional[str]]] = None, - custom_logit_processor: Optional[Union[List[str], str]] = None, - return_hidden_states: bool = False, -+ return_routed_experts: bool = False, - stream: bool = False, - bootstrap_host: Optional[Union[List[str], str]] = None, - bootstrap_port: Optional[Union[List[int], int]] = None, -@@ -213,6 +214,7 @@ class Engine(EngineBase): - lora_path=lora_path, - custom_logit_processor=custom_logit_processor, - return_hidden_states=return_hidden_states, -+ return_routed_experts=return_routed_experts, - stream=stream, - bootstrap_host=bootstrap_host, - bootstrap_port=bootstrap_port, -diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py -index 76e999498..c098f1d37 100644 ---- a/python/sglang/srt/entrypoints/http_server.py -+++ b/python/sglang/srt/entrypoints/http_server.py -@@ -403,6 +403,10 @@ async def validate_json_request(raw_request: Request): - - - @app.get("/health") -+async def health(request: Request) -> Response: -+ return Response(status_code=200) -+ -+ - @app.get("/health_generate") - async def health_generate(request: Request) -> Response: - """ -diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py -index 7569f2b97..94084bb64 100644 ---- a/python/sglang/srt/layers/layernorm.py -+++ b/python/sglang/srt/layers/layernorm.py -@@ -88,15 +88,12 @@ class RMSNorm(CustomOp): - eps: float = 1e-6, - var_hidden_size: Optional[int] = None, - cast_x_before_out_mul: bool = False, -- fp32_residual: bool = False, -- weight_dtype: Optional = None, -- override_orig_dtype: Optional = None, -+ fp32_residual: bool = True, - ) -> None: - super().__init__() - self.cast_x_before_out_mul = cast_x_before_out_mul - self.fp32_residual = fp32_residual -- self.override_orig_dtype = override_orig_dtype -- self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) -+ self.weight = nn.Parameter(torch.ones(hidden_size)) - self.variance_epsilon = eps - self.hidden_size = hidden_size - self.variance_size_override = ( -@@ -195,14 +192,15 @@ class RMSNorm(CustomOp): - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - if not x.is_contiguous(): - x = x.contiguous() -- orig_dtype = self.override_orig_dtype or x.dtype -+ orig_dtype = x.dtype -+ -+ if residual is not None and not self.fp32_residual: -+ x = x + residual -+ residual = x.clone() - x = x.to(torch.float32) -- if residual is not None: -+ if residual is not None and self.fp32_residual: - x = x + residual.to(torch.float32) -- if self.fp32_residual: -- residual = x.clone() -- else: -- residual = x.to(orig_dtype) -+ residual = x.to(orig_dtype) - - hidden_size = x.shape[-1] - if hidden_size != self.hidden_size: -diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py -index e2c7d2ab6..de44951b5 100644 ---- a/python/sglang/srt/layers/logits_processor.py -+++ b/python/sglang/srt/layers/logits_processor.py -@@ -824,11 +824,6 @@ class LogitsProcessor(nn.Module): - None, # bias - True, # is_vnni - ) -- elif get_global_server_args().rl_on_policy_target is not None: -- # Due to tie-weight, we may not be able to change lm_head's weight dtype -- logits = torch.matmul( -- hidden_states.bfloat16(), lm_head.weight.T.bfloat16() -- ) - else: - logits = torch.matmul( - hidden_states.to(lm_head.weight.dtype), lm_head.weight.T -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -index b92f2159e..1846128be 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -@@ -13,6 +13,7 @@ import torch - import triton.language as tl - - from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - cpu_has_amx_support, - direct_register_custom_op, -@@ -607,7 +608,10 @@ def fused_experts_impl( - ).squeeze(dim=1) - else: - # According to micro benchmark results, torch.compile can get better performance for small token. -- if tokens_in_chunk <= 32: -+ if ( -+ not get_global_server_args().enable_deterministic_inference -+ and tokens_in_chunk <= 32 -+ ): - moe_sum_reduce_torch_compile( - intermediate_cache3.view(*intermediate_cache3.shape), - out_hidden_states[begin_chunk_idx:end_chunk_idx], -diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py -new file mode 100644 -index 000000000..a15b73501 ---- /dev/null -+++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py -@@ -0,0 +1,206 @@ -+import logging -+from abc import ABC -+from typing import Optional -+ -+import numpy as np -+import torch -+ -+from sglang.srt.configs.model_config import ModelConfig -+from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -+from sglang.srt.server_args import get_global_server_args -+ -+logger = logging.getLogger(__name__) -+ -+_GB = 1024 * 1024 * 1024 -+_MB = 1024 * 1024 -+ -+ -+def get_tensor_size_bytes(t: torch.Tensor): -+ return np.prod(t.shape) * t.dtype.itemsize -+ -+ -+class _RoutedExpertsDeviceCache: -+ def __init__( -+ self, model_config: ModelConfig, max_running_requests: int, device: str -+ ) -> None: -+ self.buffer = torch.zeros( -+ ( -+ max( -+ get_global_server_args().chunked_prefill_size, max_running_requests -+ ), -+ model_config.hf_text_config.num_hidden_layers, -+ model_config.hf_text_config.num_experts_per_tok, -+ ), -+ dtype=torch.int32, -+ device=device, -+ ) -+ self._finalize_allocation_log() -+ -+ def get_buffer_size_bytes(self): -+ assert hasattr(self, "buffer") -+ return get_tensor_size_bytes(self.buffer) -+ -+ def capture_fwd_routed_experts(self, layer_id: int, topk_ids: torch.Tensor): -+ assert layer_id is not None, "capturing routing experts but get layer_id None" -+ batch, _ = topk_ids.shape -+ self.buffer[:batch, layer_id, :] = topk_ids -+ -+ def _finalize_allocation_log(self): -+ """Common logging and memory usage computation for captured experts buffers.""" -+ buffer_size_MB = self.get_buffer_size_bytes() / _MB -+ logger.info( -+ f"Routing experts device buffer allocated. #shape: {tuple(self.buffer.shape)}, size: {buffer_size_MB:.2f} MB" -+ ) -+ -+ -+class _RoutedExpertsHostCache: -+ def __init__( -+ self, -+ model_config: ModelConfig, -+ num_tokens: int, -+ ) -> None: -+ self.num_tokens = num_tokens -+ self.buffer = torch.zeros( -+ ( -+ num_tokens, -+ model_config.hf_text_config.num_hidden_layers, -+ model_config.hf_text_config.num_experts_per_tok, -+ ), -+ dtype=torch.int32, -+ device="cpu", -+ ) -+ self._finalize_allocation_log() -+ -+ def get_buffer_size_bytes(self): -+ assert hasattr(self, "buffer") -+ return get_tensor_size_bytes(self.buffer) -+ -+ def set_experts_buffer(self, layer_id: int, loc: torch.Tensor, top_k: torch.Tensor): -+ self.buffer[layer_id, loc, :] = top_k.cpu() -+ -+ def _finalize_allocation_log(self): -+ """Common logging and memory usage computation for captured experts buffers.""" -+ buffer_size_GB = self.get_buffer_size_bytes() / _GB -+ logger.info( -+ f"Routing experts host buffer allocated. #tokens: {self.num_tokens}, size: {buffer_size_GB:.2f} GB" -+ ) -+ -+ -+class RoutedExpertsCapturer(ABC): -+ @staticmethod -+ def create( -+ enable: bool, -+ model_config: ModelConfig, -+ num_tokens: int, -+ max_running_requests: int, -+ device: str, -+ ): -+ if enable: -+ return _RoutedExpertsCapturerReal( -+ model_config, -+ num_tokens=num_tokens, -+ max_running_requests=max_running_requests, -+ device=device, -+ ) -+ else: -+ return _RoutedExpertsCapturerNoop() -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ raise NotImplementedError -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ raise NotImplementedError -+ -+ def sync_fwd_experts_buffer_DtoH(self, batch: int, loc: torch.Tensor): -+ raise NotImplementedError -+ -+ def get_host_cache(self): -+ raise NotImplementedError -+ -+ def get_device_cache(self): -+ raise NotImplementedError -+ -+ -+class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): -+ """Capturer for routed experts with host buffer""" -+ -+ def __init__( -+ self, -+ model_config: ModelConfig, -+ num_tokens: int, -+ max_running_requests: int, -+ device: str, -+ ): -+ -+ self.host_cache = _RoutedExpertsHostCache(model_config, num_tokens) -+ -+ self.device_cache = _RoutedExpertsDeviceCache( -+ model_config, max_running_requests, device -+ ) -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) -+ -+ def sync_fwd_experts_buffer_DtoH(self, loc: torch.Tensor): -+ batch = loc.shape[0] -+ self.host_cache.buffer[loc] = self.device_cache.buffer[:batch].cpu() -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ cache_pool_idx = ( -+ req_to_token_pool.req_to_token[req_pool_idx][:seqlen].cpu().clone() -+ ) -+ -+ return self.get_host_cache().buffer[cache_pool_idx].tolist() -+ -+ def get_host_cache(self): -+ return self.host_cache -+ -+ def get_device_cache(self): -+ return self.device_cache -+ -+ -+class _RoutedExpertsCapturerNoop(RoutedExpertsCapturer): -+ def __init__(self): -+ pass -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ pass -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ pass -+ -+ def sync_fwd_experts_buffer_DtoH(self, loc: torch.Tensor): -+ pass -+ -+ def get_host_cache(self): -+ pass -+ -+ def get_device_cache(self): -+ pass -+ -+ -+_global_expert_capturer: Optional[RoutedExpertsCapturer] = _RoutedExpertsCapturerNoop() -+ -+ -+def get_global_experts_capturer(): -+ return _global_expert_capturer -+ -+ -+def set_global_experts_capturer(capturer: RoutedExpertsCapturer): -+ global _global_expert_capturer -+ _global_expert_capturer = capturer -diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py -index 203cd5f41..dad8515f5 100644 ---- a/python/sglang/srt/layers/moe/topk.py -+++ b/python/sglang/srt/layers/moe/topk.py -@@ -44,6 +44,7 @@ from sglang.srt.eplb.expert_location_dispatch import ( - ) - from sglang.srt.layers.dp_attention import is_allocation_symmetric - from sglang.srt.layers.moe import get_moe_runner_backend -+from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer - from sglang.srt.utils import ( - cpu_has_amx_support, - get_bool_env_var, -@@ -195,6 +196,7 @@ class TopK(CustomOp): - self, - top_k: int, - *, -+ layer_id: Optional[int] = None, - use_grouped_topk: bool = False, - topk_group: Optional[int] = None, - num_expert_group: Optional[int] = None, -@@ -215,6 +217,7 @@ class TopK(CustomOp): - if use_grouped_topk: - assert num_expert_group is not None and topk_group is not None - -+ self.layer_id = layer_id - self.topk_config = TopKConfig( - top_k=top_k, - use_grouped_topk=use_grouped_topk, -@@ -240,6 +243,7 @@ class TopK(CustomOp): - self.topk_config.torch_native = True - return select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -289,6 +293,7 @@ class TopK(CustomOp): - ): - topk_output = select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -306,6 +311,7 @@ class TopK(CustomOp): - ) -> TopKOutput: - return select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -387,6 +393,7 @@ class TopK(CustomOp): - self.topk_config.torch_native = True - return select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -823,6 +830,7 @@ def select_experts( - router_logits: torch.Tensor, - topk_config: TopKConfig, - *, -+ layer_id: Optional[int] = None, - num_token_non_padded: Optional[torch.Tensor] = None, - expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None, - ) -> StandardTopKOutput: -@@ -920,7 +928,10 @@ def select_experts( - ) - - get_global_expert_distribution_recorder().on_select_experts(topk_ids=topk_ids) -- -+ get_global_experts_capturer().capture( -+ layer_id=layer_id, -+ topk_ids=topk_ids, -+ ) - return StandardTopKOutput(topk_weights, topk_ids, router_logits) - - -diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py -index 51981da81..7b54569c4 100644 ---- a/python/sglang/srt/layers/rotary_embedding.py -+++ b/python/sglang/srt/layers/rotary_embedding.py -@@ -129,9 +129,6 @@ class RotaryEmbedding(CustomOp): - - if get_global_server_args().rl_on_policy_target is not None: - self._forward_method = self.forward_native -- self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( -- self._apply_rotary_emb_wrapped -- ) - - def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: - """Compute the inverse frequency.""" -diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py -index 59a0f3bb9..ca641831a 100644 ---- a/python/sglang/srt/layers/sampler.py -+++ b/python/sglang/srt/layers/sampler.py -@@ -102,16 +102,11 @@ class Sampler(nn.Module): - if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB: - probs_without_temp_scaling = torch.softmax(logits, dim=-1) - -- if get_global_server_args().rl_on_policy_target is not None: -- logits_div_temperature = ( -- logits.bfloat16().div(sampling_info.temperatures).bfloat16() -- ) -- logprobs_via_logsoftmax_kernel = torch.log_softmax( -- logits_div_temperature, dim=-1 -- ) -- - # Post process logits - logits.div_(sampling_info.temperatures) -+ if get_global_server_args().rl_on_policy_target is not None: -+ logprobs_via_logsoftmax_kernel = torch.log_softmax(logits, dim=-1) -+ - logits[:] = torch.softmax(logits, dim=-1) - probs = logits - del logits -diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py -index 9399bbbea..91fbf80ab 100644 ---- a/python/sglang/srt/managers/detokenizer_manager.py -+++ b/python/sglang/srt/managers/detokenizer_manager.py -@@ -273,6 +273,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): - output_token_ids_logprobs_idx=recv_obj.output_token_ids_logprobs_idx, - output_token_entropy_val=recv_obj.output_token_entropy_val, - output_hidden_states=recv_obj.output_hidden_states, -+ output_routed_experts=recv_obj.output_routed_experts, - placeholder_tokens_idx=None, - placeholder_tokens_val=None, - retraction_counts=recv_obj.retraction_counts, -diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index b22f98fbd..3513a9f14 100644 ---- a/python/sglang/srt/managers/io_struct.py -+++ b/python/sglang/srt/managers/io_struct.py -@@ -175,6 +175,8 @@ class GenerateReqInput(BaseReq): - log_metrics: bool = True - # Whether to return hidden states - return_hidden_states: Union[List[bool], bool] = False -+ # Whether to return captured routed experts -+ return_routed_experts: bool = False - - # The modalities of the image data [image, multi-images, video] - modalities: Optional[List[str]] = None -@@ -592,6 +594,7 @@ class GenerateReqInput(BaseReq): - if isinstance(self.return_hidden_states, list) - else self.return_hidden_states - ), -+ return_routed_experts=self.return_routed_experts, - modalities=self.modalities[i] if self.modalities else None, - session_params=self.session_params, - lora_path=self.lora_path[i] if self.lora_path is not None else None, -@@ -655,6 +658,9 @@ class TokenizedGenerateReqInput(BaseReq): - # Whether to return hidden states - return_hidden_states: bool = False - -+ # Whether to return captured routed experts -+ return_routed_experts: bool = False -+ - # The input embeds - input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None - -@@ -910,6 +916,9 @@ class BatchTokenIDOutput( - # Hidden states - output_hidden_states: List[List[float]] - -+ # The routed experts for each output token -+ output_routed_experts: List[List[int]] -+ - # The information of placeholder tokens (e.g., image token) - # idx is the index of the token in the prompt after expansion. - # val is the length of padded tokens after expansion. -@@ -989,6 +998,9 @@ class BatchStrOutput( - # Hidden states - output_hidden_states: List[List[float]] - -+ # The routed experts for each output token -+ output_routed_experts: List[List[int]] -+ - # The information of placeholder tokens (e.g., image token) - # idx is the index of the token in the prompt after expansion. - # val is the length of padded tokens after expansion. -diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index 326b010b2..b7e24cfcf 100644 ---- a/python/sglang/srt/managers/schedule_batch.py -+++ b/python/sglang/srt/managers/schedule_batch.py -@@ -451,6 +451,7 @@ class Req: - session_id: Optional[str] = None, - custom_logit_processor: Optional[str] = None, - return_hidden_states: bool = False, -+ return_routed_experts: bool = False, - eos_token_ids: Optional[Set[int]] = None, - bootstrap_host: Optional[str] = None, - bootstrap_port: Optional[int] = None, -@@ -628,6 +629,10 @@ class Req: - self.output_topk_p = None - self.output_topk_index = None - -+ # capture routed experts -+ self.return_routed_experts = return_routed_experts -+ self.routed_experts = [] -+ - # Embedding (return values) - self.embedding = None - -@@ -943,6 +948,7 @@ class Req: - self.retraction_count += 1 - - self.prefix_indices = torch.empty((0,), dtype=torch.int64) -+ self.routed_experts = [] - self.last_node = None - self.swa_uuid_for_lock = None - self.extend_input_len = 0 -@@ -1112,6 +1118,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - # Whether to return hidden states - return_hidden_states: bool = False - -+ # Whether to return captured experts -+ return_routed_experts: bool = False -+ - # Whether this batch is prefill-only (no token generation needed) - is_prefill_only: bool = False - -@@ -1155,6 +1164,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - device=req_to_token_pool.device, - spec_algorithm=spec_algorithm, - return_hidden_states=any(req.return_hidden_states for req in reqs), -+ return_routed_experts=any(req.return_routed_experts for req in reqs), - is_prefill_only=all(req.is_prefill_only for req in reqs), - chunked_req=chunked_req, - ) -@@ -1900,7 +1910,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - def __str__(self): - return ( - f"ScheduleBatch(forward_mode={self.forward_mode.name if self.forward_mode else 'None'}, " -- f"#req={(len(self.reqs))})" -+ f"#req={(len(self.reqs))}), " -+ f"#out_cache_loc={self.out_cache_loc})" - ) - - -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index 719e93691..dd9f613da 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -1250,6 +1250,7 @@ class Scheduler( - input_embeds=recv_req.input_embeds, - custom_logit_processor=recv_req.custom_logit_processor, - return_hidden_states=recv_req.return_hidden_states, -+ return_routed_experts=recv_req.return_routed_experts, - eos_token_ids=self.model_config.hf_eos_token_id, - bootstrap_host=recv_req.bootstrap_host, - bootstrap_port=recv_req.bootstrap_port, -diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -index 5f5467c4f..186216c73 100644 ---- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py -+++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -@@ -9,6 +9,7 @@ import torch - from sglang.srt.disaggregation.utils import DisaggregationMode - from sglang.srt.environ import envs - from sglang.srt.layers.logits_processor import LogitsProcessorOutput -+from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer - from sglang.srt.managers.io_struct import ( - AbortReq, - BatchEmbeddingOutput, -@@ -112,6 +113,14 @@ class SchedulerOutputProcessorMixin: - req.check_finished() - - if req.finished(): -+ req.routed_experts = ( -+ get_global_experts_capturer().get_routed_experts( -+ req_pool_idx=req.req_pool_idx, -+ seqlen=req.seqlen, -+ req_to_token_pool=self.req_to_token_pool, -+ ) -+ ) -+ - release_kv_cache(req, self.tree_cache) - req.time_stats.completion_time = time.perf_counter() - elif not batch.decoding_reqs or req not in batch.decoding_reqs: -@@ -333,6 +342,12 @@ class SchedulerOutputProcessorMixin: - req.check_finished(new_accepted_len) - - if req.finished(): -+ req.routed_experts = get_global_experts_capturer().get_routed_experts( -+ req_pool_idx=req.req_pool_idx, -+ seqlen=req.seqlen, -+ req_to_token_pool=self.req_to_token_pool, -+ ) -+ - if self.server_args.disaggregation_decode_enable_offload_kvcache: - # Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes - if not self.decode_offload_manager.offload_kv_cache(req): -@@ -721,6 +736,7 @@ class SchedulerOutputProcessorMixin: - spec_accepted_tokens = [] - retraction_counts = [] - output_hidden_states = None -+ output_routed_experts = None - - queue_times = [] - forward_entry_times = [] -@@ -925,6 +941,10 @@ class SchedulerOutputProcessorMixin: - if output_hidden_states is None: - output_hidden_states = [] - output_hidden_states.append(req.hidden_states) -+ if req.return_routed_experts: -+ if output_routed_experts is None: -+ output_routed_experts = [] -+ output_routed_experts.append(req.routed_experts) - - if ( - req.finished() -@@ -971,6 +991,7 @@ class SchedulerOutputProcessorMixin: - output_token_ids_logprobs_idx=output_token_ids_logprobs_idx, - output_token_entropy_val=None, - output_hidden_states=output_hidden_states, -+ output_routed_experts=output_routed_experts, - rids=rids, - http_worker_ipcs=http_worker_ipcs, - placeholder_tokens_idx=None, -diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py -index e6f59e5b0..c199b987b 100644 ---- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py -+++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py -@@ -138,7 +138,7 @@ class SchedulerRuntimeCheckerMixin: - f"available_size={len(self.req_to_token_pool.free_slots)}, " - f"total_size={self.req_to_token_pool.size}\n" - ) -- raise ValueError(msg) -+ # raise ValueError(msg) - - def check_memory(self: Scheduler): - if self.is_hybrid: -@@ -150,7 +150,7 @@ class SchedulerRuntimeCheckerMixin: - - if memory_leak: - msg = "token_to_kv_pool_allocator memory leak detected! " f"{token_msg}" -- raise ValueError(msg) -+ # raise ValueError(msg) - - self._check_req_pool() - -diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -index 9bed7030d..a12deed3a 100644 ---- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py -+++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -@@ -1,6 +1,7 @@ - from __future__ import annotations - - import logging -+import os - from typing import TYPE_CHECKING, Tuple - - import torch -@@ -11,6 +12,8 @@ from sglang.srt.constants import ( - GPU_MEMORY_TYPE_KV_CACHE, - GPU_MEMORY_TYPE_WEIGHTS, - ) -+from sglang.srt.distributed import get_moe_ep_group, get_moe_tp_group, get_tp_group -+from sglang.srt.layers.dp_attention import get_attention_tp_group - from sglang.srt.managers.io_struct import ( - DestroyWeightsUpdateGroupReqInput, - DestroyWeightsUpdateGroupReqOutput, -@@ -76,7 +79,8 @@ class SchedulerUpdateWeightsMixin: - - def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): - """Update the online model parameter from tensors.""" -- success, message = self.tp_worker.update_weights_from_tensor(recv_req) -+ worker = self.draft_worker or self.tp_worker -+ success, message = worker.update_weights_from_tensor(recv_req) - # TODO extract common code b/t update_weights_from_distributed and update_weights_from_tensor later - if success: - if recv_req.flush_cache: -@@ -132,6 +136,20 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_CUDA_GRAPH in tags: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_CUDA_GRAPH) - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ tp_group = get_tp_group() -+ if tp_group is not None and tp_group.pynccl_comm is not None: -+ tp_group.pynccl_comm.nccl_pause() -+ attn_tp_group = get_attention_tp_group() -+ if attn_tp_group is not None and attn_tp_group.pynccl_comm is not None: -+ attn_tp_group.pynccl_comm.nccl_pause() -+ moe_ep_group = get_moe_ep_group() -+ if moe_ep_group is not None and moe_ep_group.pynccl_comm is not None: -+ moe_ep_group.pynccl_comm.nccl_pause() -+ moe_tp_group = get_moe_tp_group() -+ if moe_tp_group is not None and moe_tp_group.pynccl_comm is not None: -+ moe_tp_group.pynccl_comm.nccl_pause() -+ - torch.cuda.synchronize() - - return ReleaseMemoryOccupationReqOutput() -@@ -150,6 +168,20 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_CUDA_GRAPH in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_CUDA_GRAPH) - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ tp_group = get_tp_group() -+ if tp_group is not None and tp_group.pynccl_comm is not None: -+ tp_group.pynccl_comm.nccl_resume() -+ attn_tp_group = get_attention_tp_group() -+ if attn_tp_group is not None and attn_tp_group.pynccl_comm is not None: -+ attn_tp_group.pynccl_comm.nccl_resume() -+ moe_ep_group = get_moe_ep_group() -+ if moe_ep_group is not None and moe_ep_group.pynccl_comm is not None: -+ moe_ep_group.pynccl_comm.nccl_resume() -+ moe_tp_group = get_moe_tp_group() -+ if moe_tp_group is not None and moe_tp_group.pynccl_comm is not None: -+ moe_tp_group.pynccl_comm.nccl_resume() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_WEIGHTS) - torch.distributed.barrier(self.tp_cpu_group) -diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index c5a2e35ed..403cfcd9b 100644 ---- a/python/sglang/srt/managers/tokenizer_manager.py -+++ b/python/sglang/srt/managers/tokenizer_manager.py -@@ -802,6 +802,7 @@ class TokenizerManager(TokenizerCommunicatorMixin): - session_params=session_params, - custom_logit_processor=obj.custom_logit_processor, - return_hidden_states=obj.return_hidden_states, -+ return_routed_experts=obj.return_routed_experts, - data_parallel_rank=obj.data_parallel_rank, - priority=obj.priority, - extra_key=obj.extra_key, -@@ -1169,6 +1170,9 @@ class TokenizerManager(TokenizerCommunicatorMixin): - async with self.is_pause_cond: - self.is_pause = True - self.abort_request(abort_all=True) -+ # do double abort to ensure all in-flight requests are aborted -+ await asyncio.sleep(1) -+ self.abort_request(abort_all=True) - - async def continue_generation(self): - async with self.is_pause_cond: -@@ -1490,6 +1494,9 @@ class TokenizerManager(TokenizerCommunicatorMixin): - if getattr(recv_obj, "output_hidden_states", None): - meta_info["hidden_states"] = recv_obj.output_hidden_states[i] - -+ if getattr(recv_obj, "output_routed_experts", None): -+ meta_info["routed_experts"] = recv_obj.output_routed_experts[i] -+ - if isinstance(recv_obj, BatchStrOutput): - state.text += recv_obj.output_strs[i] - if state.obj.stream: -@@ -1616,12 +1623,13 @@ class TokenizerManager(TokenizerCommunicatorMixin): - return - - if len(recv_obj.input_token_logprobs_val) > 0: -- state.input_token_logprobs_val.extend( -- recv_obj.input_token_logprobs_val[recv_obj_index] -- ) -- state.input_token_logprobs_idx.extend( -- recv_obj.input_token_logprobs_idx[recv_obj_index] -- ) -+ if recv_obj.input_token_logprobs_val[recv_obj_index]: -+ state.input_token_logprobs_val.extend( -+ recv_obj.input_token_logprobs_val[recv_obj_index] -+ ) -+ state.input_token_logprobs_idx.extend( -+ recv_obj.input_token_logprobs_idx[recv_obj_index] -+ ) - state.output_token_logprobs_val.extend( - recv_obj.output_token_logprobs_val[recv_obj_index] - ) -@@ -1739,6 +1747,9 @@ class TokenizerManager(TokenizerCommunicatorMixin): - meta_info["spec_accept_length"] = ( - recv_obj.completion_tokens[i] / recv_obj.spec_verify_ct[i] - ) -+ meta_info["spec_accept_token_num"] = accepted_tokens -+ meta_info["spec_draft_token_num"] = total_draft_tokens -+ meta_info["spec_verify_ct"] = recv_obj.spec_verify_ct[i] - - def _calculate_timing_metrics( - self, -diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py -index 739e28943..73d88e9e9 100644 ---- a/python/sglang/srt/mem_cache/memory_pool.py -+++ b/python/sglang/srt/mem_cache/memory_pool.py -@@ -357,12 +357,16 @@ class HybridReqToTokenPool(ReqToTokenPool): - device=device, - enable_memory_saver=enable_memory_saver, - ) -- self._init_mamba_pool( -- size=mamba_size, -- cache_params=cache_params, -- device=device, -- speculative_num_draft_tokens=speculative_num_draft_tokens, -+ memory_saver_adapter = TorchMemorySaverAdapter.create( -+ enable=enable_memory_saver - ) -+ with memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): -+ self._init_mamba_pool( -+ size=mamba_size, -+ cache_params=cache_params, -+ device=device, -+ speculative_num_draft_tokens=speculative_num_draft_tokens, -+ ) - - def _init_mamba_pool( - self, -@@ -848,6 +852,7 @@ class HybridLinearKVPool(KVCache): - enable_kvcache_transpose: bool, - device: str, - mamba_pool: MambaPool, -+ enable_memory_saver: bool, - # TODO: refactor mla related args - use_mla: bool = False, - kv_lora_rank: int = None, -@@ -879,7 +884,7 @@ class HybridLinearKVPool(KVCache): - head_dim=head_dim, - layer_num=self.full_layer_nums, - device=device, -- enable_memory_saver=False, -+ enable_memory_saver=enable_memory_saver, - ) - else: - TokenToKVPoolClass = MLATokenToKVPool -@@ -891,7 +896,7 @@ class HybridLinearKVPool(KVCache): - device=device, - kv_lora_rank=kv_lora_rank, - qk_rope_head_dim=qk_rope_head_dim, -- enable_memory_saver=False, -+ enable_memory_saver=enable_memory_saver, - ) - self.full_attention_layer_id_mapping = { - id: i for i, id in enumerate(full_attention_layer_ids) -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 19e029d60..12ae6d0eb 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -86,6 +86,11 @@ from sglang.srt.layers.dp_attention import ( - initialize_dp_attention, - ) - from sglang.srt.layers.logits_processor import LogitsProcessorOutput -+from sglang.srt.layers.moe.routed_experts_capturer import ( -+ RoutedExpertsCapturer, -+ get_global_experts_capturer, -+ set_global_experts_capturer, -+) - from sglang.srt.layers.sampler import Sampler - from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model - from sglang.srt.lora.lora_manager import LoRAManager -@@ -484,6 +489,10 @@ class ModelRunner: - server_args.max_running_requests, - server_args.max_total_tokens, - ) -+ -+ # Init routed experts capturer -+ self.init_routed_experts_capturer() -+ - if self.device == "cuda": - self.init_cublas() - self.init_attention_backend() -@@ -519,6 +528,31 @@ class ModelRunner: - - self.model.set_eagle3_layers_to_capture(eagle_aux_hidden_state_layer_ids) - -+ def init_routed_experts_capturer(self): -+ # TODO: the redundant logic with TpModelWorker -+ max_running_requests = min( -+ ( -+ self.max_total_num_tokens // 2 -+ if self.server_args.max_running_requests is None -+ else self.server_args.max_running_requests -+ // ( -+ self.server_args.dp_size -+ if self.server_args.enable_dp_attention -+ else 1 -+ ) -+ ), -+ self.req_to_token_pool.size, -+ ) -+ set_global_experts_capturer( -+ RoutedExpertsCapturer.create( -+ enable=get_global_server_args().enable_return_routed_experts, -+ model_config=self.model_config, -+ num_tokens=self.max_total_num_tokens + self.page_size, -+ max_running_requests=max_running_requests, -+ device=self.device, -+ ) -+ ) -+ - def model_specific_adjustment(self): - server_args = self.server_args - -@@ -758,7 +792,11 @@ class ModelRunner: - - with self.memory_saver_adapter.region( - GPU_MEMORY_TYPE_WEIGHTS, -- enable_cpu_backup=self.server_args.enable_weights_cpu_backup, -+ enable_cpu_backup=( -+ self.server_args.enable_weights_cpu_backup -+ if not self.is_draft_worker -+ else True -+ ), - ): - self.model = get_model( - model_config=self.model_config, -@@ -1810,6 +1848,7 @@ class ModelRunner: - enable_kvcache_transpose=False, - device=self.device, - mamba_pool=self.req_to_token_pool.mamba_pool, -+ enable_memory_saver=self.server_args.enable_memory_saver, - use_mla=self.use_mla_backend, - **extra_args, - ) -@@ -2164,6 +2203,10 @@ class ModelRunner: - reinit_attn_backend, - split_forward_count, - ) -+ # Copy cached routing experts' buffers back to CPU cache -+ get_global_experts_capturer().sync_fwd_experts_buffer_DtoH( -+ forward_batch.out_cache_loc -+ ) - - if self.eplb_manager is not None: - self.eplb_manager.on_forward_pass_end() -diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py -index 895b69105..4209a3ea6 100644 ---- a/python/sglang/srt/models/deepseek_v2.py -+++ b/python/sglang/srt/models/deepseek_v2.py -@@ -641,6 +641,7 @@ class DeepseekV2MoE(nn.Module): - - self.topk = TopK( - top_k=config.num_experts_per_tok + self.num_fused_shared_experts, -+ layer_id=self.layer_id, - renormalize=config.norm_topk_prob, - use_grouped_topk=True, - num_expert_group=config.n_group, -diff --git a/python/sglang/srt/models/ernie4.py b/python/sglang/srt/models/ernie4.py -index ab1b6576b..dffd8f09a 100644 ---- a/python/sglang/srt/models/ernie4.py -+++ b/python/sglang/srt/models/ernie4.py -@@ -87,6 +87,7 @@ class Ernie4Moe(nn.Module): - - self.topk = TopK( - top_k=config.moe_k, -+ layer_id=layer_id, - renormalize=True, - use_grouped_topk=False, - correction_bias=self.gate.e_score_correction_bias, -diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py -index 3b04422b1..3b810843e 100644 ---- a/python/sglang/srt/models/glm4_moe.py -+++ b/python/sglang/srt/models/glm4_moe.py -@@ -374,6 +374,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): - - self.topk = TopK( - top_k=self.top_k, -+ layer_id=self.layer_id, - renormalize=config.norm_topk_prob, - use_grouped_topk=True, - num_expert_group=config.n_group, -diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py -index 9474700c4..398d622ff 100644 ---- a/python/sglang/srt/models/gpt_oss.py -+++ b/python/sglang/srt/models/gpt_oss.py -@@ -113,6 +113,7 @@ class GptOssSparseMoeBlock(nn.Module): - self.topk = TopK( - top_k=config.num_experts_per_tok, - renormalize=True, -+ layer_id=layer_id, - ) - - self.top_k = config.num_experts_per_tok -diff --git a/python/sglang/srt/models/grok.py b/python/sglang/srt/models/grok.py -index 1f4a3b443..4eb23cca8 100644 ---- a/python/sglang/srt/models/grok.py -+++ b/python/sglang/srt/models/grok.py -@@ -167,6 +167,7 @@ class Grok1MoE(nn.Module): - self.topk = TopK( - top_k=top_k, - renormalize=False, -+ layer_id=layer_id, - custom_routing_function=custom_routing_function, - ) - -diff --git a/python/sglang/srt/models/hunyuan.py b/python/sglang/srt/models/hunyuan.py -index 7c6fd9e48..b20d28544 100644 ---- a/python/sglang/srt/models/hunyuan.py -+++ b/python/sglang/srt/models/hunyuan.py -@@ -150,6 +150,7 @@ class HunYuanSparseMoeBlock(nn.Module): - - self.topk = TopK( - top_k=top_k, -+ layer_id=layer_id, - renormalize=True if top_k > 1 else False, - ) - -diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py -index 84aeb8b30..19637b20c 100644 ---- a/python/sglang/srt/models/longcat_flash.py -+++ b/python/sglang/srt/models/longcat_flash.py -@@ -241,6 +241,7 @@ class LongcatFlashMoE(nn.Module): - renormalize=False, - use_grouped_topk=False, - correction_bias=self.router.e_score_correction_bias.data, -+ layer_id=layer_id, - ) - self.topk.forward = self.topk.forward_native - -diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py -index a7dbadec6..c83a41338 100644 ---- a/python/sglang/srt/models/qwen2.py -+++ b/python/sglang/srt/models/qwen2.py -@@ -90,9 +90,6 @@ class Qwen2MLP(nn.Module): - self.act_fn = SiluAndMul() - - def forward(self, x): -- if get_global_server_args().rl_on_policy_target is not None: -- x = x.bfloat16() -- - gate_up, _ = self.gate_up_proj(x) - x = self.act_fn(gate_up) - x, _ = self.down_proj(x) -@@ -279,11 +276,6 @@ class Qwen2Model(nn.Module): - quant_config=quant_config, - enable_tp=not is_dp_attention_enabled(), - prefix=add_prefix("embed_tokens", prefix), -- params_dtype=( -- torch.float32 -- if get_global_server_args().rl_on_policy_target is not None -- else None -- ), - ) - else: - self.embed_tokens = PPMissingLayer() -@@ -306,10 +298,8 @@ class Qwen2Model(nn.Module): - if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py -index 051095e61..7db06dea8 100644 ---- a/python/sglang/srt/models/qwen2_moe.py -+++ b/python/sglang/srt/models/qwen2_moe.py -@@ -151,6 +151,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module): - self.topk = TopK( - top_k=config.num_experts_per_tok, - renormalize=config.norm_topk_prob, -+ layer_id=layer_id, - ) - - self.experts = get_moe_impl_class(quant_config)( -@@ -552,7 +553,17 @@ class Qwen2MoeModel(nn.Module): - prefix=add_prefix("layers", prefix), - ) - if self.pp_group.is_last_rank: -- self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.norm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - else: - self.norm = PPMissingLayer(return_tuple=True) - -diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py -index 9a9ac4da8..da1e1f713 100644 ---- a/python/sglang/srt/models/qwen3.py -+++ b/python/sglang/srt/models/qwen3.py -@@ -91,8 +91,8 @@ class Qwen3Attention(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -@@ -167,18 +167,10 @@ class Qwen3Attention(nn.Module): - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - ) -> torch.Tensor: -- if get_global_server_args().rl_on_policy_target is not None: -- hidden_states = hidden_states.bfloat16() -- - qkv, _ = self.qkv_proj(hidden_states) - q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q, k = self._apply_qk_norm(q, k) - q, k = self.rotary_emb(positions, q, k) -- -- if get_global_server_args().rl_on_policy_target is not None: -- q = q.to(torch.bfloat16) -- k = k.to(torch.bfloat16) -- - attn_output = self.attn(q, k, v, forward_batch) - output, _ = self.o_proj(attn_output) - return output -@@ -224,10 +216,8 @@ class Qwen3DecoderLayer(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index d3acc629b..3c59c51f2 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -21,6 +21,7 @@ import logging - from typing import Any, Dict, Iterable, List, Optional, Tuple - - import torch -+import torch.nn.functional as F - from torch import nn - - from sglang.srt.distributed import ( -@@ -48,7 +49,7 @@ from sglang.srt.layers.moe import ( - ) - from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class - from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE --from sglang.srt.layers.moe.topk import TopK -+from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK - from sglang.srt.layers.quantization.base_config import QuantizationConfig - from sglang.srt.layers.radix_attention import RadixAttention - from sglang.srt.layers.rotary_embedding import MRotaryEmbedding, get_rope -@@ -100,7 +101,9 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - top_k=config.num_experts_per_tok, - renormalize=config.norm_topk_prob, - use_grouped_topk=False, -+ layer_id=layer_id, - ) -+ self.top_k = config.num_experts_per_tok - - self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts -@@ -162,7 +165,22 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - - # router_logits: (num_tokens, n_experts) - router_logits, _ = self.gate(hidden_states) -- topk_output = self.topk(hidden_states, router_logits) -+ -+ if get_global_server_args().rl_on_policy_target is not None: -+ routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) -+ routing_weights, selected_experts = torch.topk( -+ routing_weights, self.top_k, dim=-1 -+ ) -+ routing_weights /= routing_weights.sum(dim=-1, keepdim=True) -+ routing_weights = routing_weights.to(hidden_states.dtype) -+ topk_output = StandardTopKOutput( -+ topk_weights=routing_weights, -+ topk_ids=selected_experts, -+ router_logits=router_logits, -+ ) -+ else: -+ topk_output = self.topk(hidden_states, router_logits) -+ - final_hidden_states = self.experts(hidden_states, topk_output) - if ( - self.tp_size > 1 -@@ -341,7 +359,7 @@ class Qwen3MoeAttention(nn.Module): - ) - self.compatible_with_fused_kv_buffer = ( - False if isinstance(self.rotary_emb, MRotaryEmbedding) else True -- ) -+ ) and (get_global_server_args().rl_on_policy_target is None) - - self.attn = RadixAttention( - self.num_heads, -@@ -352,8 +370,16 @@ class Qwen3MoeAttention(nn.Module): - prefix=add_prefix("attn", prefix), - ) - -- self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -- self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) -+ self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) - self.alt_stream = alt_stream - - def _apply_qk_norm( -@@ -518,9 +544,19 @@ class Qwen3MoeDecoderLayer(nn.Module): - quant_config=quant_config, - prefix=add_prefix("mlp", prefix), - ) -- self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.input_layernorm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - self.post_attention_layernorm = RMSNorm( -- config.hidden_size, eps=config.rms_norm_eps -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs - ) - - self.layer_communicator = LayerCommunicator( -diff --git a/python/sglang/srt/models/step3_vl.py b/python/sglang/srt/models/step3_vl.py -index 5a9e74ab6..07a06351f 100644 ---- a/python/sglang/srt/models/step3_vl.py -+++ b/python/sglang/srt/models/step3_vl.py -@@ -129,6 +129,7 @@ class Step3TextMoEMLP(nn.Module): - top_k=config.moe_top_k, - renormalize=config.norm_expert_weight, - use_grouped_topk=False, -+ layer_id=layer_id, - ) - - self.experts = get_moe_impl_class(quant_config)( -diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index 5b9a520b9..da24facf4 100644 ---- a/python/sglang/srt/server_args.py -+++ b/python/sglang/srt/server_args.py -@@ -515,6 +515,7 @@ class ServerArgs: - disable_fast_image_processor: bool = False - keep_mm_feature_on_device: bool = False - enable_return_hidden_states: bool = False -+ enable_return_routed_experts: bool = False - scheduler_recv_interval: int = 1 - numa_node: Optional[List[int]] = None - enable_deterministic_inference: bool = False -@@ -3384,6 +3385,11 @@ class ServerArgs: - action="store_true", - help="Enable returning hidden states with responses.", - ) -+ parser.add_argument( -+ "--enable-return-routed-experts", -+ action="store_true", -+ help="Enable returning routed experts of each layer with responses.", -+ ) - parser.add_argument( - "--scheduler-recv-interval", - type=int, -diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index a2d72dc48..c18f37f1c 100644 ---- a/python/sglang/srt/speculative/eagle_info.py -+++ b/python/sglang/srt/speculative/eagle_info.py -@@ -750,6 +750,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] - self.verified_id = self.verified_id[: len(new_indices)] -+ if self.accept_length is not None: -+ self.accept_length = self.accept_length[: len(new_indices)] -+ if self.accept_length_cpu is not None: -+ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] - else: - # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` - self.topk_p = self.topk_p[new_indices] -@@ -784,6 +788,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) - self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) - self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) -+ if self.accept_length is not None and spec_info.accept_length is not None: -+ self.accept_length = torch.cat( -+ [self.accept_length, spec_info.accept_length] -+ ) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif self.accept_length is not None: -+ zeros = torch.zeros( -+ [spec_info.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([self.accept_length, zeros]) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif spec_info.accept_length is not None: -+ zeros = torch.zeros( -+ [self.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([zeros, spec_info.accept_length]) -+ self.accept_length_cpu = self.accept_length.tolist() - - - @dataclass -diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py -index 08e6516bc..37b4c3f85 100644 ---- a/python/sglang/srt/speculative/eagle_worker.py -+++ b/python/sglang/srt/speculative/eagle_worker.py -@@ -9,6 +9,7 @@ from sglang.srt.layers.dp_attention import get_attention_tp_group - from sglang.srt.layers.logits_processor import LogitsProcessorOutput - from sglang.srt.layers.moe.utils import speculative_moe_backend_context - from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs -+from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput - from sglang.srt.managers.schedule_batch import ScheduleBatch - from sglang.srt.managers.scheduler import GenerationBatchResult - from sglang.srt.managers.tp_worker import TpModelWorker -@@ -50,6 +51,7 @@ from sglang.srt.speculative.spec_utils import ( - select_top_k_tokens, - ) - from sglang.srt.utils import ( -+ MultiprocessingSerializer, - empty_context, - get_available_gpu_memory, - get_bool_env_var, -@@ -57,6 +59,7 @@ from sglang.srt.utils import ( - is_npu, - next_power_of_2, - ) -+from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions - - _is_npu = is_npu() - -@@ -984,6 +987,26 @@ class EAGLEWorker(TpModelWorker): - draft_input.topk_p, draft_input.topk_index = fast_topk(probs, self.topk, dim=-1) - draft_input.hidden_states = logits_output.hidden_states - -+ def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): -+ -+ monkey_patch_torch_reductions() -+ named_tensors = MultiprocessingSerializer.deserialize( -+ recv_req.serialized_named_tensors[self.tp_rank] -+ ) -+ success, message = self.model_runner.update_weights_from_tensor( -+ named_tensors=named_tensors, -+ load_format=recv_req.load_format, -+ ) -+ if not success: -+ return success, message -+ -+ success, message = self.target_worker.model_runner.update_weights_from_tensor( -+ named_tensors=named_tensors, -+ load_format=recv_req.load_format, -+ ) -+ -+ return success, message -+ - - @torch.compile(dynamic=True, disable=_is_npu) - def get_last_loc_large_page_size_top_k_1( -diff --git a/python/sglang/srt/weight_sync/tensor_bucket.py b/python/sglang/srt/weight_sync/tensor_bucket.py -index 44273713f..c1d592ddb 100644 ---- a/python/sglang/srt/weight_sync/tensor_bucket.py -+++ b/python/sglang/srt/weight_sync/tensor_bucket.py -@@ -22,6 +22,9 @@ class FlattenedTensorBucket: - while preserving all metadata needed for reconstruction. - """ - -+ # This field is solely for users of to check whether the class supports this feature -+ supports_multi_dtypes = True -+ - def __init__( - self, - named_tensors: List[Tuple[str, torch.Tensor]] = None, -@@ -48,7 +51,7 @@ class FlattenedTensorBucket: - flattened_tensors: List[torch.Tensor] = [None] * len(named_tensors) - - for i, (name, tensor) in enumerate(named_tensors): -- flattened = tensor.flatten() -+ flattened = tensor.flatten().view(torch.uint8) - flattened_tensors[i] = flattened - - # Store metadata -@@ -93,14 +96,12 @@ class FlattenedTensorBucket: - reconstructed = [None] * len(self.metadata) - - for i, meta in enumerate(self.metadata): -- tensor = self.flattened_tensor[meta.start_idx : meta.end_idx].reshape( -- meta.shape -+ tensor = ( -+ self.flattened_tensor[meta.start_idx : meta.end_idx] -+ .view(meta.dtype) -+ .reshape(meta.shape) - ) - -- # batch dtype conversion (if needed) -- if tensor.dtype != meta.dtype: -- tensor = tensor.to(meta.dtype) -- - reconstructed[i] = (meta.name, tensor) - - return reconstructed diff --git a/docker/patch/v0.5.6/sglang.patch b/docker/patch/v0.5.6/sglang.patch deleted file mode 100644 index de12cdd43..000000000 --- a/docker/patch/v0.5.6/sglang.patch +++ /dev/null @@ -1,2053 +0,0 @@ -diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py -index ef52bda7f..537d892dc 100644 ---- a/python/sglang/srt/disaggregation/decode.py -+++ b/python/sglang/srt/disaggregation/decode.py -@@ -296,6 +296,13 @@ class DecodePreallocQueue: - ) - return kv_manager - -+ def release_memory_occupation(self): -+ if hasattr(self.kv_manager, "close"): -+ self.kv_manager.close() -+ -+ def resume_memory_occupation(self): -+ self.kv_manager = self._init_kv_manager() -+ - def add(self, req: Req, is_retracted: bool = False) -> None: - """Add a request to the pending queue.""" - if self._check_if_req_exceed_kv_capacity(req): -diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py -index d4414d084..c5fb10155 100644 ---- a/python/sglang/srt/disaggregation/mooncake/conn.py -+++ b/python/sglang/srt/disaggregation/mooncake/conn.py -@@ -1074,6 +1074,19 @@ class MooncakeKVManager(CommonKVManager): - f"Losing connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr}), {len(affected_rooms)} requests affected" - ) - -+ def close(self): -+ # Batch deregister KV data buffers -+ if self.kv_args.kv_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.kv_data_ptrs) -+ -+ # Batch deregister auxiliary data buffers -+ if self.kv_args.aux_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.aux_data_ptrs) -+ -+ # Batch deregister state/extra pool data buffers -+ if self.kv_args.state_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.state_data_ptrs) -+ - - class MooncakeKVSender(CommonKVSender): - -diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py -index 952374ed5..239ac2571 100644 ---- a/python/sglang/srt/disaggregation/prefill.py -+++ b/python/sglang/srt/disaggregation/prefill.py -@@ -305,6 +305,13 @@ class PrefillBootstrapQueue: - else: - return bootstrapped_reqs, failed_reqs - -+ def release_memory_occupation(self): -+ if hasattr(self.kv_manager, "close"): -+ self.kv_manager.close() -+ -+ def resume_memory_occupation(self): -+ self.kv_manager = self._init_kv_manager() -+ - - class SchedulerDisaggregationPrefillMixin: - """ -diff --git a/python/sglang/srt/distributed/device_communicators/pynccl.py b/python/sglang/srt/distributed/device_communicators/pynccl.py -index 86c53f26b..52acf95b9 100644 ---- a/python/sglang/srt/distributed/device_communicators/pynccl.py -+++ b/python/sglang/srt/distributed/device_communicators/pynccl.py -@@ -380,3 +380,9 @@ class PyNcclCommunicator: - - self.disabled = old_disable - self.stream = old_stream -+ -+ def nccl_pause(self): -+ self.nccl.ncclPause(self.comm) -+ -+ def nccl_resume(self): -+ self.nccl.ncclResume(self.comm) -diff --git a/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py b/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -index 6b12f2922..7028a4e46 100644 ---- a/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -+++ b/python/sglang/srt/distributed/device_communicators/pynccl_wrapper.py -@@ -304,6 +304,17 @@ class NCCLLibrary: - Function("ncclGroupEnd", ncclResult_t, []), - ] - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ exported_functions.extend( -+ [ -+ # ncclResult_t ncclPause(ncclComm_t comm); -+ Function("ncclPause", ncclResult_t, [ncclComm_t]), -+ # ncclResult_t ncclResume(ncclComm_t comm); -+ Function("ncclResume", ncclResult_t, [ncclComm_t]), -+ Function("ncclSetGroupID", ncclResult_t, [ctypes.c_int]), -+ ] -+ ) -+ - exported_functions_symm_mem = [ - # ncclResult_t ncclCommWindowRegister(ncclComm_t comm, void* buff, size_t size, ncclWindow_t* win, int winFlags); - Function( -@@ -551,6 +562,12 @@ class NCCLLibrary: - def ncclGroupEnd(self) -> None: - self.NCCL_CHECK(self._funcs["ncclGroupEnd"]()) - -+ def ncclPause(self, comm: ncclComm_t) -> None: -+ self.NCCL_CHECK(self._funcs["ncclPause"](comm)) -+ -+ def ncclResume(self, comm: ncclComm_t) -> None: -+ self.NCCL_CHECK(self._funcs["ncclResume"](comm)) -+ - - __all__ = [ - "NCCLLibrary", -diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py -index cf90f6fe0..11d26df81 100644 ---- a/python/sglang/srt/distributed/parallel_state.py -+++ b/python/sglang/srt/distributed/parallel_state.py -@@ -1780,7 +1780,10 @@ def get_tensor_model_parallel_world_size(): - - def get_tensor_model_parallel_rank(): - """Return my rank for the tensor model parallel group.""" -- return get_tp_group().rank_in_group -+ try: -+ return get_tp_group().rank_in_group -+ except Exception: -+ return 0 - - - def get_pipeline_model_parallel_world_size(): -diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py -index 67a082ea6..390365864 100644 ---- a/python/sglang/srt/entrypoints/engine.py -+++ b/python/sglang/srt/entrypoints/engine.py -@@ -183,6 +183,7 @@ class Engine(EngineBase): - lora_path: Optional[List[Optional[str]]] = None, - custom_logit_processor: Optional[Union[List[str], str]] = None, - return_hidden_states: bool = False, -+ return_routed_experts: bool = False, - stream: bool = False, - bootstrap_host: Optional[Union[List[str], str]] = None, - bootstrap_port: Optional[Union[List[int], int]] = None, -@@ -218,6 +219,7 @@ class Engine(EngineBase): - lora_path=lora_path, - custom_logit_processor=custom_logit_processor, - return_hidden_states=return_hidden_states, -+ return_routed_experts=return_routed_experts, - stream=stream, - bootstrap_host=bootstrap_host, - bootstrap_port=bootstrap_port, -diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py -index 9f556a885..992843285 100644 ---- a/python/sglang/srt/layers/attention/vision.py -+++ b/python/sglang/srt/layers/attention/vision.py -@@ -518,11 +518,25 @@ class VisionAttention(nn.Module): - self.dummy_dim = (num_dummy_heads + num_heads) * self.head_size - - if self.qk_normalization: -+ norm_kwargs = ( -+ dict( -+ weight_dtype=torch.float32, -+ cast_x_before_out_mul=True, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) - self.q_norm = RMSNorm( -- self.dummy_dim, eps=layer_norm_eps, var_hidden_size=embed_dim -+ self.dummy_dim, -+ eps=layer_norm_eps, -+ var_hidden_size=embed_dim, -+ **norm_kwargs, - ) - self.k_norm = RMSNorm( -- self.dummy_dim, eps=layer_norm_eps, var_hidden_size=embed_dim -+ self.dummy_dim, -+ eps=layer_norm_eps, -+ var_hidden_size=embed_dim, -+ **norm_kwargs, - ) - - # Select attention backend via a unified method -@@ -648,6 +662,15 @@ class VisionAttention(nn.Module): - if x.dim() == 2: - x = x.unsqueeze(0) - assert x.dim() == 3, x.shape -+ if ( -+ get_global_server_args().rl_on_policy_target is not None -+ and position_embeddings is not None -+ ): -+ assert isinstance(position_embeddings, tuple), ( -+ "expected position_embeddings to be a tuple of two tensors,\n" -+ f"but got {type(position_embeddings)}, change if needed" -+ ) -+ position_embeddings = tuple(p.to(x.dtype) for p in position_embeddings) - x_shape = x.shape - bsz, s, _ = x_shape - head = self.num_attention_heads_per_partition -diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py -index 932f52aeb..79c6b664f 100644 ---- a/python/sglang/srt/layers/communicator.py -+++ b/python/sglang/srt/layers/communicator.py -@@ -372,6 +372,7 @@ class LayerCommunicator: - residual: torch.Tensor, - forward_batch: ForwardBatch, - quant_format: str = "", -+ post_residual_addition: Optional[torch.Tensor] = None, - ): - if get_attn_tp_context().input_scattered: - hidden_states, residual = self._tp_reduce_scatter( -@@ -453,7 +454,9 @@ class LayerCommunicator: - ) - else: - hidden_states, residual = self.input_layernorm( -- hidden_states, residual -+ hidden_states, -+ residual, -+ post_residual_addition, - ) - - hidden_states = self._communicate_simple_fn( -diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py -index 3293a8a59..a075b71ce 100644 ---- a/python/sglang/srt/layers/layernorm.py -+++ b/python/sglang/srt/layers/layernorm.py -@@ -84,15 +84,12 @@ class RMSNorm(CustomOp): - eps: float = 1e-6, - var_hidden_size: Optional[int] = None, - cast_x_before_out_mul: bool = False, -- fp32_residual: bool = False, -- weight_dtype: Optional = None, -- override_orig_dtype: Optional = None, -+ fp32_residual: bool = True, - ) -> None: - super().__init__() - self.cast_x_before_out_mul = cast_x_before_out_mul - self.fp32_residual = fp32_residual -- self.override_orig_dtype = override_orig_dtype -- self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) -+ self.weight = nn.Parameter(torch.ones(hidden_size)) - self.variance_epsilon = eps - self.hidden_size = hidden_size - self.variance_size_override = ( -@@ -105,21 +102,26 @@ class RMSNorm(CustomOp): - self, - x: torch.Tensor, - residual: Optional[torch.Tensor] = None, -+ post_residual_addition: Optional[torch.Tensor] = None, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - if self.variance_size_override is not None: -- return self.forward_native(x, residual) -+ return self.forward_native(x, residual, post_residual_addition) - if is_batch_invariant_mode_enabled(): - if ( - residual is not None - or get_global_server_args().rl_on_policy_target == "fsdp" - ): -- return self.forward_native(x, residual) -+ return self.forward_native(x, residual, post_residual_addition) - return rms_norm_batch_invariant( - x, - self.weight.data, - self.variance_epsilon, - ) - if residual is not None: -+ # TODO: Ideally we want to have (a+b)+c. but right now we can only have a+(b+c). -+ # (a+b)+c != a+(b+c), we probably need to add another parameter to fused_add_rmsnorm -+ if post_residual_addition is not None: -+ residual = residual + post_residual_addition - fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon) - return x, residual - out = rmsnorm(x, self.weight.data, self.variance_epsilon) -@@ -179,17 +181,35 @@ class RMSNorm(CustomOp): - self, - x: torch.Tensor, - residual: Optional[torch.Tensor] = None, -+ post_residual_addition: Optional[torch.Tensor] = None, - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - if not x.is_contiguous(): - x = x.contiguous() -- orig_dtype = self.override_orig_dtype or x.dtype -+ orig_dtype = x.dtype -+ -+ if residual is not None and not self.fp32_residual: -+ x = ( -+ x -+ + residual -+ + ( -+ post_residual_addition -+ if post_residual_addition is not None -+ else 0.0 -+ ) -+ ) -+ residual = x.clone() - x = x.to(torch.float32) -- if residual is not None: -- x = x + residual.to(torch.float32) -- if self.fp32_residual: -- residual = x.clone() -- else: -- residual = x.to(orig_dtype) -+ if residual is not None and self.fp32_residual: -+ x = ( -+ x -+ + residual.to(torch.float32) -+ + ( -+ post_residual_addition.to(torch.float32) -+ if post_residual_addition is not None -+ else 0.0 -+ ) -+ ) -+ residual = x.to(orig_dtype) - - hidden_size = x.shape[-1] - if hidden_size != self.hidden_size: -diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py -index 522865765..733bad5f2 100644 ---- a/python/sglang/srt/layers/logits_processor.py -+++ b/python/sglang/srt/layers/logits_processor.py -@@ -841,11 +841,6 @@ class LogitsProcessor(nn.Module): - None, # bias - True, # is_vnni - ) -- elif get_global_server_args().rl_on_policy_target is not None: -- # Due to tie-weight, we may not be able to change lm_head's weight dtype -- logits = torch.matmul( -- hidden_states.bfloat16(), lm_head.weight.T.bfloat16() -- ) - else: - logits = torch.matmul( - hidden_states.to(lm_head.weight.dtype), lm_head.weight.T -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -index e7d5a67cc..639e47163 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -@@ -14,6 +14,7 @@ import torch.nn.functional as F - import triton.language as tl - - from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - cpu_has_amx_support, - direct_register_custom_op, -@@ -626,7 +627,10 @@ def fused_experts_impl( - ).squeeze(dim=1) - else: - # According to micro benchmark results, torch.compile can get better performance for small token. -- if tokens_in_chunk <= 32: -+ if ( -+ not get_global_server_args().enable_deterministic_inference -+ and tokens_in_chunk <= 32 -+ ): - moe_sum_reduce_torch_compile( - intermediate_cache3.view(*intermediate_cache3.shape), - out_hidden_states[begin_chunk_idx:end_chunk_idx], -diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py -new file mode 100644 -index 000000000..e16817f1f ---- /dev/null -+++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py -@@ -0,0 +1,279 @@ -+import logging -+from abc import ABC -+from contextlib import contextmanager -+from typing import Optional -+ -+import numpy as np -+import torch -+ -+from sglang.srt.configs.model_config import ModelConfig -+from sglang.srt.layers.dp_attention import ( -+ get_attention_dp_rank, -+ get_dp_local_info, -+ is_dp_attention_enabled, -+) -+from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -+from sglang.srt.server_args import get_global_server_args -+ -+logger = logging.getLogger(__name__) -+ -+_GB = 1024 * 1024 * 1024 -+_MB = 1024 * 1024 -+ -+ -+def get_tensor_size_bytes(t: torch.Tensor): -+ return np.prod(t.shape) * t.dtype.itemsize -+ -+ -+class _RoutedExpertsDeviceCache: -+ def __init__( -+ self, -+ max_running_requests: int, -+ num_hidden_layers: int, -+ num_experts_per_tok: int, -+ num_fused_shared_experts: int, -+ device: str, -+ ) -> None: -+ self.buffer = torch.zeros( -+ ( -+ max( -+ get_global_server_args().chunked_prefill_size -+ * get_global_server_args().dp_size, -+ max_running_requests, -+ ), -+ num_hidden_layers, -+ num_experts_per_tok + num_fused_shared_experts, -+ ), -+ dtype=torch.int32, -+ device=device, -+ ) -+ self._finalize_allocation_log() -+ -+ def get_buffer_size_bytes(self): -+ assert hasattr(self, "buffer") -+ return get_tensor_size_bytes(self.buffer) -+ -+ def capture_fwd_routed_experts(self, layer_id: int, topk_ids: torch.Tensor): -+ assert layer_id is not None, "capturing routing experts but get layer_id None" -+ batch, _ = topk_ids.shape -+ self.buffer[:batch, layer_id, :] = topk_ids -+ -+ def _finalize_allocation_log(self): -+ """Common logging and memory usage computation for captured experts buffers.""" -+ buffer_size_MB = self.get_buffer_size_bytes() / _MB -+ logger.info( -+ f"Routing experts device buffer allocated. #shape: {tuple(self.buffer.shape)}, size: {buffer_size_MB:.2f} MB" -+ ) -+ -+ -+class _RoutedExpertsHostCache: -+ def __init__( -+ self, -+ num_tokens: int, -+ num_hidden_layers: int, -+ num_experts_per_tok: int, -+ ) -> None: -+ self.num_tokens = num_tokens -+ self.buffer = torch.zeros( -+ ( -+ num_tokens, -+ num_hidden_layers, -+ num_experts_per_tok, -+ ), -+ dtype=torch.int32, -+ device="cpu", -+ pin_memory=True, -+ ) -+ self._finalize_allocation_log() -+ -+ def get_buffer_size_bytes(self): -+ assert hasattr(self, "buffer") -+ return get_tensor_size_bytes(self.buffer) -+ -+ def set_experts_buffer(self, layer_id: int, loc: torch.Tensor, top_k: torch.Tensor): -+ self.buffer[layer_id, loc, :] = top_k.to(device="cpu", non_blocking=True) -+ -+ def _finalize_allocation_log(self): -+ """Common logging and memory usage computation for captured experts buffers.""" -+ buffer_size_GB = self.get_buffer_size_bytes() / _GB -+ logger.info( -+ f"Routing experts host buffer allocated. #tokens: {self.num_tokens}, size: {buffer_size_GB:.2f} GB" -+ ) -+ -+ -+class RoutedExpertsCapturer(ABC): -+ @staticmethod -+ def create( -+ enable: bool, -+ model_config: ModelConfig, -+ num_fused_shared_experts: int, -+ num_tokens: int, -+ max_running_requests: int, -+ device: str, -+ ): -+ if enable: -+ return _RoutedExpertsCapturerReal( -+ model_config, -+ num_tokens=num_tokens, -+ max_running_requests=max_running_requests, -+ num_fused_shared_experts=num_fused_shared_experts, -+ device=device, -+ ) -+ else: -+ return _RoutedExpertsCapturerNoop() -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ raise NotImplementedError -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ raise NotImplementedError -+ -+ def sync_fwd_experts_buffer_DtoH( -+ self, -+ device_loc: torch.Tensor, -+ cpu_loc: torch.Tensor, -+ can_run_graph: bool, -+ cuda_graph_batch: int, -+ ): -+ raise NotImplementedError -+ -+ @contextmanager -+ def with_forward(self, forward_batch): -+ yield -+ -+ def get_host_cache(self): -+ raise NotImplementedError -+ -+ def get_device_cache(self): -+ raise NotImplementedError -+ -+ -+class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): -+ """Capturer for routed experts with host buffer""" -+ -+ def __init__( -+ self, -+ model_config: ModelConfig, -+ num_tokens: int, -+ max_running_requests: int, -+ num_fused_shared_experts: int, -+ device: str, -+ ): -+ self.forward_batch = None -+ self.num_fused_shared_experts = num_fused_shared_experts -+ self.num_hidden_layers = model_config.hf_text_config.num_hidden_layers -+ self.num_experts_per_tok = model_config.hf_text_config.num_experts_per_tok -+ -+ self.host_cache = _RoutedExpertsHostCache( -+ num_tokens=num_tokens, -+ num_hidden_layers=self.num_hidden_layers, -+ num_experts_per_tok=self.num_experts_per_tok, -+ ) -+ -+ self.device_cache = _RoutedExpertsDeviceCache( -+ max_running_requests=max_running_requests, -+ num_hidden_layers=self.num_hidden_layers, -+ num_experts_per_tok=self.num_experts_per_tok, -+ num_fused_shared_experts=self.num_fused_shared_experts, -+ device=device, -+ ) -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) -+ -+ def sync_fwd_experts_buffer_DtoH( -+ self, -+ device_loc: torch.Tensor, -+ cpu_loc: torch.Tensor, -+ can_run_graph: bool, -+ cuda_graph_batch: int, -+ ): -+ if is_dp_attention_enabled(): -+ local_start_pos, local_num_tokens = get_dp_local_info(self.forward_batch) -+ # handle with cuda graph padding -+ if can_run_graph: -+ local_start_pos = get_attention_dp_rank() * cuda_graph_batch -+ local_end_pos = local_start_pos + local_num_tokens -+ else: -+ local_end_pos = local_start_pos + local_num_tokens -+ else: -+ local_start_pos = 0 -+ local_end_pos = device_loc.shape[0] -+ -+ self.host_cache.buffer[cpu_loc] = self.device_cache.buffer[ -+ local_start_pos:local_end_pos, :, : self.num_experts_per_tok -+ ].cpu() -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ cache_pool_idx = ( -+ req_to_token_pool.req_to_token[req_pool_idx][: seqlen - 1].cpu().clone() -+ ) -+ return self.get_host_cache().buffer[cache_pool_idx] -+ -+ @contextmanager -+ def with_forward(self, forward_batch): -+ self.forward_batch = forward_batch -+ yield -+ -+ def get_host_cache(self): -+ return self.host_cache -+ -+ def get_device_cache(self): -+ return self.device_cache -+ -+ -+class _RoutedExpertsCapturerNoop(RoutedExpertsCapturer): -+ def __init__(self): -+ pass -+ -+ def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ pass -+ -+ def get_routed_experts( -+ self, -+ req_pool_idx: int, -+ seqlen: int, -+ req_to_token_pool: ReqToTokenPool, -+ ): -+ pass -+ -+ def sync_fwd_experts_buffer_DtoH( -+ self, -+ device_loc: torch.Tensor, -+ cpu_loc: torch.Tensor, -+ can_run_graph: bool, -+ cuda_graph_batch: int, -+ ): -+ pass -+ -+ @contextmanager -+ def with_forward(self, forward_batch): -+ yield -+ -+ def get_host_cache(self): -+ pass -+ -+ def get_device_cache(self): -+ pass -+ -+ -+_global_expert_capturer: Optional[RoutedExpertsCapturer] = _RoutedExpertsCapturerNoop() -+ -+ -+def get_global_experts_capturer(): -+ return _global_expert_capturer -+ -+ -+def set_global_experts_capturer(capturer: RoutedExpertsCapturer): -+ global _global_expert_capturer -+ _global_expert_capturer = capturer -\ No newline at end of file -diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py -index a802647e8..0fd550c0c 100644 ---- a/python/sglang/srt/layers/moe/topk.py -+++ b/python/sglang/srt/layers/moe/topk.py -@@ -48,6 +48,7 @@ from sglang.srt.eplb.expert_location_dispatch import ( - ) - from sglang.srt.layers.dp_attention import is_allocation_symmetric - from sglang.srt.layers.moe import get_moe_runner_backend -+from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer - from sglang.srt.utils import ( - cpu_has_amx_support, - get_bool_env_var, -@@ -212,6 +213,7 @@ class TopK(CustomOp): - self, - top_k: int, - *, -+ layer_id: Optional[int] = None, - use_grouped_topk: bool = False, - topk_group: Optional[int] = None, - num_expert_group: Optional[int] = None, -@@ -233,6 +235,7 @@ class TopK(CustomOp): - if use_grouped_topk: - assert num_expert_group is not None and topk_group is not None - -+ self.layer_id = layer_id - self.topk_config = TopKConfig( - top_k=top_k, - use_grouped_topk=use_grouped_topk, -@@ -260,6 +263,7 @@ class TopK(CustomOp): - self.topk_config.torch_native = True - return select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -309,6 +313,7 @@ class TopK(CustomOp): - ): - topk_output = select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -326,6 +331,7 @@ class TopK(CustomOp): - ) -> TopKOutput: - return select_experts( - hidden_states=hidden_states, -+ layer_id=self.layer_id, - router_logits=router_logits, - topk_config=self.topk_config, - num_token_non_padded=num_token_non_padded, -@@ -856,6 +862,7 @@ def select_experts( - router_logits: torch.Tensor, - topk_config: TopKConfig, - *, -+ layer_id: Optional[int] = None, - num_token_non_padded: Optional[torch.Tensor] = None, - expert_location_dispatch_info: Optional[ExpertLocationDispatchInfo] = None, - ) -> StandardTopKOutput: -@@ -983,7 +990,10 @@ def select_experts( - ) - - get_global_expert_distribution_recorder().on_select_experts(topk_ids=topk_ids) -- -+ get_global_experts_capturer().capture( -+ layer_id=layer_id, -+ topk_ids=topk_ids, -+ ) - return StandardTopKOutput(topk_weights, topk_ids, router_logits) - - -diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py -index 70466bb20..cd85fc2f2 100644 ---- a/python/sglang/srt/layers/moe/utils.py -+++ b/python/sglang/srt/layers/moe/utils.py -@@ -284,7 +284,7 @@ def speculative_moe_a2a_backend_context(): - global MOE_A2A_BACKEND - original_backend = MOE_A2A_BACKEND - try: -- MOE_A2A_BACKEND = MoeA2ABackend.NONE -+ MOE_A2A_BACKEND = get_speculative_moe_a2a_backend() - yield - finally: - MOE_A2A_BACKEND = original_backend -diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py -index 0cdb7e1ae..df8860409 100644 ---- a/python/sglang/srt/layers/rotary_embedding.py -+++ b/python/sglang/srt/layers/rotary_embedding.py -@@ -15,7 +15,6 @@ from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - cpu_has_amx_support, - get_bool_env_var, -- get_compiler_backend, - is_cpu, - is_cuda, - is_hip, -@@ -132,9 +131,7 @@ class RotaryEmbedding(CustomOp): - - if get_global_server_args().rl_on_policy_target is not None: - self._forward_method = self.forward_native -- self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( -- self._apply_rotary_emb_wrapped -- ) -+ - self.position_cos, self.position_sin = None, None - - def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: -@@ -1423,6 +1420,9 @@ class MRotaryEmbedding(RotaryEmbedding): - f"Corrected mrope_section: {self.mrope_section} (sum={sum(self.mrope_section)})" - ) - -+ if get_global_server_args().rl_on_policy_target is not None: -+ self._forward_method = self.forward_native -+ - def _match_cos_sin_cache_dtype(self, query: torch.Tensor) -> None: - # __setattr__ in nn.Module (called by `self.cos_sin_cache = ...`) - # is expensive, so avoid calling it if possible -@@ -1432,8 +1432,7 @@ class MRotaryEmbedding(RotaryEmbedding): - ): - self.cos_sin_cache = self.cos_sin_cache.to(query.device, dtype=query.dtype) - -- @torch.compile(dynamic=True, backend=get_compiler_backend()) -- def _forward_native( -+ def forward_native( - self, - positions: torch.Tensor, - query: torch.Tensor, -@@ -1490,7 +1489,7 @@ class MRotaryEmbedding(RotaryEmbedding): - key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape) - return query, key - -- def forward( -+ def forward_cuda( - self, - positions: torch.Tensor, - query: torch.Tensor, -@@ -1507,14 +1506,12 @@ class MRotaryEmbedding(RotaryEmbedding): - """ - assert positions.ndim == 1 or positions.ndim == 2 - -- if positions.ndim == 2 and self.mrope_section and _is_cuda: -- return self._forward_triton(positions, query, key) -- elif _is_npu: -- return self._forward_npu(positions, query, key) -- else: -- return self._forward_native(positions, query, key) -+ # Use Triton kernel for multimodal (2D positions) with mrope -+ if positions.ndim == 2 and self.mrope_section: -+ return self.forward_triton(positions, query, key) -+ return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) - -- def _forward_triton( -+ def forward_triton( - self, - positions: torch.Tensor, - query: torch.Tensor, -@@ -1563,15 +1560,19 @@ class MRotaryEmbedding(RotaryEmbedding): - key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape) - return query, key - -- def _forward_npu( -+ def forward_npu( - self, - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, -+ fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: -+ assert ( -+ fused_set_kv_buffer_arg is None -+ ), "fused_set_kv_buffer_arg is not supported for npu implementation" - # TODO: remove this when npu_mrope supports QNumHeads * QHeadSize > 4096 - if query.shape[1] > 4096: -- return self._forward_native(positions, query, key) -+ return self.forward_native(positions, query, key, fused_set_kv_buffer_arg) - rotary_mode = "half" - if self.is_neox_style: - rotary_mode = "half" -diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py -index 7f6f6a010..c4a673145 100644 ---- a/python/sglang/srt/layers/sampler.py -+++ b/python/sglang/srt/layers/sampler.py -@@ -105,16 +105,11 @@ class Sampler(nn.Module): - if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB: - probs_without_temp_scaling = torch.softmax(logits, dim=-1) - -- if get_global_server_args().rl_on_policy_target is not None: -- logits_div_temperature = ( -- logits.bfloat16().div(sampling_info.temperatures).bfloat16() -- ) -- logprobs_via_logsoftmax_kernel = torch.log_softmax( -- logits_div_temperature, dim=-1 -- ) -- - # Post process logits - logits.div_(sampling_info.temperatures) -+ if get_global_server_args().rl_on_policy_target is not None: -+ logprobs_via_logsoftmax_kernel = torch.log_softmax(logits, dim=-1) -+ - # For ascend backend, softmax is not needed before sampling - if not get_global_server_args().sampling_backend == "ascend" or ( - return_logprob and not SGLANG_RETURN_ORIGINAL_LOGPROB -diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py -index 87922077e..8cb6bad8d 100644 ---- a/python/sglang/srt/managers/detokenizer_manager.py -+++ b/python/sglang/srt/managers/detokenizer_manager.py -@@ -247,6 +247,16 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): - s.sent_offset = len(output_str) - output_strs.append(incremental_output) - -+ output_routed_experts = [] -+ if recv_obj.output_routed_experts is not None: -+ output_routed_experts = [ -+ ( -+ output_routed_experts.tolist() -+ if output_routed_experts is not None -+ else [] -+ ) -+ for output_routed_experts in recv_obj.output_routed_experts -+ ] - return BatchStrOutput( - rids=recv_obj.rids, - http_worker_ipcs=recv_obj.http_worker_ipcs, -@@ -272,6 +282,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): - output_token_ids_logprobs_idx=recv_obj.output_token_ids_logprobs_idx, - output_token_entropy_val=recv_obj.output_token_entropy_val, - output_hidden_states=recv_obj.output_hidden_states, -+ output_routed_experts=output_routed_experts, - placeholder_tokens_idx=None, - placeholder_tokens_val=None, - retraction_counts=recv_obj.retraction_counts, -diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index e34736cc4..5e5997a1a 100644 ---- a/python/sglang/srt/managers/io_struct.py -+++ b/python/sglang/srt/managers/io_struct.py -@@ -23,6 +23,8 @@ from dataclasses import dataclass, field - from enum import Enum - from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union - -+import torch -+ - from sglang.srt.lora.lora_registry import LoRARef - from sglang.srt.managers.schedule_batch import BaseFinishReason - from sglang.srt.multimodal.mm_utils import has_valid_data -@@ -175,6 +177,8 @@ class GenerateReqInput(BaseReq): - log_metrics: bool = True - # Whether to return hidden states - return_hidden_states: Union[List[bool], bool] = False -+ # Whether to return captured routed experts -+ return_routed_experts: bool = False - - # The modalities of the image data [image, multi-images, video] - modalities: Optional[List[str]] = None -@@ -592,6 +596,7 @@ class GenerateReqInput(BaseReq): - if isinstance(self.return_hidden_states, list) - else self.return_hidden_states - ), -+ return_routed_experts=self.return_routed_experts, - modalities=self.modalities[i] if self.modalities else None, - session_params=self.session_params, - lora_path=self.lora_path[i] if self.lora_path is not None else None, -@@ -655,6 +660,9 @@ class TokenizedGenerateReqInput(BaseReq): - # Whether to return hidden states - return_hidden_states: bool = False - -+ # Whether to return captured routed experts -+ return_routed_experts: bool = False -+ - # The input embeds - input_embeds: Optional[Union[List[List[List[float]]], List[List[float]]]] = None - -@@ -910,6 +918,9 @@ class BatchTokenIDOutput( - # Hidden states - output_hidden_states: List[List[float]] - -+ # The routed experts for each output token -+ output_routed_experts: List[torch.Tensor] -+ - # The information of placeholder tokens (e.g., image token) - # idx is the index of the token in the prompt after expansion. - # val is the length of padded tokens after expansion. -@@ -989,6 +1000,9 @@ class BatchStrOutput( - # Hidden states - output_hidden_states: List[List[float]] - -+ # The routed experts for each output token -+ output_routed_experts: List[List[int]] -+ - # The information of placeholder tokens (e.g., image token) - # idx is the index of the token in the prompt after expansion. - # val is the length of padded tokens after expansion. -diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index c4c5a9ebb..1450c5fd8 100644 ---- a/python/sglang/srt/managers/schedule_batch.py -+++ b/python/sglang/srt/managers/schedule_batch.py -@@ -450,6 +450,7 @@ class Req: - session_id: Optional[str] = None, - custom_logit_processor: Optional[str] = None, - return_hidden_states: bool = False, -+ return_routed_experts: bool = False, - eos_token_ids: Optional[Set[int]] = None, - bootstrap_host: Optional[str] = None, - bootstrap_port: Optional[int] = None, -@@ -629,6 +630,12 @@ class Req: - self.output_topk_p = None - self.output_topk_index = None - -+ # capture routed experts -+ self.return_routed_experts = return_routed_experts -+ self.routed_experts: Optional[torch.Tensor] = ( -+ None # cpu tensor: shape (seqlen, topk) -+ ) -+ - # Embedding (return values) - self.embedding = None - -@@ -992,6 +999,7 @@ class Req: - self.retraction_count += 1 - - self.prefix_indices = torch.empty((0,), dtype=torch.int64) -+ self.routed_experts = [] - self.last_node = None - self.swa_uuid_for_lock = None - self.extend_input_len = 0 -@@ -1159,6 +1167,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - # Whether to return hidden states - return_hidden_states: bool = False - -+ # Whether to return captured experts -+ return_routed_experts: bool = False -+ - # Whether this batch is prefill-only (no token generation needed) - is_prefill_only: bool = False - -@@ -1206,6 +1217,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - device=req_to_token_pool.device, - spec_algorithm=spec_algorithm, - return_hidden_states=any(req.return_hidden_states for req in reqs), -+ return_routed_experts=any(req.return_routed_experts for req in reqs), - is_prefill_only=all(req.is_prefill_only for req in reqs), - chunked_req=chunked_req, - dllm_config=dllm_config, -@@ -1457,6 +1469,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - self.req_pool_indices = req_pool_indices_tensor - self.orig_seq_lens = orig_seq_lens_tensor - self.out_cache_loc = out_cache_loc -+ self.out_cache_loc_cpu = out_cache_loc.cpu() - self.input_embeds = ( - torch.tensor(input_embeds).to(self.device, non_blocking=True) - if input_embeds -@@ -1508,10 +1521,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - - input_ids = torch.cat([self.input_ids, running_batch.input_ids]) - out_cache_loc = torch.cat([self.out_cache_loc, running_batch.out_cache_loc]) -+ out_cache_loc_cpu = torch.cat( -+ [self.out_cache_loc_cpu, running_batch.out_cache_loc_cpu] -+ ) - - self.merge_batch(running_batch) - self.input_ids = input_ids - self.out_cache_loc = out_cache_loc -+ self.out_cache_loc_cpu = out_cache_loc_cpu - - # For overlap scheduler, the output_ids has one step delay - delta = 0 if self.enable_overlap else -1 -@@ -1677,6 +1694,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - self.seq_lens_cpu = torch.empty(0, dtype=torch.int64) - self.orig_seq_lens = torch.empty(0, dtype=torch.int32, device=self.device) - self.out_cache_loc = torch.empty(0, dtype=torch.int64, device=self.device) -+ self.out_cache_loc_cpu = torch.empty(0, dtype=torch.int64, device="cpu") - self.req_pool_indices = torch.empty(0, dtype=torch.int32, device=self.device) - self.seq_lens_sum = 0 - self.extend_num_tokens = 0 -@@ -1736,6 +1754,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - - # Allocate memory - self.out_cache_loc = alloc_for_decode(self, token_per_req=1) -+ self.out_cache_loc_cpu = self.out_cache_loc.to("cpu", non_blocking=True) - - # Update req-level memory management fields - for req in self.reqs: -@@ -1807,6 +1826,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - self.seq_lens_cpu = self.seq_lens_cpu[keep_indices] - self.orig_seq_lens = self.orig_seq_lens[keep_indices_device] - self.out_cache_loc = None -+ self.out_cache_loc_cpu = None - self.seq_lens_sum = self.seq_lens.sum().item() - self.output_ids = self.output_ids[keep_indices_device] - self.return_logprob = any(req.return_logprob for req in self.reqs) -@@ -1852,6 +1872,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - self.seq_lens_cpu = torch.cat([self.seq_lens_cpu, other.seq_lens_cpu]) - self.orig_seq_lens = torch.cat([self.orig_seq_lens, other.orig_seq_lens]) - self.out_cache_loc = None -+ self.out_cache_loc_cpu = None - self.seq_lens_sum += other.seq_lens_sum - if self.output_ids is not None: - self.output_ids = torch.cat([self.output_ids, other.output_ids]) -@@ -1903,6 +1924,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - seq_lens=self.seq_lens, - orig_seq_lens=self.orig_seq_lens, - out_cache_loc=self.out_cache_loc, -+ out_cache_loc_cpu=self.out_cache_loc_cpu, - seq_lens_cpu=seq_lens_cpu, - seq_lens_sum=self.seq_lens_sum, - return_logprob=self.return_logprob, -@@ -1983,7 +2005,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - def __str__(self): - return ( - f"ScheduleBatch(forward_mode={self.forward_mode.name if self.forward_mode else 'None'}, " -- f"#req={(len(self.reqs))})" -+ f"#req={(len(self.reqs))}), " -+ f"#out_cache_loc={self.out_cache_loc})" - ) - - -@@ -2038,6 +2061,9 @@ class ModelWorkerBatch: - # Sampling info - sampling_info: SamplingBatchInfo - -+ # cpu copy of out_cache_loc -+ out_cache_loc_cpu: Optional[torch.Tensor] = None -+ - # The original sequence lengths, Qwen-1M related - orig_seq_lens: Optional[torch.Tensor] = None - -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index b801fd8f8..9e27cc825 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -1305,6 +1305,7 @@ class Scheduler( - input_embeds=recv_req.input_embeds, - custom_logit_processor=recv_req.custom_logit_processor, - return_hidden_states=recv_req.return_hidden_states, -+ return_routed_experts=recv_req.return_routed_experts, - eos_token_ids=self.model_config.hf_eos_token_id, - bootstrap_host=recv_req.bootstrap_host, - bootstrap_port=recv_req.bootstrap_port, -diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -index c48f5f893..a9796c25f 100644 ---- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py -+++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -@@ -9,6 +9,7 @@ import torch - from sglang.srt.disaggregation.utils import DisaggregationMode - from sglang.srt.environ import envs - from sglang.srt.layers.logits_processor import LogitsProcessorOutput -+from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer - from sglang.srt.managers.io_struct import ( - AbortReq, - BatchEmbeddingOutput, -@@ -112,6 +113,14 @@ class SchedulerOutputProcessorMixin: - req.check_finished() - - if req.finished(): -+ req.routed_experts = ( -+ get_global_experts_capturer().get_routed_experts( -+ req_pool_idx=req.req_pool_idx, -+ seqlen=req.seqlen, -+ req_to_token_pool=self.req_to_token_pool, -+ ) -+ ) -+ - release_kv_cache(req, self.tree_cache) - req.time_stats.completion_time = time.perf_counter() - elif not batch.decoding_reqs or req not in batch.decoding_reqs: -@@ -362,6 +371,12 @@ class SchedulerOutputProcessorMixin: - req.check_finished(new_accepted_len) - - if req.finished(): -+ req.routed_experts = get_global_experts_capturer().get_routed_experts( -+ req_pool_idx=req.req_pool_idx, -+ seqlen=req.seqlen, -+ req_to_token_pool=self.req_to_token_pool, -+ ) -+ - if self.server_args.disaggregation_decode_enable_offload_kvcache: - # Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes - if not self.decode_offload_manager.offload_kv_cache(req): -@@ -756,6 +771,7 @@ class SchedulerOutputProcessorMixin: - spec_accepted_tokens = [] - retraction_counts = [] - output_hidden_states = None -+ output_routed_experts = None - - queue_times = [] - forward_entry_times = [] -@@ -946,6 +962,10 @@ class SchedulerOutputProcessorMixin: - if output_hidden_states is None: - output_hidden_states = [] - output_hidden_states.append(req.hidden_states) -+ if req.return_routed_experts: -+ if output_routed_experts is None: -+ output_routed_experts = [] -+ output_routed_experts.append(req.routed_experts) - - if ( - req.finished() -@@ -994,6 +1014,7 @@ class SchedulerOutputProcessorMixin: - output_token_ids_logprobs_idx=output_token_ids_logprobs_idx, - output_token_entropy_val=None, - output_hidden_states=output_hidden_states, -+ output_routed_experts=output_routed_experts, - placeholder_tokens_idx=None, - placeholder_tokens_val=None, - retraction_counts=retraction_counts, -diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -index f8ebfc1f4..a05449fac 100644 ---- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py -+++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -@@ -1,6 +1,7 @@ - from __future__ import annotations - - import logging -+import os - import traceback - from typing import TYPE_CHECKING, Tuple - -@@ -12,6 +13,9 @@ from sglang.srt.constants import ( - GPU_MEMORY_TYPE_KV_CACHE, - GPU_MEMORY_TYPE_WEIGHTS, - ) -+from sglang.srt.disaggregation.utils import DisaggregationMode -+from sglang.srt.distributed import get_moe_ep_group, get_moe_tp_group, get_tp_group -+from sglang.srt.layers.dp_attention import get_attention_tp_group - from sglang.srt.managers.io_struct import ( - CheckWeightsReqInput, - CheckWeightsReqOutput, -@@ -127,6 +131,13 @@ class SchedulerUpdateWeightsMixin: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) - self.flush_cache() - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.release_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.release_memory_occupation() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.stashed_model_static_state = _export_static_state( - self.tp_worker.model_runner.model -@@ -137,6 +148,20 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_CUDA_GRAPH in tags: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_CUDA_GRAPH) - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ tp_group = get_tp_group() -+ if tp_group is not None and tp_group.pynccl_comm is not None: -+ tp_group.pynccl_comm.nccl_pause() -+ attn_tp_group = get_attention_tp_group() -+ if attn_tp_group is not None and attn_tp_group.pynccl_comm is not None: -+ attn_tp_group.pynccl_comm.nccl_pause() -+ moe_ep_group = get_moe_ep_group() -+ if moe_ep_group is not None and moe_ep_group.pynccl_comm is not None: -+ moe_ep_group.pynccl_comm.nccl_pause() -+ moe_tp_group = get_moe_tp_group() -+ if moe_tp_group is not None and moe_tp_group.pynccl_comm is not None: -+ moe_tp_group.pynccl_comm.nccl_pause() -+ - torch.get_device_module().synchronize() - - return ReleaseMemoryOccupationReqOutput() -@@ -155,6 +180,20 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_CUDA_GRAPH in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_CUDA_GRAPH) - -+ if os.environ.get("AMEM_ENABLE", "0") == "1": -+ tp_group = get_tp_group() -+ if tp_group is not None and tp_group.pynccl_comm is not None: -+ tp_group.pynccl_comm.nccl_resume() -+ attn_tp_group = get_attention_tp_group() -+ if attn_tp_group is not None and attn_tp_group.pynccl_comm is not None: -+ attn_tp_group.pynccl_comm.nccl_resume() -+ moe_ep_group = get_moe_ep_group() -+ if moe_ep_group is not None and moe_ep_group.pynccl_comm is not None: -+ moe_ep_group.pynccl_comm.nccl_resume() -+ moe_tp_group = get_moe_tp_group() -+ if moe_tp_group is not None and moe_tp_group.pynccl_comm is not None: -+ moe_tp_group.pynccl_comm.nccl_resume() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_WEIGHTS) - torch.distributed.barrier(self.tp_cpu_group) -@@ -167,6 +206,13 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_KV_CACHE in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.resume_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.resume_memory_occupation() -+ - return ResumeMemoryOccupationReqOutput() - - def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): -diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index b90cf0616..98d71d896 100644 ---- a/python/sglang/srt/managers/tokenizer_manager.py -+++ b/python/sglang/srt/managers/tokenizer_manager.py -@@ -888,6 +888,7 @@ class TokenizerManager(TokenizerCommunicatorMixin): - session_params=session_params, - custom_logit_processor=obj.custom_logit_processor, - return_hidden_states=obj.return_hidden_states, -+ return_routed_experts=obj.return_routed_experts, - data_parallel_rank=obj.data_parallel_rank, - priority=obj.priority, - extra_key=obj.extra_key, -@@ -1621,6 +1622,9 @@ class TokenizerManager(TokenizerCommunicatorMixin): - if getattr(recv_obj, "output_hidden_states", None): - meta_info["hidden_states"] = recv_obj.output_hidden_states[i] - -+ if getattr(recv_obj, "output_routed_experts", None): -+ meta_info["routed_experts"] = recv_obj.output_routed_experts[i] -+ - if isinstance(recv_obj, BatchStrOutput): - state.text += recv_obj.output_strs[i] - if self.server_args.stream_output and state.obj.stream: -@@ -1747,12 +1751,13 @@ class TokenizerManager(TokenizerCommunicatorMixin): - return - - if len(recv_obj.input_token_logprobs_val) > 0: -- state.input_token_logprobs_val.extend( -- recv_obj.input_token_logprobs_val[recv_obj_index] -- ) -- state.input_token_logprobs_idx.extend( -- recv_obj.input_token_logprobs_idx[recv_obj_index] -- ) -+ if recv_obj.input_token_logprobs_val[recv_obj_index]: -+ state.input_token_logprobs_val.extend( -+ recv_obj.input_token_logprobs_val[recv_obj_index] -+ ) -+ state.input_token_logprobs_idx.extend( -+ recv_obj.input_token_logprobs_idx[recv_obj_index] -+ ) - state.output_token_logprobs_val.extend( - recv_obj.output_token_logprobs_val[recv_obj_index] - ) -diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py -index 3a85e6a7e..2859dafa1 100644 ---- a/python/sglang/srt/model_executor/forward_batch_info.py -+++ b/python/sglang/srt/model_executor/forward_batch_info.py -@@ -51,6 +51,7 @@ from sglang.srt.layers.dp_attention import ( - set_dp_buffer_len, - set_is_extend_in_batch, - ) -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import get_compiler_backend, is_npu, support_triton - from sglang.srt.utils.common import ceil_align - -@@ -214,6 +215,9 @@ class ForwardBatch: - # The sum of all sequence lengths - seq_lens_sum: int - -+ # cpu copy of out_cache_loc -+ out_cache_loc_cpu: Optional[torch.Tensor] = None -+ - # The original sequence length without being chunked. Qwen-1M related. - orig_seq_lens: Optional[torch.Tensor] = None - -@@ -368,6 +372,7 @@ class ForwardBatch: - req_pool_indices=batch.req_pool_indices, - seq_lens=batch.seq_lens, - out_cache_loc=batch.out_cache_loc, -+ out_cache_loc_cpu=batch.out_cache_loc_cpu, - mm_inputs=batch.multimodal_inputs, - encoder_cached=batch.encoder_cached, - encoder_lens=batch.encoder_lens, -@@ -623,7 +628,10 @@ class ForwardBatch: - mm_input = batch.multimodal_inputs[batch_idx] - if self.forward_mode.is_decode(): - # 3 * N -- if mm_input is None: -+ if ( -+ mm_input is None -+ or get_global_server_args().rl_on_policy_target is not None -+ ): - mrope_positions_list[batch_idx] = torch.full( - (3, 1), - self.seq_lens[batch_idx] - 1, -@@ -640,7 +648,10 @@ class ForwardBatch: - batch.extend_seq_lens[batch_idx], - batch.extend_prefix_lens[batch_idx], - ) -- if mm_input is None: -+ if ( -+ mm_input is None -+ or get_global_server_args().rl_on_policy_target is not None -+ ): - # text only - mrope_positions = torch.tensor( - [ -@@ -823,6 +834,10 @@ class ForwardBatch: - ) - - self.out_cache_loc = self._pad_tensor_to_size(self.out_cache_loc, num_tokens) -+ if self.out_cache_loc_cpu is not None: -+ self.out_cache_loc_cpu = self._pad_tensor_to_size( -+ self.out_cache_loc_cpu, num_tokens -+ ) - if self.encoder_lens is not None: - self.encoder_lens = self._pad_tensor_to_size(self.encoder_lens, bs) - self.positions = self._pad_tensor_to_size(self.positions, num_tokens) -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 4d58278b7..8f50dc430 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -94,6 +94,11 @@ from sglang.srt.layers.dp_attention import ( - set_is_extend_in_batch, - ) - from sglang.srt.layers.logits_processor import LogitsProcessorOutput -+from sglang.srt.layers.moe.routed_experts_capturer import ( -+ RoutedExpertsCapturer, -+ get_global_experts_capturer, -+ set_global_experts_capturer, -+) - from sglang.srt.layers.pooler import EmbeddingPoolerOutput - from sglang.srt.layers.sampler import Sampler - from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model -@@ -502,6 +507,10 @@ class ModelRunner: - server_args.max_running_requests, - server_args.max_total_tokens, - ) -+ -+ # Init routed experts capturer -+ self.init_routed_experts_capturer() -+ - if self.device == "cuda": - self.init_cublas() - self.init_attention_backend() -@@ -545,6 +554,40 @@ class ModelRunner: - # Initialize piecewise CUDA graph - self.init_piecewise_cuda_graphs() - -+ def init_routed_experts_capturer(self): -+ # TODO: the redundant logic with TpModelWorker -+ max_running_requests = min( -+ ( -+ self.max_total_num_tokens // 2 -+ if self.server_args.max_running_requests is None -+ else self.server_args.max_running_requests -+ // ( -+ self.server_args.dp_size -+ if self.server_args.enable_dp_attention -+ else 1 -+ ) -+ ), -+ self.req_to_token_pool.size, -+ ) -+ -+ if not self.server_args.disable_shared_experts_fusion and hasattr( -+ self.model, "num_fused_shared_experts" -+ ): -+ num_fused_shared_experts = self.model.num_fused_shared_experts -+ else: -+ num_fused_shared_experts = 0 -+ -+ set_global_experts_capturer( -+ RoutedExpertsCapturer.create( -+ enable=get_global_server_args().enable_return_routed_experts, -+ model_config=self.model_config, -+ num_fused_shared_experts=num_fused_shared_experts, -+ num_tokens=self.max_total_num_tokens + self.page_size, -+ max_running_requests=max_running_requests, -+ device=self.device, -+ ) -+ ) -+ - def model_specific_adjustment(self): - server_args = self.server_args - -@@ -792,7 +835,11 @@ class ModelRunner: - ) - with self.memory_saver_adapter.region( - GPU_MEMORY_TYPE_WEIGHTS, -- enable_cpu_backup=enable_cpu_backup, -+ enable_cpu_backup=( -+ self.server_args.enable_weights_cpu_backup -+ if not self.is_draft_worker -+ else True -+ ), - ): - self.model = get_model( - model_config=self.model_config, -@@ -2645,9 +2692,12 @@ class ModelRunner: - ) -> Tuple[Union[LogitsProcessorOutput, PPProxyTensors], bool]: - self.forward_pass_id += 1 - -- with get_global_expert_distribution_recorder().with_forward_pass( -- self.forward_pass_id, -- forward_batch, -+ with ( -+ get_global_expert_distribution_recorder().with_forward_pass( -+ self.forward_pass_id, -+ forward_batch, -+ ), -+ get_global_experts_capturer().with_forward(forward_batch), - ): - output = self._forward_raw( - forward_batch, -@@ -2656,6 +2706,13 @@ class ModelRunner: - reinit_attn_backend, - split_forward_count, - ) -+ # Copy cached routing experts' buffers back to CPU cache -+ get_global_experts_capturer().sync_fwd_experts_buffer_DtoH( -+ device_loc=forward_batch.out_cache_loc, -+ cpu_loc=forward_batch.out_cache_loc_cpu, -+ can_run_graph=output[1], -+ cuda_graph_batch=getattr(self.graph_runner, "bs", None), -+ ) - - if self.eplb_manager is not None: - self.eplb_manager.on_forward_pass_end() -diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py -index dc30b4f0a..f29dc4b71 100644 ---- a/python/sglang/srt/models/deepseek_v2.py -+++ b/python/sglang/srt/models/deepseek_v2.py -@@ -667,6 +667,7 @@ class DeepseekV2MoE(nn.Module): - - self.topk = TopK( - top_k=config.num_experts_per_tok + self.num_fused_shared_experts, -+ layer_id=self.layer_id, - renormalize=config.norm_topk_prob, - use_grouped_topk=True, - num_expert_group=config.n_group, -diff --git a/python/sglang/srt/models/ernie4.py b/python/sglang/srt/models/ernie4.py -index ab1b6576b..dffd8f09a 100644 ---- a/python/sglang/srt/models/ernie4.py -+++ b/python/sglang/srt/models/ernie4.py -@@ -87,6 +87,7 @@ class Ernie4Moe(nn.Module): - - self.topk = TopK( - top_k=config.moe_k, -+ layer_id=layer_id, - renormalize=True, - use_grouped_topk=False, - correction_bias=self.gate.e_score_correction_bias, -diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py -index a9689b8f2..bc8538da8 100644 ---- a/python/sglang/srt/models/glm4_moe.py -+++ b/python/sglang/srt/models/glm4_moe.py -@@ -379,6 +379,17 @@ class Glm4MoeSparseMoeBlock(nn.Module): - - self.gate = Glm4MoeGate(config=config, prefix=add_prefix("gate", prefix)) - -+ self.topk = TopK( -+ top_k=self.top_k, -+ layer_id=self.layer_id, -+ renormalize=config.norm_topk_prob, -+ use_grouped_topk=True, -+ num_expert_group=config.n_group, -+ topk_group=config.topk_group, -+ correction_bias=self.gate.e_score_correction_bias, -+ routed_scaling_factor=self.routed_scaling_factor, -+ ) -+ - self.experts = get_moe_impl_class(quant_config)( - num_experts=config.n_routed_experts + self.num_fused_shared_experts, - num_fused_shared_experts=self.num_fused_shared_experts, -diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py -index 9474700c4..398d622ff 100644 ---- a/python/sglang/srt/models/gpt_oss.py -+++ b/python/sglang/srt/models/gpt_oss.py -@@ -113,6 +113,7 @@ class GptOssSparseMoeBlock(nn.Module): - self.topk = TopK( - top_k=config.num_experts_per_tok, - renormalize=True, -+ layer_id=layer_id, - ) - - self.top_k = config.num_experts_per_tok -diff --git a/python/sglang/srt/models/grok.py b/python/sglang/srt/models/grok.py -index fd513060a..a089475b7 100644 ---- a/python/sglang/srt/models/grok.py -+++ b/python/sglang/srt/models/grok.py -@@ -142,6 +142,7 @@ class Grok1MoE(nn.Module): - self.topk = TopK( - top_k=top_k, - renormalize=False, -+ layer_id=layer_id, - custom_routing_function=custom_routing_function, - ) - -diff --git a/python/sglang/srt/models/hunyuan.py b/python/sglang/srt/models/hunyuan.py -index 7c6fd9e48..b20d28544 100644 ---- a/python/sglang/srt/models/hunyuan.py -+++ b/python/sglang/srt/models/hunyuan.py -@@ -150,6 +150,7 @@ class HunYuanSparseMoeBlock(nn.Module): - - self.topk = TopK( - top_k=top_k, -+ layer_id=layer_id, - renormalize=True if top_k > 1 else False, - ) - -diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py -index 3530609ba..01c89e893 100644 ---- a/python/sglang/srt/models/longcat_flash.py -+++ b/python/sglang/srt/models/longcat_flash.py -@@ -245,6 +245,7 @@ class LongcatFlashMoE(nn.Module): - renormalize=False, - use_grouped_topk=False, - correction_bias=self.router.e_score_correction_bias.data, -+ layer_id=layer_id, - ) - self.topk.forward = self.topk.forward_native - -diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py -index a7dbadec6..c83a41338 100644 ---- a/python/sglang/srt/models/qwen2.py -+++ b/python/sglang/srt/models/qwen2.py -@@ -90,9 +90,6 @@ class Qwen2MLP(nn.Module): - self.act_fn = SiluAndMul() - - def forward(self, x): -- if get_global_server_args().rl_on_policy_target is not None: -- x = x.bfloat16() -- - gate_up, _ = self.gate_up_proj(x) - x = self.act_fn(gate_up) - x, _ = self.down_proj(x) -@@ -279,11 +276,6 @@ class Qwen2Model(nn.Module): - quant_config=quant_config, - enable_tp=not is_dp_attention_enabled(), - prefix=add_prefix("embed_tokens", prefix), -- params_dtype=( -- torch.float32 -- if get_global_server_args().rl_on_policy_target is not None -- else None -- ), - ) - else: - self.embed_tokens = PPMissingLayer() -@@ -306,10 +298,8 @@ class Qwen2Model(nn.Module): - if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py -index ea33e81ef..561934dce 100644 ---- a/python/sglang/srt/models/qwen2_moe.py -+++ b/python/sglang/srt/models/qwen2_moe.py -@@ -161,6 +161,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module): - self.topk = TopK( - top_k=config.num_experts_per_tok, - renormalize=config.norm_topk_prob, -+ layer_id=layer_id, - ) - - self.experts = get_moe_impl_class(quant_config)( -@@ -581,7 +582,17 @@ class Qwen2MoeModel(nn.Module): - prefix=add_prefix("layers", prefix), - ) - if self.pp_group.is_last_rank: -- self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.norm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - else: - self.norm = PPMissingLayer(return_tuple=True) - -diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py -index 30b92acbd..a0d14895f 100644 ---- a/python/sglang/srt/models/qwen3.py -+++ b/python/sglang/srt/models/qwen3.py -@@ -90,8 +90,8 @@ class Qwen3Attention(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -@@ -256,10 +256,8 @@ class Qwen3DecoderLayer(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -@@ -289,10 +287,14 @@ class Qwen3DecoderLayer(nn.Module): - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - residual: Optional[torch.Tensor], -+ post_residual_addition: Optional[torch.Tensor] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: - # Self Attention - hidden_states, residual = self.layer_communicator.prepare_attn( -- hidden_states, residual, forward_batch -+ hidden_states, -+ residual, -+ forward_batch, -+ post_residual_addition=post_residual_addition, - ) - if hidden_states.shape[0] != 0: - hidden_states = self.self_attn( -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index 9737ac719..09c756918 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -22,6 +22,7 @@ import math - from typing import Any, Dict, Iterable, List, Optional, Tuple, TypeVar - - import torch -+import torch.nn.functional as F - from torch import nn - from transformers import PretrainedConfig - -@@ -50,7 +51,7 @@ from sglang.srt.layers.moe import ( - ) - from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class - from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE --from sglang.srt.layers.moe.topk import TopK -+from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK - from sglang.srt.layers.moe.utils import RoutingMethodType - from sglang.srt.layers.quantization.base_config import QuantizationConfig - from sglang.srt.layers.radix_attention import RadixAttention -@@ -227,7 +228,9 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - top_k=config.num_experts_per_tok, - renormalize=config.norm_topk_prob, - use_grouped_topk=False, -+ layer_id=layer_id, - ) -+ self.top_k = config.num_experts_per_tok - - self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts -@@ -293,7 +296,22 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - - # router_logits: (num_tokens, n_experts) - router_logits, _ = self.gate(hidden_states) -- topk_output = self.topk(hidden_states, router_logits) -+ -+ if get_global_server_args().rl_on_policy_target is not None: -+ routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) -+ routing_weights, selected_experts = torch.topk( -+ routing_weights, self.top_k, dim=-1 -+ ) -+ routing_weights /= routing_weights.sum(dim=-1, keepdim=True) -+ routing_weights = routing_weights.to(hidden_states.dtype) -+ topk_output = StandardTopKOutput( -+ topk_weights=routing_weights, -+ topk_ids=selected_experts, -+ router_logits=router_logits, -+ ) -+ else: -+ topk_output = self.topk(hidden_states, router_logits) -+ - final_hidden_states = self.experts(hidden_states, topk_output) - if ( - self.tp_size > 1 -@@ -474,13 +492,14 @@ class Qwen3MoeAttention(nn.Module): - ) - self.compatible_with_fused_kv_buffer = ( - False if isinstance(self.rotary_emb, MRotaryEmbedding) else True -- ) -+ ) and (get_global_server_args().rl_on_policy_target is None) - self.compatible_with_fused_qk_norm_rope = ( - not isinstance(self.rotary_emb, MRotaryEmbedding) - ) and self.head_dim in (64, 128, 256) - self.use_fused_qk_norm_rope = ( - get_global_server_args().enable_fused_qk_norm_rope - and self.compatible_with_fused_qk_norm_rope -+ and (get_global_server_args().rl_on_policy_target is None) - ) - self._used_fused_qk_norm_rope_last_call = False - -@@ -493,8 +512,16 @@ class Qwen3MoeAttention(nn.Module): - prefix=add_prefix("attn", prefix), - ) - -- self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -- self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) -+ self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) - self.alt_stream = alt_stream - - def _apply_qk_norm( -@@ -751,9 +778,19 @@ class Qwen3MoeDecoderLayer(nn.Module): - quant_config=quant_config, - prefix=add_prefix("mlp", prefix), - ) -- self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.input_layernorm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - self.post_attention_layernorm = RMSNorm( -- config.hidden_size, eps=config.rms_norm_eps -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs - ) - - self.layer_communicator = LayerCommunicator( -diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py -index ed52f7ff4..8ce9fab9d 100644 ---- a/python/sglang/srt/models/qwen3_vl.py -+++ b/python/sglang/srt/models/qwen3_vl.py -@@ -18,7 +18,6 @@ import re - from functools import lru_cache, partial - from typing import Callable, Iterable, List, Optional, Tuple, Union - --import numpy as np - import torch - import torch.nn as nn - from einops import rearrange -@@ -349,83 +348,65 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): - return rotary_pos_emb - - def fast_pos_embed_interpolate(self, grid_thw): -+ grid_ts, grid_hs, grid_ws = grid_thw[:, 0], grid_thw[:, 1], grid_thw[:, 2] - num_grid_per_side = int(self.num_position_embeddings**0.5) -+ device = self.pos_embed.weight.device - - idx_list = [[] for _ in range(4)] - weight_list = [[] for _ in range(4)] - -- # TODO: use torch instand of np -- for t, h, w in grid_thw: -- h_idxs = np.linspace(0, num_grid_per_side - 1, h) -- w_idxs = np.linspace(0, num_grid_per_side - 1, w) -+ for t, h, w in zip(grid_ts, grid_hs, grid_ws): -+ h_idxs = torch.linspace(0, num_grid_per_side - 1, h) -+ w_idxs = torch.linspace(0, num_grid_per_side - 1, w) - -- h_idxs_floor = h_idxs.astype(int) -- w_idxs_floor = w_idxs.astype(int) -- h_idxs_ceil = (h_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1) -- w_idxs_ceil = (w_idxs.astype(int) + 1).clip(max=num_grid_per_side - 1) -+ h_idxs_floor = h_idxs.int() -+ w_idxs_floor = w_idxs.int() -+ h_idxs_ceil = (h_idxs.int() + 1).clip(max=num_grid_per_side - 1) -+ w_idxs_ceil = (w_idxs.int() + 1).clip(max=num_grid_per_side - 1) - - dh = h_idxs - h_idxs_floor - dw = w_idxs - w_idxs_floor - -- idx_list[0].extend( -- ((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_floor[None]) -- .flatten() -- .tolist() -- * t -- ) -- idx_list[1].extend( -- ((h_idxs_floor * num_grid_per_side)[None].T + w_idxs_ceil[None]) -- .flatten() -- .tolist() -- * t -- ) -- idx_list[2].extend( -- ((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_floor[None]) -- .flatten() -- .tolist() -- * t -- ) -- idx_list[3].extend( -- ((h_idxs_ceil * num_grid_per_side)[None].T + w_idxs_ceil[None]) -- .flatten() -- .tolist() -- * t -- ) -+ base_h = h_idxs_floor * num_grid_per_side -+ base_h_ceil = h_idxs_ceil * num_grid_per_side - -- weight_list[0].extend( -- ((1 - dh)[None].T * (1 - dw)[None]).flatten().tolist() * t -- ) -- weight_list[1].extend(((1 - dh)[None].T * dw[None]).flatten().tolist() * t) -- weight_list[2].extend((dh[None].T * (1 - dw)[None]).flatten().tolist() * t) -- weight_list[3].extend((dh[None].T * dw[None]).flatten().tolist() * t) -+ indices = [ -+ (base_h[None].T + w_idxs_floor[None]).flatten(), -+ (base_h[None].T + w_idxs_ceil[None]).flatten(), -+ (base_h_ceil[None].T + w_idxs_floor[None]).flatten(), -+ (base_h_ceil[None].T + w_idxs_ceil[None]).flatten(), -+ ] - -- device = self.pos_embed.weight.device -- dtype = self.pos_embed.weight.dtype -+ weights = [ -+ ((1 - dh)[None].T * (1 - dw)[None]).flatten(), -+ ((1 - dh)[None].T * dw[None]).flatten(), -+ (dh[None].T * (1 - dw)[None]).flatten(), -+ (dh[None].T * dw[None]).flatten(), -+ ] - -- p0 = ( -- self.pos_embed(torch.tensor(idx_list[0], dtype=torch.long, device=device)) -- * torch.tensor(weight_list[0], dtype=dtype, device=device)[:, None] -- ) -- p1 = ( -- self.pos_embed(torch.tensor(idx_list[1], dtype=torch.long, device=device)) -- * torch.tensor(weight_list[1], dtype=dtype, device=device)[:, None] -- ) -- p2 = ( -- self.pos_embed(torch.tensor(idx_list[2], dtype=torch.long, device=device)) -- * torch.tensor(weight_list[2], dtype=dtype, device=device)[:, None] -+ for i in range(4): -+ idx_list[i].extend(indices[i].tolist()) -+ weight_list[i].extend(weights[i].tolist()) -+ -+ idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=device) -+ weight_tensor = torch.tensor( -+ weight_list, dtype=self.pos_embed.weight.dtype, device=device - ) -- p3 = ( -- self.pos_embed(torch.tensor(idx_list[3], dtype=torch.long, device=device)) -- * torch.tensor(weight_list[3], dtype=dtype, device=device)[:, None] -+ pos_embeds = self.pos_embed(idx_tensor).to(device) * weight_tensor[:, :, None] -+ patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] -+ -+ patch_pos_embeds = patch_pos_embeds.split( -+ [h * w for h, w in zip(grid_hs, grid_ws)] - ) - -- patch_pos_embeds = p0 + p1 + p2 + p3 -- patch_pos_embeds = patch_pos_embeds.split([t * h * w for t, h, w in grid_thw]) - patch_pos_embeds_permute = [] -- m_size = self.spatial_merge_size -- for pos_embed, (t, h, w) in zip(patch_pos_embeds, grid_thw): -+ merge_size = self.spatial_merge_size -+ for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): -+ pos_embed = pos_embed.repeat(t, 1) - pos_embed = ( -- pos_embed.view(t, h // m_size, m_size, w // m_size, m_size, -1) -+ pos_embed.view( -+ t, h // merge_size, merge_size, w // merge_size, merge_size, -1 -+ ) - .permute(0, 1, 3, 2, 4, 5) - .flatten(0, 4) - ) -@@ -555,21 +536,27 @@ class Qwen3LLMModel(Qwen3Model): - hidden_states + residual if residual is not None else hidden_states - ) - -+ deepstack_embeds = None -+ if input_deepstack_embeds is not None: -+ prev_layer_idx = layer_idx - 1 -+ if prev_layer_idx in self.deepstack_embed_to_decoder_layer: -+ sep = self.hidden_size * prev_layer_idx -+ deepstack_embeds = input_deepstack_embeds[ -+ :, sep : sep + self.hidden_size -+ ] -+ -+ # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. -+ # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 -+ # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack -+ # The order matters because addition with different tensors is not associative in practice. - hidden_states, residual = layer( - positions, - hidden_states, - forward_batch, - residual, -+ post_residual_addition=deepstack_embeds, - ) - -- # process deepstack -- if ( -- input_deepstack_embeds is not None -- and layer_idx in self.deepstack_embed_to_decoder_layer -- ): -- sep = self.hidden_size * layer_idx -- hidden_states += input_deepstack_embeds[:, sep : sep + self.hidden_size] -- - if not self.pp_group.is_last_rank: - return PPProxyTensors( - { -diff --git a/python/sglang/srt/models/step3_vl.py b/python/sglang/srt/models/step3_vl.py -index 4474f62d5..0e537c398 100644 ---- a/python/sglang/srt/models/step3_vl.py -+++ b/python/sglang/srt/models/step3_vl.py -@@ -129,6 +129,7 @@ class Step3TextMoEMLP(nn.Module): - top_k=config.moe_top_k, - renormalize=config.norm_expert_weight, - use_grouped_topk=False, -+ layer_id=layer_id, - ) - - self.experts = get_moe_impl_class(quant_config)( -diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py -index 370aec2b6..47666d8f3 100644 ---- a/python/sglang/srt/multimodal/processors/base_processor.py -+++ b/python/sglang/srt/multimodal/processors/base_processor.py -@@ -13,6 +13,7 @@ from PIL import Image - from transformers import BaseImageProcessorFast - - from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - get_bool_env_var, - is_npu, -@@ -260,7 +261,9 @@ class BaseMultimodalProcessor(ABC): - and isinstance(processor.image_processor, BaseImageProcessorFast) - and not self.server_args.disable_fast_image_processor - ): -- if not _is_npu: -+ if get_global_server_args().rl_on_policy_target is not None: -+ kwargs["device"] = "cpu" -+ elif not _is_npu: - kwargs["device"] = "cuda" - elif processor.__class__.__name__ not in { - "Qwen2_5_VLProcessor", -diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index 8e7753dab..323788f39 100644 ---- a/python/sglang/srt/server_args.py -+++ b/python/sglang/srt/server_args.py -@@ -535,6 +535,7 @@ class ServerArgs: - disable_fast_image_processor: bool = False - keep_mm_feature_on_device: bool = False - enable_return_hidden_states: bool = False -+ enable_return_routed_experts: bool = False - scheduler_recv_interval: int = 1 - numa_node: Optional[List[int]] = None - enable_deterministic_inference: bool = False -@@ -1966,6 +1967,9 @@ class ServerArgs: - "Enable deterministic inference because of rl_on_policy_target." - ) - self.enable_deterministic_inference = True -+ -+ # For VLM -+ os.environ["SGLANG_VLM_CACHE_SIZE_MB"] = "0" - # TODO remove this environment variable as a whole - os.environ["SGLANG_ENABLE_DETERMINISTIC_INFERENCE"] = "1" - -@@ -3705,6 +3709,11 @@ class ServerArgs: - action="store_true", - help="Enable returning hidden states with responses.", - ) -+ parser.add_argument( -+ "--enable-return-routed-experts", -+ action="store_true", -+ help="Enable returning routed experts of each layer with responses.", -+ ) - parser.add_argument( - "--scheduler-recv-interval", - type=int, -diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index b3d72df05..ddfe0b178 100644 ---- a/python/sglang/srt/speculative/eagle_info.py -+++ b/python/sglang/srt/speculative/eagle_info.py -@@ -746,6 +746,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] - self.verified_id = self.verified_id[: len(new_indices)] -+ if self.accept_length is not None: -+ self.accept_length = self.accept_length[: len(new_indices)] -+ if self.accept_length_cpu is not None: -+ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] - else: - # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` - self.topk_p = self.topk_p[new_indices] -@@ -777,6 +781,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) - self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) - self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) -+ if self.accept_length is not None and spec_info.accept_length is not None: -+ self.accept_length = torch.cat( -+ [self.accept_length, spec_info.accept_length] -+ ) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif self.accept_length is not None: -+ zeros = torch.zeros( -+ [spec_info.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([self.accept_length, zeros]) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif spec_info.accept_length is not None: -+ zeros = torch.zeros( -+ [self.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([zeros, spec_info.accept_length]) -+ self.accept_length_cpu = self.accept_length.tolist() - - - @dataclass diff --git a/docker/patch/v0.5.7/sglang.patch b/docker/patch/v0.5.7/sglang.patch deleted file mode 100644 index 8c2b46fb7..000000000 --- a/docker/patch/v0.5.7/sglang.patch +++ /dev/null @@ -1,1290 +0,0 @@ -diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py -index aa10cb08d..d41c31a09 100644 ---- a/python/sglang/srt/configs/model_config.py -+++ b/python/sglang/srt/configs/model_config.py -@@ -268,6 +268,12 @@ class ModelConfig: - ): - self.hf_config.architectures[0] = "DeepseekV3ForCausalLMNextN" - -+ if ( -+ is_draft_model -+ and self.hf_config.architectures[0] == "DeepseekV32ForCausalLM" -+ ): -+ self.hf_config.architectures[0] = "DeepseekV3ForCausalLMNextN" -+ - if is_draft_model and self.hf_config.architectures[0] == "Glm4MoeForCausalLM": - self.hf_config.architectures[0] = "Glm4MoeForCausalLMNextN" - -diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py -index 51af67636..54716de5c 100644 ---- a/python/sglang/srt/disaggregation/decode.py -+++ b/python/sglang/srt/disaggregation/decode.py -@@ -315,6 +315,13 @@ class DecodePreallocQueue: - ) - return kv_manager - -+ def release_memory_occupation(self): -+ if hasattr(self.kv_manager, "close"): -+ self.kv_manager.close() -+ -+ def resume_memory_occupation(self): -+ self.kv_manager = self._init_kv_manager() -+ - def add(self, req: Req, is_retracted: bool = False) -> None: - """Add a request to the pending queue.""" - if self._check_if_req_exceed_kv_capacity(req): -diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py -index 32e8c0b69..df913da7b 100644 ---- a/python/sglang/srt/disaggregation/mooncake/conn.py -+++ b/python/sglang/srt/disaggregation/mooncake/conn.py -@@ -1079,6 +1079,19 @@ class MooncakeKVManager(CommonKVManager): - f"Losing connection with prefill instance (bootstrap_addr: {failed_bootstrap_addr}), {len(affected_rooms)} requests affected" - ) - -+ def close(self): -+ # Batch deregister KV data buffers -+ if self.kv_args.kv_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.kv_data_ptrs) -+ -+ # Batch deregister auxiliary data buffers -+ if self.kv_args.aux_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.aux_data_ptrs) -+ -+ # Batch deregister state/extra pool data buffers -+ if self.kv_args.state_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.state_data_ptrs) -+ - - class MooncakeKVSender(CommonKVSender): - -diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py -index a6eed743a..0124d8917 100644 ---- a/python/sglang/srt/disaggregation/prefill.py -+++ b/python/sglang/srt/disaggregation/prefill.py -@@ -306,6 +306,13 @@ class PrefillBootstrapQueue: - else: - return bootstrapped_reqs, failed_reqs - -+ def release_memory_occupation(self): -+ if hasattr(self.kv_manager, "close"): -+ self.kv_manager.close() -+ -+ def resume_memory_occupation(self): -+ self.kv_manager = self._init_kv_manager() -+ - - class SchedulerDisaggregationPrefillMixin: - """ -diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py -index 0478526ef..cfb1aa669 100644 ---- a/python/sglang/srt/distributed/parallel_state.py -+++ b/python/sglang/srt/distributed/parallel_state.py -@@ -1797,7 +1797,10 @@ def get_tensor_model_parallel_world_size(): - - def get_tensor_model_parallel_rank(): - """Return my rank for the tensor model parallel group.""" -- return get_tp_group().rank_in_group -+ try: -+ return get_tp_group().rank_in_group -+ except Exception: -+ return 0 - - - def get_pipeline_model_parallel_world_size(): -diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py -index 6f69fd19b..da20ac2ed 100644 ---- a/python/sglang/srt/entrypoints/engine.py -+++ b/python/sglang/srt/entrypoints/engine.py -@@ -49,6 +49,7 @@ from sglang.srt.managers.io_struct import ( - InitWeightsUpdateGroupReqInput, - LoadLoRAAdapterReqInput, - MultimodalDataInputFormat, -+ PostProcessWeightsReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, - RpcReqInput, -@@ -593,6 +594,24 @@ class Engine(EngineBase): - self.tokenizer_manager.update_weights_from_ipc(obj, None) - ) - -+ def post_process_weights( -+ self, -+ restore_weights_before_load: bool = False, -+ post_process_quantization: bool = False, -+ ): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ Should be called after weight update is finished. -+ """ -+ obj = PostProcessWeightsReqInput( -+ restore_weights_before_load=restore_weights_before_load, -+ post_process_quantization=post_process_quantization, -+ ) -+ -+ return self.loop.run_until_complete( -+ self.tokenizer_manager.post_process_weights(obj, None) -+ ) -+ - def get_weights_by_name(self, name: str, truncate_size: int = 100): - """Get weights by parameter name.""" - obj = GetWeightsByNameReqInput(name=name, truncate_size=truncate_size) -diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py -index 88705cc35..c8dc052f1 100644 ---- a/python/sglang/srt/entrypoints/http_server.py -+++ b/python/sglang/srt/entrypoints/http_server.py -@@ -107,6 +107,7 @@ from sglang.srt.managers.io_struct import ( - OpenSessionReqInput, - ParseFunctionCallReq, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -957,6 +958,21 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re - else: - return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) - -+@app.post("/post_process_weights") -+async def post_process_weights(req: PostProcessWeightsReqInput, request: Request): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ This should be called selectively after `update_weights_from_distributed/update_weights_from_tensor`. -+ """ -+ success, message = await _global_state.tokenizer_manager.post_process_weights( -+ req, request -+ ) -+ -+ content = {"success": success, "message": message} -+ return ORJSONResponse( -+ content, status_code=200 if success else HTTPStatus.BAD_REQUEST -+ ) -+ - - @app.post("/update_weight_version") - async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request): -diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -index c9e82e4b1..f2584546a 100644 ---- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -+++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -@@ -3,6 +3,7 @@ from __future__ import annotations - from abc import ABC, abstractmethod - from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple - -+import os - import torch - from einops import rearrange - -@@ -178,7 +179,7 @@ class Indexer(MultiPlatformOp): - max_position=max_position_embeddings, - base=rope_theta, # type: ignore - rope_scaling=rope_scaling, -- is_neox_style=True, -+ is_neox_style=True if os.environ.get("INDEXER_ROPE_NEOX_STYLE", "1") == "1" else False, - device=get_global_server_args().device, - ) - self.block_size = block_size -@@ -188,6 +189,9 @@ class Indexer(MultiPlatformOp): - @torch.compile(dynamic=True) - def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor): - weights, _ = self.weights_proj(x.float()) -+ if weights.shape[1] < 32: -+ assert 32 % weights.shape[1] == 0 -+ weights = weights.repeat_interleave(32 // weights.shape[1], dim=1) - weights = weights * self.n_heads**-0.5 - weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale - return weights -@@ -837,6 +841,9 @@ class Indexer(MultiPlatformOp): - query, key = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch - ) -+ if query.shape[1] < 32: -+ assert 32 % query.shape[1] == 0 -+ query = query.repeat_interleave(32//query.shape[1], dim=1) - - if enable_dual_stream: - current_stream = torch.cuda.current_stream() -diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py -index 7bef9d2ab..f588cbdb0 100644 ---- a/python/sglang/srt/layers/layernorm.py -+++ b/python/sglang/srt/layers/layernorm.py -@@ -83,15 +83,12 @@ class RMSNorm(MultiPlatformOp): - eps: float = 1e-6, - var_hidden_size: Optional[int] = None, - cast_x_before_out_mul: bool = False, -- fp32_residual: bool = False, -- weight_dtype: Optional = None, -- override_orig_dtype: Optional = None, -+ fp32_residual: bool = True, - ) -> None: - super().__init__() - self.cast_x_before_out_mul = cast_x_before_out_mul - self.fp32_residual = fp32_residual -- self.override_orig_dtype = override_orig_dtype -- self.weight = nn.Parameter(torch.ones(hidden_size, dtype=weight_dtype)) -+ self.weight = nn.Parameter(torch.ones(hidden_size)) - self.variance_epsilon = eps - self.hidden_size = hidden_size - self.variance_size_override = ( -@@ -193,10 +190,22 @@ class RMSNorm(MultiPlatformOp): - ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: - if not x.is_contiguous(): - x = x.contiguous() -- orig_dtype = self.override_orig_dtype or x.dtype -+ orig_dtype = x.dtype - post_residual_addition = kwargs.get("post_residual_addition") -+ -+ if residual is not None and not self.fp32_residual: -+ x = ( -+ x -+ + residual -+ + ( -+ post_residual_addition -+ if post_residual_addition is not None -+ else 0.0 -+ ) -+ ) -+ residual = x.clone() - x = x.to(torch.float32) -- if residual is not None: -+ if residual is not None and self.fp32_residual: - x = ( - x - + residual.to(torch.float32) -@@ -206,10 +215,7 @@ class RMSNorm(MultiPlatformOp): - else 0.0 - ) - ) -- if self.fp32_residual: -- residual = x.clone() -- else: -- residual = x.to(orig_dtype) -+ residual = x.to(orig_dtype) - - hidden_size = x.shape[-1] - if hidden_size != self.hidden_size: -diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py -index fa7431048..cd33ea735 100644 ---- a/python/sglang/srt/layers/logits_processor.py -+++ b/python/sglang/srt/layers/logits_processor.py -@@ -878,11 +878,6 @@ class LogitsProcessor(nn.Module): - None, # bias - True, # is_vnni - ) -- elif get_global_server_args().rl_on_policy_target is not None: -- # Due to tie-weight, we may not be able to change lm_head's weight dtype -- logits = torch.matmul( -- hidden_states.bfloat16(), lm_head.weight.T.bfloat16() -- ) - else: - logits = torch.matmul( - hidden_states.to(lm_head.weight.dtype), lm_head.weight.T -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -index a1885fade..14d692365 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/fused_moe.py -@@ -14,6 +14,7 @@ import torch.nn.functional as F - import triton.language as tl - - from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig -+from sglang.srt.server_args import get_global_server_args - from sglang.srt.utils import ( - cpu_has_amx_support, - get_bool_env_var, -@@ -573,7 +574,10 @@ def fused_experts_impl( - ).squeeze(dim=1) - else: - # According to micro benchmark results, torch.compile can get better performance for small token. -- if tokens_in_chunk <= 32: -+ if ( -+ not get_global_server_args().enable_deterministic_inference -+ and tokens_in_chunk <= 32 -+ ): - moe_sum_reduce_torch_compile( - intermediate_cache3.view(*intermediate_cache3.shape), - out_hidden_states[begin_chunk_idx:end_chunk_idx], -diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py -index 00bd68755..5a3ca8a67 100644 ---- a/python/sglang/srt/layers/moe/routed_experts_capturer.py -+++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py -@@ -1,5 +1,6 @@ - import logging - from abc import ABC -+from contextlib import contextmanager - from typing import Optional - - import numpy as np -@@ -8,13 +9,18 @@ import torch - - from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.layers.dp_attention import ( -+ attn_tp_all_gather_into_tensor, - get_attention_dp_rank, -+ get_attention_tp_size, - get_dp_local_info, - is_dp_attention_enabled, - ) - from sglang.srt.mem_cache.memory_pool import ReqToTokenPool - from sglang.srt.model_executor.forward_batch_info import ForwardBatch - from sglang.srt.server_args import get_global_server_args -+from sglang.srt.layers.moe import ( -+ get_moe_a2a_backend, -+) - - logger = logging.getLogger(__name__) - -@@ -181,13 +187,26 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - device=device, - ) - -+ if get_moe_a2a_backend().is_deepep(): -+ attn_tp_size = get_attention_tp_size() if is_dp_attention_enabled() else 1 -+ self.gather_buffer = torch.empty( -+ ( -+ self.device_cache.buffer.shape[0] * attn_tp_size, -+ self.device_cache.buffer.shape[2], -+ ), -+ dtype=torch.int32, -+ device=device, -+ ) -+ - def _sync_fwd_experts_buffer_DtoH( - self, - forward_batch: ForwardBatch, - can_run_graph: bool, - cuda_graph_batch: int, - ): -- if is_dp_attention_enabled(): -+ # When DeepEP is enabled, capture() already does all_gather, so device_cache.buffer -+ # contains data from all DP ranks. We should not slice by DP rank in this case. -+ if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep(): - local_start_pos, local_num_tokens = get_dp_local_info(forward_batch) - # handle with cuda graph padding - if can_run_graph: -@@ -206,6 +225,12 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - ].cpu() - - def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ if get_moe_a2a_backend().is_deepep(): -+ local_topk_ids = topk_ids -+ topk_ids = self.gather_buffer[ -+ : local_topk_ids.size(0) * get_attention_tp_size() -+ ] -+ attn_tp_all_gather_into_tensor(topk_ids, local_topk_ids) - self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) - - def get_routed_experts( -diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py -index c5e5a11fc..dd321fa13 100644 ---- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py -+++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py -@@ -1016,13 +1016,37 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod): - layer.a2_scale = None - layer.marlin_state = GPTQMarlinState.REPACK - -+ if not hasattr(layer, "_original_shapes"): -+ layer._original_shapes = {} -+ -+ # Force record: these are the target GPTQ shapes for rollback. -+ layer._original_shapes["w13_weight_packed"] = tuple(w13_weight.shape) -+ layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) -+ -+ # Also record the shapes of the scales. -+ layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) -+ layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) -+ - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: -+ # Skip if the layer is already converted to Marlin format to prevent double-packing. -+ if getattr(layer, "is_marlin_converted", False): -+ return -+ -+ if not hasattr(layer, "_original_shapes"): -+ layer._original_shapes = {} - - def replace_tensor(name, new_t): -+ target_attr = getattr(layer, name) -+ -+ # Only save if the key doesn't exist to prevent overwriting with Marlin shapes. -+ if name not in layer._original_shapes: -+ # This is a safety check; `create_weights` usually handles this already. -+ layer._original_shapes[name] = tuple(target_attr.shape) -+ - # It is important to use resize_() here since it ensures - # the same buffer is reused -- getattr(layer, name).resize_(new_t.shape) -- getattr(layer, name).copy_(new_t) -+ target_attr.resize_(new_t.shape) -+ target_attr.copy_(new_t) - del new_t - - num_experts = layer.w13_weight_g_idx.shape[0] -@@ -1078,7 +1102,7 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod): - layer.w13_weight_packed.shape[2], - self.num_bits, - ) -- replace_parameter(layer, "w13_weight_packed", marlin_w13_qweight) -+ replace_tensor("w13_weight_packed", marlin_w13_qweight) - marlin_w2_qweight = gptq_marlin_moe_repack( - layer.w2_weight_packed, - layer.w2_g_idx_sort_indices, -@@ -1086,7 +1110,7 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod): - layer.w2_weight_packed.shape[2], - self.num_bits, - ) -- replace_parameter(layer, "w2_weight_packed", marlin_w2_qweight) -+ replace_tensor("w2_weight_packed", marlin_w2_qweight) - # Repack scales - marlin_w13_scales = marlin_moe_permute_scales( - layer.w13_weight_scale, -@@ -1094,7 +1118,7 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod): - layer.w13_weight_scale.shape[2], - self.group_size, - ) -- replace_parameter(layer, "w13_weight_scale", marlin_w13_scales) -+ replace_tensor("w13_weight_scale", marlin_w13_scales) - - marlin_w2_scales = marlin_moe_permute_scales( - layer.w2_weight_scale, -@@ -1103,7 +1127,22 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod): - layer.w2_weight_scale.shape[2], - self.group_size, - ) -- replace_parameter(layer, "w2_weight_scale", marlin_w2_scales) -+ replace_tensor("w2_weight_scale", marlin_w2_scales) -+ -+ layer.is_marlin_converted = True -+ -+ def restore_weights_before_loading(self, layer: torch.nn.Module): -+ """Forcibly resize parameters back to their original shapes (e.g., GPTQ format) before loading weights.""" -+ if not hasattr(layer, "_original_shapes"): -+ return -+ -+ for name, orig_shape in layer._original_shapes.items(): -+ param = getattr(layer, name, None) -+ -+ if param is not None and param.shape != orig_shape: -+ param.resize_(orig_shape) -+ -+ layer.is_marlin_converted = False - - def create_moe_runner( - self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig -diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py -index 480579e01..dd8ca7d4f 100644 ---- a/python/sglang/srt/layers/rotary_embedding.py -+++ b/python/sglang/srt/layers/rotary_embedding.py -@@ -136,9 +136,7 @@ class RotaryEmbedding(MultiPlatformOp): - - if get_global_server_args().rl_on_policy_target is not None: - self._forward_method = self.forward_native -- self._apply_rotary_emb_wrapped = torch.compile(dynamic=True)( -- self._apply_rotary_emb_wrapped -- ) -+ - self.position_cos, self.position_sin = None, None - - def _compute_inv_freq(self, base: Union[int, float]) -> torch.Tensor: -@@ -1578,6 +1576,9 @@ class MRotaryEmbedding(RotaryEmbedding): - key: torch.Tensor, - fused_set_kv_buffer_arg: Optional[FusedSetKVBufferArg] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: -+ assert ( -+ fused_set_kv_buffer_arg is None -+ ), "fused_set_kv_buffer_arg is not supported for npu implementation" - # TODO: remove this when npu_mrope supports QNumHeads * QHeadSize > 4096 - assert ( - fused_set_kv_buffer_arg is None -diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py -index 55bef5652..35ad68b1c 100644 ---- a/python/sglang/srt/layers/sampler.py -+++ b/python/sglang/srt/layers/sampler.py -@@ -108,16 +108,11 @@ class Sampler(nn.Module): - if return_logprob and SGLANG_RETURN_ORIGINAL_LOGPROB: - probs_without_temp_scaling = torch.softmax(logits, dim=-1) - -- if get_global_server_args().rl_on_policy_target is not None: -- logits_div_temperature = ( -- logits.bfloat16().div(sampling_info.temperatures).bfloat16() -- ) -- logprobs_via_logsoftmax_kernel = torch.log_softmax( -- logits_div_temperature, dim=-1 -- ) -- - # Post process logits - logits.div_(sampling_info.temperatures) -+ if get_global_server_args().rl_on_policy_target is not None: -+ logprobs_via_logsoftmax_kernel = torch.log_softmax(logits, dim=-1) -+ - # For ascend backend, softmax is not needed before sampling - if not get_global_server_args().sampling_backend == "ascend" or ( - return_logprob and not SGLANG_RETURN_ORIGINAL_LOGPROB -diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index 2ecd8542f..89ef8200d 100644 ---- a/python/sglang/srt/managers/io_struct.py -+++ b/python/sglang/srt/managers/io_struct.py -@@ -1292,6 +1292,19 @@ class UpdateWeightsFromIPCReqOutput(BaseReq): - success: bool - message: str - -+@dataclass -+class PostProcessWeightsReqInput(BaseReq): -+ # Whether to restore weights before loading new weights -+ restore_weights_before_load: bool = False -+ # Whether to enable quantization post-processing -+ post_process_quantization: bool = False -+ -+ -+@dataclass -+class PostProcessWeightsReqOutput(BaseReq): -+ success: bool -+ message: str -+ - - @dataclass - class InitWeightsSendGroupForRemoteInstanceReqOutput(BaseReq): -diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index d423e61d7..9156d543c 100644 ---- a/python/sglang/srt/managers/schedule_batch.py -+++ b/python/sglang/srt/managers/schedule_batch.py -@@ -2186,7 +2186,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - def __str__(self): - return ( - f"ScheduleBatch(forward_mode={self.forward_mode.name if self.forward_mode else 'None'}, " -- f"#req={(len(self.reqs))})" -+ f"#req={(len(self.reqs))}), " -+ f"#out_cache_loc={self.out_cache_loc})" - ) - - -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index 92d286897..43bfab691 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -98,6 +98,7 @@ from sglang.srt.managers.io_struct import ( - OpenSessionReqInput, - OpenSessionReqOutput, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReq, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -1060,6 +1061,7 @@ class Scheduler( - ), - (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor), - (UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc), -+ (PostProcessWeightsReqInput, self.post_process_weights), - (GetWeightsByNameReqInput, self.get_weights_by_name), - (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), - (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), -diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -index e40586c24..243e2b0c2 100644 ---- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py -+++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -@@ -10,6 +10,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode - from sglang.srt.environ import envs - from sglang.srt.layers.logits_processor import LogitsProcessorOutput - from sglang.srt.layers.moe.routed_experts_capturer import get_global_experts_capturer -+ - from sglang.srt.managers.io_struct import ( - AbortReq, - BatchEmbeddingOutput, -@@ -1070,7 +1071,7 @@ class SchedulerOutputProcessorMixin: - req.log_time_stats() - - # Send to detokenizer -- if reqs or is_idle_batch: -+ if rids or is_idle_batch: - if self.model_config.is_multimodal_gen: - return - -diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -index 293a84350..c3a618bcc 100644 ---- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py -+++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -@@ -1,6 +1,7 @@ - from __future__ import annotations - - import logging -+import os - import traceback - from typing import TYPE_CHECKING, Tuple - -@@ -12,6 +13,9 @@ from sglang.srt.constants import ( - GPU_MEMORY_TYPE_KV_CACHE, - GPU_MEMORY_TYPE_WEIGHTS, - ) -+from sglang.srt.disaggregation.utils import DisaggregationMode -+from sglang.srt.distributed import get_moe_ep_group, get_moe_tp_group, get_tp_group -+from sglang.srt.layers.dp_attention import get_attention_tp_group - from sglang.srt.managers.io_struct import ( - CheckWeightsReqInput, - CheckWeightsReqOutput, -@@ -21,6 +25,8 @@ from sglang.srt.managers.io_struct import ( - GetWeightsByNameReqOutput, - InitWeightsUpdateGroupReqInput, - InitWeightsUpdateGroupReqOutput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ReleaseMemoryOccupationReqInput, - ReleaseMemoryOccupationReqOutput, - ResumeMemoryOccupationReqInput, -@@ -114,6 +120,11 @@ class SchedulerUpdateWeightsMixin: - torch.distributed.barrier(group=self.tp_cpu_group) - return UpdateWeightsFromIPCReqOutput(success, message) - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Optional post-processing for updated weights (e.g., Marlin conversion).""" -+ success, message = self.tp_worker.post_process_weights(recv_req) -+ return PostProcessWeightsReqOutput(success, message) -+ - def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): - parameter = self.tp_worker.get_weights_by_name(recv_req) - return GetWeightsByNameReqOutput(parameter) -@@ -137,6 +148,13 @@ class SchedulerUpdateWeightsMixin: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) - self.flush_cache() - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.release_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.release_memory_occupation() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.stashed_model_static_state = _export_static_state( - self.tp_worker.model_runner.model -@@ -177,6 +195,13 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_KV_CACHE in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.resume_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.resume_memory_occupation() -+ - return ResumeMemoryOccupationReqOutput() - - def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): -diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -index e5d42bed8..412293b30 100644 ---- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py -+++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -@@ -49,6 +49,8 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqOutput, - LoRAUpdateOutput, - OpenSessionReqInput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ProfileReq, - ProfileReqOutput, - ProfileReqType, -@@ -177,6 +179,9 @@ class TokenizerCommunicatorMixin: - self.update_weights_from_ipc_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -+ self.post_process_weights_communicator = _Communicator( -+ self.send_to_scheduler, server_args.dp_size -+ ) - self.get_weights_by_name_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -@@ -250,6 +255,10 @@ class TokenizerCommunicatorMixin: - UpdateWeightsFromIPCReqOutput, - self.update_weights_from_ipc_communicator.handle_recv, - ), -+ ( -+ PostProcessWeightsReqOutput, -+ self.post_process_weights_communicator.handle_recv, -+ ), - ( - GetWeightsByNameReqOutput, - self.get_weights_by_name_communicator.handle_recv, -@@ -433,6 +442,17 @@ class TokenizerCommunicatorMixin: - - return success, message - -+ async def post_process_weights( -+ self: TokenizerManager, -+ obj: PostProcessWeightsReqInput, -+ request: Optional[fastapi.Request] = None, -+ ) -> Tuple[bool, str]: -+ """Trigger post-processing hooks for weights after loading (e.g., Marlin conversion).""" -+ self.auto_create_handle_loop() -+ async with self.model_update_lock.writer_lock: -+ results = await self.post_process_weights_communicator(obj) -+ return _Communicator.merge_results(results) -+ - async def init_weights_send_group_for_remote_instance( - self, - obj: InitWeightsSendGroupForRemoteInstanceReqInput, -diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py -index 49f63a198..e4cd0ff2b 100644 ---- a/python/sglang/srt/managers/tp_worker.py -+++ b/python/sglang/srt/managers/tp_worker.py -@@ -27,6 +27,7 @@ from sglang.srt.managers.io_struct import ( - InitWeightsSendGroupForRemoteInstanceReqInput, - InitWeightsUpdateGroupReqInput, - LoadLoRAAdapterReqInput, -+ PostProcessWeightsReqInput, - SendWeightsToRemoteInstanceReqInput, - UnloadLoRAAdapterReqInput, - UpdateWeightFromDiskReqInput, -@@ -175,6 +176,11 @@ class BaseTpWorker(ABC): - success, message = self.model_runner.update_weights_from_ipc(recv_req) - return success, message - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Perform optional post-processing on the updated model weights (e.g., Marlin conversion).""" -+ success, message = self.model_runner.post_process_weights(recv_req) -+ return success, message -+ - def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): - parameter = self.model_runner.get_weights_by_name( - recv_req.name, recv_req.truncate_size -diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py -index 65d562a27..b00a20e95 100644 ---- a/python/sglang/srt/mem_cache/memory_pool.py -+++ b/python/sglang/srt/mem_cache/memory_pool.py -@@ -1678,7 +1678,8 @@ class NSATokenToKVPool(MLATokenToKVPool): - with ( - torch.cuda.use_mem_pool(self.custom_mem_pool) - if self.custom_mem_pool -- else nullcontext() -+ else nullcontext(), -+ self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE), - ): - self.index_k_with_scale_buffer = [ - torch.zeros( -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 1d69c0582..d984c2e12 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -558,7 +558,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): - ) - - # Init routed experts capturer -- self.init_routed_experts_capturer() -+ if not self.is_draft_worker: -+ self.init_routed_experts_capturer() - - if self.device == "cuda": - self.init_cublas() -@@ -2224,11 +2225,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): - output.expert_distribution_metrics = recorder_outputs.get("metrics") - - # Copy cached routing experts' buffers back to CPU cache -- get_global_experts_capturer().on_forward_end( -- forward_batch=forward_batch, -- can_run_graph=output.can_run_graph, -- cuda_graph_batch=getattr(self.graph_runner, "bs", None), -- ) -+ if not self.is_draft_worker: -+ # In speculative decoding, num_tokens_per_bs > 1, so we need to pass -+ # the actual number of tokens per dp rank in cuda graph, not batch size. -+ cuda_graph_num_tokens = None -+ if getattr(self.graph_runner, "bs", None): -+ cuda_graph_num_tokens = ( -+ self.graph_runner.bs * self.graph_runner.num_tokens_per_bs -+ ) -+ get_global_experts_capturer().on_forward_end( -+ forward_batch=forward_batch, -+ can_run_graph=output.can_run_graph, -+ cuda_graph_batch=cuda_graph_num_tokens, -+ ) - - if self.eplb_manager is not None: - self.eplb_manager.on_forward_pass_end() -@@ -2436,6 +2445,41 @@ class ModelRunner(ModelRunnerKVCacheMixin): - logger.error(f"IPC weight update failed: {e}") - return False, str(e) - -+ def post_process_weights(self, recv_req): -+ """ -+ Execute post-processing logic for model weights, such as Marlin quantization format conversion. -+ """ -+ from sglang.srt.model_loader.loader import device_loading_context -+ -+ target_device = torch.device("cuda", torch.cuda.current_device()) -+ -+ if recv_req.restore_weights_before_load: -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports restoring weights -+ if quant_method is not None and hasattr( -+ quant_method, "restore_weights_before_loading" -+ ): -+ -+ with device_loading_context(module, target_device): -+ quant_method.restore_weights_before_loading(module) -+ -+ if recv_req.post_process_quantization: -+ # Iterate through all modules to apply specific post-loading processing -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports quantization post-processing -+ if quant_method is not None and hasattr( -+ quant_method, "process_weights_after_loading" -+ ): -+ -+ # Apply the post-processing (e.g., repacking weights for Marlin kernel) -+ with device_loading_context(module, target_device): -+ quant_method.process_weights_after_loading(module) -+ -+ return True, "Success" - - def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]): - params_dict = dict(model.named_parameters()) -diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py -index ed8cc7ada..d44c8aaa0 100644 ---- a/python/sglang/srt/models/deepseek_v2.py -+++ b/python/sglang/srt/models/deepseek_v2.py -@@ -2704,7 +2704,11 @@ class DeepseekV2AttentionMLA(nn.Module): - ): - k = k_nope.new_empty(*k_shape) - concat_mla_k(k=k, k_nope=k_nope, k_rope=k_pe) -- elif _is_cuda: -+ elif _is_cuda and all( -+ # (i.bit_count() == 1) == (is_power_of_two(i)) -+ i.bit_count() == 1 -+ for i in (k_shape[1], k_nope.shape[-1], k_pe.shape[-1]) -+ ): - # fa3 mha support fp8 inputs - if ( - self.current_attention_backend == "fa3" -diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py -index a7dbadec6..c83a41338 100644 ---- a/python/sglang/srt/models/qwen2.py -+++ b/python/sglang/srt/models/qwen2.py -@@ -90,9 +90,6 @@ class Qwen2MLP(nn.Module): - self.act_fn = SiluAndMul() - - def forward(self, x): -- if get_global_server_args().rl_on_policy_target is not None: -- x = x.bfloat16() -- - gate_up, _ = self.gate_up_proj(x) - x = self.act_fn(gate_up) - x, _ = self.down_proj(x) -@@ -279,11 +276,6 @@ class Qwen2Model(nn.Module): - quant_config=quant_config, - enable_tp=not is_dp_attention_enabled(), - prefix=add_prefix("embed_tokens", prefix), -- params_dtype=( -- torch.float32 -- if get_global_server_args().rl_on_policy_target is not None -- else None -- ), - ) - else: - self.embed_tokens = PPMissingLayer() -@@ -306,10 +298,8 @@ class Qwen2Model(nn.Module): - if self.pp_group.is_last_rank: - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py -index 3ad9f6736..0b9c7f499 100644 ---- a/python/sglang/srt/models/qwen2_moe.py -+++ b/python/sglang/srt/models/qwen2_moe.py -@@ -586,7 +586,17 @@ class Qwen2MoeModel(nn.Module): - prefix=add_prefix("layers", prefix), - ) - if self.pp_group.is_last_rank: -- self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.norm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - else: - self.norm = PPMissingLayer(return_tuple=True) - -diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py -index 9220831f6..47a1a4e4c 100644 ---- a/python/sglang/srt/models/qwen3.py -+++ b/python/sglang/srt/models/qwen3.py -@@ -90,8 +90,8 @@ class Qwen3Attention(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -@@ -242,10 +242,8 @@ class Qwen3DecoderLayer(nn.Module): - - norm_kwargs = ( - dict( -- weight_dtype=torch.float32, - cast_x_before_out_mul=True, -- override_orig_dtype=torch.float32, -- fp32_residual=True, -+ fp32_residual=False, - ) - if get_global_server_args().rl_on_policy_target is not None - else {} -diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py -index e11678a9e..e277d46f2 100644 ---- a/python/sglang/srt/models/qwen3_moe.py -+++ b/python/sglang/srt/models/qwen3_moe.py -@@ -22,6 +22,7 @@ import math - from typing import Any, Dict, Iterable, List, Optional, Tuple, TypeVar - - import torch -+import torch.nn.functional as F - from torch import nn - from transformers import PretrainedConfig - -@@ -50,7 +51,7 @@ from sglang.srt.layers.moe import ( - ) - from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class - from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE --from sglang.srt.layers.moe.topk import TopK -+from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK - from sglang.srt.layers.moe.utils import RoutingMethodType - from sglang.srt.layers.quantization.base_config import QuantizationConfig - from sglang.srt.layers.radix_attention import RadixAttention -@@ -229,6 +230,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - use_grouped_topk=False, - layer_id=layer_id, - ) -+ self.top_k = config.num_experts_per_tok - - self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts -@@ -294,7 +296,22 @@ class Qwen3MoeSparseMoeBlock(nn.Module): - - # router_logits: (num_tokens, n_experts) - router_logits, _ = self.gate(hidden_states) -- topk_output = self.topk(hidden_states, router_logits) -+ -+ if get_global_server_args().rl_on_policy_target is not None: -+ routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float) -+ routing_weights, selected_experts = torch.topk( -+ routing_weights, self.top_k, dim=-1 -+ ) -+ routing_weights /= routing_weights.sum(dim=-1, keepdim=True) -+ routing_weights = routing_weights.to(hidden_states.dtype) -+ topk_output = StandardTopKOutput( -+ topk_weights=routing_weights, -+ topk_ids=selected_experts, -+ router_logits=router_logits, -+ ) -+ else: -+ topk_output = self.topk(hidden_states, router_logits) -+ - final_hidden_states = self.experts(hidden_states, topk_output) - if ( - self.tp_size > 1 -@@ -475,13 +492,14 @@ class Qwen3MoeAttention(nn.Module): - ) - self.compatible_with_fused_kv_buffer = ( - False if isinstance(self.rotary_emb, MRotaryEmbedding) else True -- ) -+ ) and (get_global_server_args().rl_on_policy_target is None) - self.compatible_with_fused_qk_norm_rope = ( - not isinstance(self.rotary_emb, MRotaryEmbedding) - ) and self.head_dim in (64, 128, 256) - self.use_fused_qk_norm_rope = ( - get_global_server_args().enable_fused_qk_norm_rope - and self.compatible_with_fused_qk_norm_rope -+ and (get_global_server_args().rl_on_policy_target is None) - ) - self._used_fused_qk_norm_rope_last_call = False - -@@ -494,8 +512,16 @@ class Qwen3MoeAttention(nn.Module): - prefix=add_prefix("attn", prefix), - ) - -- self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -- self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) -+ self.k_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) - self.alt_stream = alt_stream - - def op_prepare(self, state): -@@ -736,9 +762,19 @@ class Qwen3MoeDecoderLayer(nn.Module): - quant_config=quant_config, - prefix=add_prefix("mlp", prefix), - ) -- self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) -+ norm_kwargs = ( -+ dict( -+ cast_x_before_out_mul=True, -+ fp32_residual=False, -+ ) -+ if get_global_server_args().rl_on_policy_target is not None -+ else {} -+ ) -+ self.input_layernorm = RMSNorm( -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs -+ ) - self.post_attention_layernorm = RMSNorm( -- config.hidden_size, eps=config.rms_norm_eps -+ config.hidden_size, eps=config.rms_norm_eps, **norm_kwargs - ) - - self.layer_communicator = LayerCommunicator( -diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py -index 079f45843..218e32362 100644 ---- a/python/sglang/srt/models/qwen3_vl.py -+++ b/python/sglang/srt/models/qwen3_vl.py -@@ -397,28 +397,68 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): - return cos_combined, sin_combined - - def fast_pos_embed_interpolate(self, grid_thw): -- patch_pos_embeds_permute = [] -- m_size = self.spatial_merge_size -+ grid_ts, grid_hs, grid_ws = grid_thw[:, 0], grid_thw[:, 1], grid_thw[:, 2] -+ num_grid_per_side = int(self.num_position_embeddings**0.5) -+ device = self.pos_embed.weight.device -+ -+ idx_list = [[] for _ in range(4)] -+ weight_list = [[] for _ in range(4)] -+ -+ for t, h, w in zip(grid_ts, grid_hs, grid_ws): -+ h_idxs = torch.linspace(0, num_grid_per_side - 1, h) -+ w_idxs = torch.linspace(0, num_grid_per_side - 1, w) -+ -+ h_idxs_floor = h_idxs.int() -+ w_idxs_floor = w_idxs.int() -+ h_idxs_ceil = (h_idxs.int() + 1).clip(max=num_grid_per_side - 1) -+ w_idxs_ceil = (w_idxs.int() + 1).clip(max=num_grid_per_side - 1) -+ -+ dh = h_idxs - h_idxs_floor -+ dw = w_idxs - w_idxs_floor -+ -+ base_h = h_idxs_floor * num_grid_per_side -+ base_h_ceil = h_idxs_ceil * num_grid_per_side -+ -+ indices = [ -+ (base_h[None].T + w_idxs_floor[None]).flatten(), -+ (base_h[None].T + w_idxs_ceil[None]).flatten(), -+ (base_h_ceil[None].T + w_idxs_floor[None]).flatten(), -+ (base_h_ceil[None].T + w_idxs_ceil[None]).flatten(), -+ ] -+ -+ weights = [ -+ ((1 - dh)[None].T * (1 - dw)[None]).flatten(), -+ ((1 - dh)[None].T * dw[None]).flatten(), -+ (dh[None].T * (1 - dw)[None]).flatten(), -+ (dh[None].T * dw[None]).flatten(), -+ ] - -- embeds = torch.arange(self.num_grid, device=self.pos_embed.weight.device) -- embeds = ( -- self.pos_embed(embeds) -- .permute(1, 0) -- .reshape(1, -1, self.num_grid_per_side, self.num_grid_per_side) -+ for i in range(4): -+ idx_list[i].extend(indices[i].tolist()) -+ weight_list[i].extend(weights[i].tolist()) -+ -+ idx_tensor = torch.tensor(idx_list, dtype=torch.long, device=device) -+ weight_tensor = torch.tensor( -+ weight_list, dtype=self.pos_embed.weight.dtype, device=device - ) -- for t, h, w in grid_thw: -- pos_embed = torch.nn.functional.interpolate( -- embeds, size=(h, w), mode="bilinear", align_corners=self.align_corners -- ) -- pos_embed = pos_embed.reshape( -- -1, -- h // self.spatial_merge_size, -- self.spatial_merge_size, -- w // self.spatial_merge_size, -- self.spatial_merge_size, -+ pos_embeds = self.pos_embed(idx_tensor).to(device) * weight_tensor[:, :, None] -+ patch_pos_embeds = pos_embeds[0] + pos_embeds[1] + pos_embeds[2] + pos_embeds[3] -+ -+ patch_pos_embeds = patch_pos_embeds.split( -+ [h * w for h, w in zip(grid_hs, grid_ws)] -+ ) -+ -+ patch_pos_embeds_permute = [] -+ merge_size = self.spatial_merge_size -+ for pos_embed, t, h, w in zip(patch_pos_embeds, grid_ts, grid_hs, grid_ws): -+ pos_embed = pos_embed.repeat(t, 1) -+ pos_embed = ( -+ pos_embed.view( -+ t, h // merge_size, merge_size, w // merge_size, merge_size, -1 -+ ) -+ .permute(0, 1, 3, 2, 4, 5) -+ .flatten(0, 4) - ) -- pos_embed = pos_embed.permute(1, 3, 2, 4, 0) -- pos_embed = pos_embed.flatten(0, 3).repeat(t, 1) - patch_pos_embeds_permute.append(pos_embed) - return torch.cat(patch_pos_embeds_permute) - -@@ -610,14 +650,19 @@ class Qwen3LLMModel(Qwen3Model): - hidden_states + residual if residual is not None else hidden_states - ) - -+ deepstack_embeds = None -+ if input_deepstack_embeds is not None: -+ prev_layer_idx = layer_idx - 1 -+ if prev_layer_idx in self.deepstack_embed_to_decoder_layer: -+ sep = self.hidden_size * prev_layer_idx -+ deepstack_embeds = input_deepstack_embeds[ -+ :, sep : sep + self.hidden_size -+ ] -+ - # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. - # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 - # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack - # The order matters because addition with different tensors is not associative in practice. -- # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. -- deepstack_embeds = self.get_deepstack_embeds( -- layer_idx - 1, input_deepstack_embeds -- ) - hidden_states, residual = layer( - positions, - hidden_states, -diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index a2b26e0e0..72db29801 100644 ---- a/python/sglang/srt/server_args.py -+++ b/python/sglang/srt/server_args.py -@@ -527,6 +527,7 @@ class ServerArgs: - cuda_graph_max_bs: Optional[int] = None - cuda_graph_bs: Optional[List[int]] = None - disable_cuda_graph: bool = False -+ disable_draft_cuda_graph: bool = False - disable_cuda_graph_padding: bool = False - enable_profile_cuda_graph: bool = False - enable_cudagraph_gc: bool = False -@@ -3980,6 +3981,11 @@ class ServerArgs: - action="store_true", - help="Disable cuda graph.", - ) -+ parser.add_argument( -+ "--disable-draft-cuda-graph", -+ action="store_true", -+ help="Disable cuda graph for draft model in speculative decoding.", -+ ) - parser.add_argument( - "--disable-cuda-graph-padding", - action="store_true", -diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -index 5fe45086c..c95fbd0f6 100644 ---- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -+++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -@@ -341,7 +341,10 @@ class EAGLEDraftCudaGraphRunner: - self.seq_lens.fill_(self.seq_len_fill_value) - self.out_cache_loc.zero_() - self.positions.zero_() -- -+ self.topk_p.zero_() -+ self.topk_index.zero_() -+ self.hidden_states.zero_() -+ self.req_pool_indices.zero_() - num_tokens = bs * self.num_tokens_per_bs - - # Common inputs -@@ -350,8 +353,8 @@ class EAGLEDraftCudaGraphRunner: - forward_batch.out_cache_loc - ) - self.positions[:raw_num_token].copy_(forward_batch.positions) -- self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) -- self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) -+ self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p.clamp(0, 1)) -+ self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index.clamp(0, self.model_runner.model_config.vocab_size - 1)) - self.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) - self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) - -diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index 1bf3816e9..b5b41dba4 100644 ---- a/python/sglang/srt/speculative/eagle_info.py -+++ b/python/sglang/srt/speculative/eagle_info.py -@@ -778,6 +778,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] - self.verified_id = self.verified_id[: len(new_indices)] -+ if self.accept_length is not None: -+ self.accept_length = self.accept_length[: len(new_indices)] -+ if self.accept_length_cpu is not None: -+ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] - else: - # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` - self.topk_p = self.topk_p[new_indices] -@@ -809,6 +813,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) - self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) - self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) -+ if self.accept_length is not None and spec_info.accept_length is not None: -+ self.accept_length = torch.cat( -+ [self.accept_length, spec_info.accept_length] -+ ) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif self.accept_length is not None: -+ zeros = torch.zeros( -+ [spec_info.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([self.accept_length, zeros]) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif spec_info.accept_length is not None: -+ zeros = torch.zeros( -+ [self.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([zeros, spec_info.accept_length]) -+ self.accept_length_cpu = self.accept_length.tolist() - - - @dataclass -diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py -index a702df4f8..61d9ae366 100644 ---- a/python/sglang/srt/speculative/eagle_worker.py -+++ b/python/sglang/srt/speculative/eagle_worker.py -@@ -231,7 +231,7 @@ class EAGLEWorker(TpModelWorker): - self.cuda_graph_runner = None - self.cuda_graph_runner_for_draft_extend = None - -- if self.server_args.disable_cuda_graph: -+ if self.server_args.disable_cuda_graph or self.server_args.disable_draft_cuda_graph: - return - - Device2DraftCudaGraphRunner = { -diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py -index 8560246c6..13db860dc 100644 ---- a/python/sglang/srt/utils/common.py -+++ b/python/sglang/srt/utils/common.py -@@ -2224,6 +2224,8 @@ class SafeUnpickler(pickle.Unpickler): - "sglang.srt.model_executor.model_runner.", - "sglang.srt.layers.", - "sglang.srt.utils.", -+ # --- slime --- -+ "slime.", - } - - DENY_CLASSES = { diff --git a/docker/patch/v0.5.9/sglang.patch b/docker/patch/v0.5.9/sglang.patch deleted file mode 100644 index e9145702e..000000000 --- a/docker/patch/v0.5.9/sglang.patch +++ /dev/null @@ -1,3216 +0,0 @@ -diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py -index 6fbd1db823..f80ec11bb4 100644 ---- a/python/sglang/srt/configs/model_config.py -+++ b/python/sglang/srt/configs/model_config.py -@@ -274,6 +274,7 @@ class ModelConfig: - - if is_draft_model and self.hf_config.architectures[0] in [ - "DeepseekV3ForCausalLM", -+ "DeepseekV32ForCausalLM", - "GlmMoeDsaForCausalLM", - ]: - self.hf_config.architectures[0] = "DeepseekV3ForCausalLMNextN" -diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py -index da4629e525..c03f98231a 100644 ---- a/python/sglang/srt/disaggregation/base/conn.py -+++ b/python/sglang/srt/disaggregation/base/conn.py -@@ -17,6 +17,7 @@ class KVArgs: - kv_data_ptrs: List[int] - kv_data_lens: List[int] - kv_item_lens: List[int] -+ aux_buffer_names: List[str] - aux_data_ptrs: List[int] - aux_data_lens: List[int] - aux_item_lens: List[int] -diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py -index 67fe82ad67..ed5fa7b0e3 100644 ---- a/python/sglang/srt/disaggregation/common/conn.py -+++ b/python/sglang/srt/disaggregation/common/conn.py -@@ -333,6 +333,10 @@ class CommonKVReceiver(BaseKVReceiver): - self.required_dst_info_num = ( - self.kv_mgr.attn_tp_size // self.prefill_attn_tp_size - ) -+ # With attention DP, one request is routed to one decode rank. -+ # Waiting for all TP shards to pre-allocate the same bootstrap room would stall forever. -+ if self.kv_mgr.attn_dp_size > 1: -+ self.required_dst_info_num = 1 - self.required_prefill_response_num = 1 * ( - self.prefill_pp_size // self.kv_mgr.pp_size - ) -@@ -422,6 +426,7 @@ class CommonKVReceiver(BaseKVReceiver): - f"Could not fetch bootstrap info for engine rank: {self.kv_mgr.kv_args.engine_rank} and target_dp_group: {self.target_dp_group} and target_pp_rank {target_pp_rank}", - ) - self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed) -+ self.bootstrap_infos = None - return - - self.bootstrap_infos = bootstrap_infos -@@ -610,8 +615,12 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer): - and int(target_dp_group) == -1 - and int(target_pp_rank) == -1 - ): -+ inferred_attn_tp_size = max( -+ (len(v) for v in self.prefill_port_table.values()), -+ default=self.attn_tp_size, -+ ) - prefill_parallel_info = { -- "prefill_attn_tp_size": self.attn_tp_size, -+ "prefill_attn_tp_size": inferred_attn_tp_size, - "prefill_dp_size": self.dp_size, - "prefill_pp_size": self.pp_size, - "prefill_page_size": self.page_size, -diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py -index 1d8baf0028..1ebb959298 100644 ---- a/python/sglang/srt/disaggregation/decode.py -+++ b/python/sglang/srt/disaggregation/decode.py -@@ -21,6 +21,7 @@ Life cycle of a request in the decode server - from __future__ import annotations - - import logging -+import os - import time - from collections import deque - from dataclasses import dataclass -@@ -40,8 +41,10 @@ from sglang.srt.disaggregation.utils import ( - MetadataBuffers, - ReqToMetadataIdxAllocator, - TransferBackend, -+ apply_prefill_timing_payload, - get_kv_class, - is_mla_backend, -+ is_slime_profiling_enabled, - kv_to_page_indices, - poll_and_all_reduce, - prepare_abort, -@@ -295,6 +298,7 @@ class DecodePreallocQueue: - kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = ( - self.metadata_buffers.get_buf_infos() - ) -+ kv_args.aux_buffer_names = self.metadata_buffers.get_aux_buffer_names() - - if hasattr(self.token_to_kv_pool, "get_state_buf_infos"): - state_data_ptrs, state_data_lens, state_item_lens = ( -@@ -336,6 +340,16 @@ class DecodePreallocQueue: - ) - return kv_manager - -+ def release_memory_occupation(self): -+ self.queue.clear() -+ self.retracted_queue.clear() -+ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): -+ self.kv_manager.deregister_buffer_to_engine() -+ -+ def resume_memory_occupation(self): -+ if hasattr(self.kv_manager, "register_buffer_to_engine"): -+ self.kv_manager.register_buffer_to_engine() -+ - def add(self, req: Req, is_retracted: bool = False) -> None: - """Add a request to the pending queue.""" - if self._check_if_req_exceed_kv_capacity(req): -@@ -440,12 +454,37 @@ class DecodePreallocQueue: - [decode_req.kv_receiver for decode_req in self.queue], self.gloo_group - ) - -+ # Bootstrap timeout: if a request has been stuck in Bootstrapping for too long, treat it as failed. -+ bootstrap_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): - if rids_to_check is not None and decode_req.req.rid not in rids_to_check: - continue - - if poll == KVPoll.Bootstrapping: -- pass -+ # Check for bootstrap timeout -+ entry_time = getattr( -+ decode_req.req.time_stats, -+ "decode_prealloc_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > bootstrap_timeout: -+ error_message = ( -+ f"Decode bootstrap timed out after {now - entry_time:.1f}s " -+ f"for request rank={self.tp_rank} " -+ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ prepare_abort( -+ decode_req.req, -+ error_message, -+ status_code=HTTPStatus.GATEWAY_TIMEOUT, -+ ) -+ if self.scheduler.enable_metrics: -+ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() - elif poll == KVPoll.WaitingForInput: - decode_req.waiting_for_input = True - elif poll == KVPoll.Failed: -@@ -590,6 +629,7 @@ class DecodePreallocQueue: - self.req_to_metadata_buffer_idx_allocator.alloc() - ) - assert decode_req.metadata_buffer_index is not None -+ self.metadata_buffers.clear_profiling_buf(decode_req.metadata_buffer_index) - page_indices = kv_to_page_indices(kv_indices, page_size) - decode_req.kv_receiver.init( - page_indices, decode_req.metadata_buffer_index, state_indices -@@ -751,6 +791,7 @@ class DecodeTransferQueue: - output_topk_index, - output_hidden_states, - output_bootstrap_room, -+ output_prefill_timing, - ) = self.metadata_buffers.get_buf(idx) - - # Validate bootstrap_room to detect context corruption -@@ -813,6 +854,14 @@ class DecodeTransferQueue: - output_top_logprobs_idx[: decode_req.req.top_logprobs_num].tolist() - ) - -+ # Inject prefill-side PD timing forwarded from the P instance. -+ # Layout: [bootstrap_queue, forward, transfer_queue, bootstrap, -+ # alloc_waiting, transfer_speed, transfer_mb, retry_count] -+ if is_slime_profiling_enabled(): -+ apply_prefill_timing_payload( -+ decode_req.req.time_stats, output_prefill_timing -+ ) -+ - decode_req.kv_receiver.clear() - decode_req.kv_receiver = None - trace_slice_end( -@@ -830,6 +879,13 @@ class DecodeTransferQueue: - [decode_req.kv_receiver for decode_req in self.queue], self.gloo_group - ) - -+ # Transfer timeout: if a request has been in the transfer queue for too long -+ # (e.g., stuck in Bootstrapping/WaitingForInput/Transferring), treat it as failed. -+ transfer_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - transferred_reqs = [] - indices_to_remove = set() - for i, (decode_req, poll) in enumerate(zip(self.queue, polls)): -@@ -877,7 +933,20 @@ class DecodeTransferQueue: - KVPoll.WaitingForInput, - KVPoll.Transferring, - ]: -- pass -+ # Check for transfer timeout -+ entry_time = getattr( -+ decode_req.req.time_stats, -+ "decode_transfer_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > transfer_timeout: -+ error_message = ( -+ f"Decode transfer timed out after {now - entry_time:.1f}s " -+ f"(state={poll}) for request rank={self.tp_rank} " -+ f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ decode_req.kv_receiver.abort() - else: - raise ValueError(f"Unexpected poll case: {poll}") - -@@ -893,6 +962,14 @@ class DecodeTransferQueue: - - return transferred_reqs - -+ def release_memory_occupation(self): -+ """Clean up all in-flight transfers before releasing GPU memory.""" -+ self.queue.clear() -+ -+ def resume_memory_occupation(self): -+ """Resume after GPU memory re-allocation. Queue was already cleared on release.""" -+ pass -+ - - class SchedulerDisaggregationDecodeMixin: - -@@ -1072,7 +1149,15 @@ class SchedulerDisaggregationDecodeMixin: - resumed_reqs = self.disagg_decode_prealloc_queue.resume_retracted_reqs() - self.waiting_queue.extend(resumed_reqs) - if len(self.disagg_decode_prealloc_queue.retracted_queue) > 0: -- # if there are still retracted requests, we do not allocate new requests -+ # Still have retracted requests that couldn't resume (not enough memory). -+ # Don't accept new requests (pop_preallocated) — they would consume memory -+ # that retracted requests need. -+ # But DO drain completed transfers: their KV is already committed, and -+ # moving them to waiting_queue frees the reserved-decode-token budget -+ # in _allocatable_tokens(), which may unblock resume on the next iteration. -+ # Without this, completed transfers hold memory indefinitely → deadlock. -+ alloc_reqs = self.disagg_decode_transfer_queue.pop_transferred() -+ self.waiting_queue.extend(alloc_reqs) - return - - if not hasattr(self, "polling_count"): -diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py -index d0d4efd958..b3a207063e 100644 ---- a/python/sglang/srt/disaggregation/mooncake/conn.py -+++ b/python/sglang/srt/disaggregation/mooncake/conn.py -@@ -30,7 +30,7 @@ from sglang.srt.disaggregation.common.utils import ( - from sglang.srt.disaggregation.mooncake.utils import ( - check_mooncake_custom_mem_pool_enabled, - ) --from sglang.srt.disaggregation.utils import DisaggregationMode -+from sglang.srt.disaggregation.utils import DisaggregationMode, iter_aux_transfer_specs - from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine - from sglang.srt.environ import envs - from sglang.srt.server_args import ServerArgs -@@ -260,6 +260,19 @@ class MooncakeKVManager(CommonKVManager): - self.kv_args.state_data_ptrs, self.kv_args.state_data_lens - ) - -+ def deregister_buffer_to_engine(self): -+ # Batch deregister KV data buffers -+ if self.kv_args.kv_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.kv_data_ptrs) -+ -+ # Batch deregister auxiliary data buffers -+ if self.kv_args.aux_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.aux_data_ptrs) -+ -+ # Batch deregister state/extra pool data buffers -+ if self.kv_args.state_data_ptrs: -+ self.engine.batch_deregister(self.kv_args.state_data_ptrs) -+ - def _transfer_data(self, mooncake_session_id, transfer_blocks): - if not transfer_blocks: - return 0 -@@ -524,10 +537,14 @@ class MooncakeKVManager(CommonKVManager): - prefill_aux_ptrs = self.kv_args.aux_data_ptrs - prefill_aux_item_lens = self.kv_args.aux_item_lens - -- for i, dst_aux_ptr in enumerate(dst_aux_ptrs): -- length = prefill_aux_item_lens[i] -- src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -- dst_addr = dst_aux_ptrs[i] + length * req.dst_aux_index -+ for _, src_addr, dst_addr, length in iter_aux_transfer_specs( -+ self.kv_args.aux_buffer_names, -+ prefill_aux_ptrs, -+ prefill_aux_item_lens, -+ dst_aux_ptrs, -+ prefill_aux_index, -+ req.dst_aux_index, -+ ): - transfer_blocks.append((src_addr, dst_addr, length)) - - return self._transfer_data(req.mooncake_session_id, transfer_blocks) -@@ -541,9 +558,14 @@ class MooncakeKVManager(CommonKVManager): - prefill_aux_ptrs = self.kv_args.aux_data_ptrs - prefill_aux_item_lens = self.kv_args.aux_item_lens - -- for i in range(len(prefill_aux_ptrs)): -- length = prefill_aux_item_lens[i] -- src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -+ for i, src_addr, _, length in iter_aux_transfer_specs( -+ self.kv_args.aux_buffer_names, -+ prefill_aux_ptrs, -+ prefill_aux_item_lens, -+ dst_aux_ptrs, -+ prefill_aux_index, -+ req.dst_aux_index, -+ ): - data = AuxDataCodec.serialize_data_from_buffer(src_addr, length) - - self.send_aux_data_to_endpoint( -@@ -643,13 +665,13 @@ class MooncakeKVManager(CommonKVManager): - raise RuntimeError( - f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." - ) -- if len(prefill_state_indices) < len(req.dst_state_indices): -- logger.warning( -- f"len(prefill_state_indices) = {len(prefill_state_indices)}, len(dst_state_indices) = {len(req.dst_state_indices)}" -+ if len(prefill_state_indices) != len(req.dst_state_indices): -+ logger.error( -+ "PD extra-state index mismatch, reject transfer to avoid corrupted outputs: " -+ f"len(prefill_state_indices)={len(prefill_state_indices)}, " -+ f"len(dst_state_indices)={len(req.dst_state_indices)}" - ) -- prefill_state_indices = prefill_state_indices[ -- : len(req.dst_state_indices) -- ] -+ return -1 - # Reuse _send_kvcache_generic interface to send extra pool data - prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) - dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32) -@@ -858,12 +880,6 @@ class MooncakeKVManager(CommonKVManager): - if ret != 0: - with self.session_lock: - self.session_failures[req.mooncake_session_id] += 1 -- # Failures should never happen if the session is not dead, if the session fails once, mark it as failed -- if self.session_failures[req.mooncake_session_id] >= 1: -- self.failed_sessions.add(req.mooncake_session_id) -- logger.error( -- f"Session {req.mooncake_session_id} failed." -- ) - self.record_failure( - kv_chunk.room, - f"Failed to send kv chunk of {kv_chunk.room} to {req.endpoint}:{req.dst_port}", -@@ -880,13 +896,31 @@ class MooncakeKVManager(CommonKVManager): - - if kv_chunk.is_last: - if kv_chunk.state_indices is not None: -- self.maybe_send_extra( -+ ret = self.maybe_send_extra( - req, - kv_chunk.state_indices, - target_rank_registration_info.dst_state_data_ptrs, - executor, - target_rank_registration_info, - ) -+ if ret != 0: -+ with self.session_lock: -+ self.session_failures[ -+ req.mooncake_session_id -+ ] += 1 -+ self.record_failure( -+ kv_chunk.room, -+ f"Failed to send extra state chunk of {kv_chunk.room} to {req.endpoint}:{req.dst_port}", -+ ) -+ self.update_status(kv_chunk.room, KVPoll.Failed) -+ self.sync_status_to_decode_endpoint( -+ req.endpoint, -+ req.dst_port, -+ req.room, -+ KVPoll.Failed, -+ local_rank, -+ ) -+ break - - # Only the last chunk we need to send the aux data - ret = self.send_aux( -@@ -895,6 +929,11 @@ class MooncakeKVManager(CommonKVManager): - target_rank_registration_info.dst_aux_ptrs, - ) - polls.append(True if ret == 0 else False) -+ if ret != 0: -+ # Mark session as failed to avoid hanging -+ # on subsequent batch_transfer_sync calls -+ with self.session_lock: -+ self.session_failures[req.mooncake_session_id] += 1 - dst_ranks_infos.append( - (req.endpoint, req.dst_port, req.room) - ) -@@ -977,15 +1016,20 @@ class MooncakeKVManager(CommonKVManager): - - if status == KVPoll.Success: - if bootstrap_room in self.request_status: -- self.prefill_response_tracker[bootstrap_room].add(prefill_rank) -+ # Guard against TOCTOU race: clear() may remove the entry -+ # between the request_status check and dict access here. - expected_response_num = ( -- self.required_prefill_response_num_table[bootstrap_room] -+ self.required_prefill_response_num_table.get(bootstrap_room) - ) -- arrived_response_num = len( -- self.prefill_response_tracker[bootstrap_room] -- ) -- if arrived_response_num == expected_response_num: -- self.update_status(bootstrap_room, KVPoll.Success) -+ if expected_response_num is not None: -+ self.prefill_response_tracker[bootstrap_room].add( -+ prefill_rank -+ ) -+ arrived_response_num = len( -+ self.prefill_response_tracker[bootstrap_room] -+ ) -+ if arrived_response_num == expected_response_num: -+ self.update_status(bootstrap_room, KVPoll.Success) - elif status == KVPoll.Failed: - self.record_failure( - bootstrap_room, -@@ -1266,7 +1310,10 @@ class MooncakeKVReceiver(CommonKVReceiver): - super().__init__(mgr, bootstrap_addr, bootstrap_room, prefill_dp_rank) - - self.kv_mgr.addr_to_rooms_tracker[self.bootstrap_addr].add(self.bootstrap_room) -- self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput) -+ # Only transition to WaitingForInput if bootstrap succeeded; -+ # if super().__init__() set status to Failed, do not override it. -+ if self.bootstrap_infos is not None: -+ self.kv_mgr.update_status(self.bootstrap_room, KVPoll.WaitingForInput) - - def _register_kv_args(self): - for bootstrap_info in self.bootstrap_infos: -diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py -index fbc8016351..7de53ba292 100644 ---- a/python/sglang/srt/disaggregation/prefill.py -+++ b/python/sglang/srt/disaggregation/prefill.py -@@ -20,6 +20,7 @@ Life cycle of a request in the prefill server - from __future__ import annotations - - import logging -+import os - import time - from collections import deque - from http import HTTPStatus -@@ -167,6 +168,7 @@ class PrefillBootstrapQueue: - kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = ( - self.metadata_buffers.get_buf_infos() - ) -+ kv_args.aux_buffer_names = self.metadata_buffers.get_aux_buffer_names() - kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device - kv_args.gpu_id = self.scheduler.gpu_id - -@@ -276,6 +278,12 @@ class PrefillBootstrapQueue: - [req.disagg_kv_sender for req in self.queue], self.gloo_group - ) - -+ # Bootstrap timeout: if a request has been stuck in Bootstrapping for too long, treat it as failed. -+ bootstrap_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - for i, (req, poll) in enumerate(zip(self.queue, polls)): - if rids_to_check is not None: - # if req not in reqs_info_to_check, skip -@@ -283,6 +291,27 @@ class PrefillBootstrapQueue: - continue - - if poll == KVPoll.Bootstrapping: -+ # Check for bootstrap timeout -+ entry_time = getattr( -+ req.time_stats, -+ "prefill_bootstrap_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > bootstrap_timeout: -+ error_message = ( -+ f"Prefill bootstrap timed out after {now - entry_time:.1f}s " -+ f"for request rank={self.tp_rank} " -+ f"{req.rid=} {req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ prepare_abort( -+ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT -+ ) -+ self.scheduler.stream_output([req], req.return_logprob) -+ indices_to_remove.add(i) -+ failed_reqs.append(req) -+ if self.scheduler.enable_metrics: -+ self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() - continue - elif poll == KVPoll.Failed: - error_message = f"Prefill bootstrap failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}" -@@ -335,6 +364,15 @@ class PrefillBootstrapQueue: - else: - return bootstrapped_reqs, failed_reqs - -+ def release_memory_occupation(self): -+ self.queue.clear() -+ if hasattr(self.kv_manager, "deregister_buffer_to_engine"): -+ self.kv_manager.deregister_buffer_to_engine() -+ -+ def resume_memory_occupation(self): -+ if hasattr(self.kv_manager, "register_buffer_to_engine"): -+ self.kv_manager.register_buffer_to_engine() -+ - - class SchedulerDisaggregationPrefillMixin: - """ -@@ -547,6 +585,18 @@ class SchedulerDisaggregationPrefillMixin: - - self.maybe_send_health_check_signal() - -+ if ( -+ self.current_scheduler_metrics_enabled -+ and hasattr(batch, "prefill_stats") -+ and batch.prefill_stats is not None -+ ): -+ can_run_cuda_graph = getattr(result, "can_run_cuda_graph", False) -+ self.log_prefill_stats( -+ prefill_stats=batch.prefill_stats, -+ can_run_cuda_graph=can_run_cuda_graph, -+ dp_cooperation_info=getattr(batch, "dp_cooperation_info", None), -+ ) -+ - def process_disagg_prefill_inflight_queue( - self: Scheduler, rids_to_check: Optional[List[str]] = None - ) -> List[Req]: -@@ -564,6 +614,13 @@ class SchedulerDisaggregationPrefillMixin: - self.attn_tp_cpu_group, - ) - -+ # Transfer timeout: if a request has been in the inflight queue for too long -+ # (e.g., stuck in WaitingForInput/Transferring), treat it as failed. -+ transfer_timeout = float( -+ os.environ.get("SGLANG_DISAGGREGATION_TRANSFER_TIMEOUT", "600") -+ ) -+ now = time.perf_counter() -+ - undone_reqs: List[Req] = [] - # Check .poll() for the reqs in disagg_prefill_inflight_queue. If Success, respond to the client and remove it from the queue - for req, poll in zip(self.disagg_prefill_inflight_queue, polls): -@@ -573,10 +630,35 @@ class SchedulerDisaggregationPrefillMixin: - undone_reqs.append(req) - continue - -- assert poll == KVPoll.Success or poll == KVPoll.Failed -+ if poll not in (KVPoll.Success, KVPoll.Failed): -+ undone_reqs.append(req) -+ continue - - if poll in [KVPoll.WaitingForInput, KVPoll.Transferring]: -- undone_reqs.append(req) -+ # Check for transfer timeout -+ entry_time = getattr( -+ req.time_stats, -+ "prefill_transfer_queue_entry_time", -+ None, -+ ) -+ if entry_time is not None and (now - entry_time) > transfer_timeout: -+ error_message = ( -+ f"Prefill transfer timed out after {now - entry_time:.1f}s " -+ f"(state={poll}) for request rank={self.tp_rank} " -+ f"{req.rid=} {req.bootstrap_room=}" -+ ) -+ logger.error(error_message) -+ release_kv_cache(req, self.tree_cache) # unlock the tree -+ prepare_abort( -+ req, error_message, status_code=HTTPStatus.GATEWAY_TIMEOUT -+ ) -+ if hasattr(req.disagg_kv_sender, "clear"): -+ req.disagg_kv_sender.clear() -+ done_reqs.append(req) -+ if self.enable_metrics: -+ self.metrics_collector.increment_transfer_failed_reqs() -+ else: -+ undone_reqs.append(req) - elif poll == KVPoll.Success: # transfer done - release_kv_cache(req, self.tree_cache) # unlock the tree - req.finished_reason = FINISH_LENGTH(length=0) -diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py -index 6d58f415a7..84723c342c 100644 ---- a/python/sglang/srt/disaggregation/utils.py -+++ b/python/sglang/srt/disaggregation/utils.py -@@ -21,6 +21,17 @@ if TYPE_CHECKING: - # Constants & Enums - ######################### - FAKE_BOOTSTRAP_HOST = "2.2.2.2" -+PREFILL_TIMING_AUX_BUFFER_NAME = "prefill_timing" -+PREFILL_TIMING_DEST_ATTRS = ( -+ ("fwd_prefill_bootstrap_queue_duration", float), -+ ("fwd_prefill_forward_duration", float), -+ ("fwd_prefill_transfer_queue_duration", float), -+ ("fwd_bootstrap_duration", float), -+ ("fwd_alloc_waiting_duration", float), -+ ("fwd_transfer_speed_gb_s", float), -+ ("fwd_transfer_total_mb", float), -+ ("fwd_prefill_retry_count", int), -+) - - - class DisaggregationMode(Enum): -@@ -139,46 +150,35 @@ class MetadataBuffers: - self.bootstrap_room = torch.zeros( - (size, 8), dtype=torch.uint64, device=device - ) -+ # Prefill-side PD timing (8 floats, padded to 16 for RDMA alignment). -+ # Layout: [bootstrap_queue, forward, transfer_queue, bootstrap, -+ # alloc_waiting, transfer_speed, transfer_mb, retry_count] -+ self.prefill_timing = torch.zeros( -+ (size, 16), dtype=torch.float32, device=device -+ ) -+ self.aux_buffers = [ -+ ("output_ids", self.output_ids), -+ ("cached_tokens", self.cached_tokens), -+ ("output_token_logprobs_val", self.output_token_logprobs_val), -+ ("output_token_logprobs_idx", self.output_token_logprobs_idx), -+ ("output_top_logprobs_val", self.output_top_logprobs_val), -+ ("output_top_logprobs_idx", self.output_top_logprobs_idx), -+ ("output_topk_p", self.output_topk_p), -+ ("output_topk_index", self.output_topk_index), -+ ("output_hidden_states", self.output_hidden_states), -+ ("bootstrap_room", self.bootstrap_room), -+ (PREFILL_TIMING_AUX_BUFFER_NAME, self.prefill_timing), -+ ] - - def get_buf_infos(self): -- ptrs = [ -- self.output_ids.data_ptr(), -- self.cached_tokens.data_ptr(), -- self.output_token_logprobs_val.data_ptr(), -- self.output_token_logprobs_idx.data_ptr(), -- self.output_top_logprobs_val.data_ptr(), -- self.output_top_logprobs_idx.data_ptr(), -- self.output_topk_p.data_ptr(), -- self.output_topk_index.data_ptr(), -- self.output_hidden_states.data_ptr(), -- self.bootstrap_room.data_ptr(), -- ] -- data_lens = [ -- self.output_ids.nbytes, -- self.cached_tokens.nbytes, -- self.output_token_logprobs_val.nbytes, -- self.output_token_logprobs_idx.nbytes, -- self.output_top_logprobs_val.nbytes, -- self.output_top_logprobs_idx.nbytes, -- self.output_topk_p.nbytes, -- self.output_topk_index.nbytes, -- self.output_hidden_states.nbytes, -- self.bootstrap_room.nbytes, -- ] -- item_lens = [ -- self.output_ids[0].nbytes, -- self.cached_tokens[0].nbytes, -- self.output_token_logprobs_val[0].nbytes, -- self.output_token_logprobs_idx[0].nbytes, -- self.output_top_logprobs_val[0].nbytes, -- self.output_top_logprobs_idx[0].nbytes, -- self.output_topk_p[0].nbytes, -- self.output_topk_index[0].nbytes, -- self.output_hidden_states[0].nbytes, -- self.bootstrap_room[0].nbytes, -- ] -+ ptrs = [buffer.data_ptr() for _, buffer in self.aux_buffers] -+ data_lens = [buffer.nbytes for _, buffer in self.aux_buffers] -+ item_lens = [buffer[0].nbytes for _, buffer in self.aux_buffers] - return ptrs, data_lens, item_lens - -+ def get_aux_buffer_names(self): -+ return [name for name, _ in self.aux_buffers] -+ - def get_buf(self, idx: int): - return ( - self.output_ids[idx], -@@ -191,8 +191,12 @@ class MetadataBuffers: - self.output_topk_index[idx], - self.output_hidden_states[idx], - self.bootstrap_room[idx], -+ self.prefill_timing[idx], - ) - -+ def clear_profiling_buf(self, idx: int): -+ self.prefill_timing[idx].zero_() -+ - def set_buf(self, req: Req): - - self.output_ids[req.metadata_buffer_index][0] = req.output_ids[0] -@@ -237,6 +241,84 @@ class MetadataBuffers: - self.bootstrap_room[req.metadata_buffer_index, 0] = ( - req.bootstrap_room if req.bootstrap_room is not None else 0 - ) -+ # Pack prefill-side PD timing durations for transfer to decode instance. -+ # Note: set_buf is called at the START of the last KV chunk send, so -+ # completion_time and prefill_transfer_queue_entry_time are not yet set. -+ # We use time.perf_counter() as the "forward just completed" timestamp. -+ import time -+ -+ ts = req.time_stats -+ timing = self.prefill_timing[req.metadata_buffer_index] -+ self.clear_profiling_buf(req.metadata_buffer_index) -+ if not is_slime_profiling_enabled(): -+ return -+ for idx, value in enumerate( -+ build_prefill_timing_payload(ts, now=time.perf_counter()) -+ ): -+ if value > 0: -+ timing[idx] = value -+ -+ -+def is_slime_profiling_enabled() -> bool: -+ return envs.SLIME_ENABLE_PROFILING.get() -+ -+ -+def build_prefill_timing_payload(time_stats, now: float) -> tuple[float, ...]: -+ bootstrap_queue_duration = 0.0 -+ if ( -+ time_stats.prefill_bootstrap_queue_entry_time > 0 -+ and time_stats.wait_queue_entry_time > 0 -+ ): -+ bootstrap_queue_duration = ( -+ time_stats.wait_queue_entry_time -+ - time_stats.prefill_bootstrap_queue_entry_time -+ ) -+ -+ prefill_forward_duration = ( -+ now - time_stats.forward_entry_time -+ if time_stats.forward_entry_time > 0 -+ else 0.0 -+ ) -+ -+ return ( -+ bootstrap_queue_duration, -+ prefill_forward_duration, -+ 0.0, -+ max(0.0, time_stats.bootstrap_duration), -+ max(0.0, time_stats.alloc_waiting_duration), -+ max(0.0, time_stats.transfer_speed_gb_s), -+ max(0.0, time_stats.transfer_total_mb), -+ float(max(0, time_stats.prefill_retry_count)), -+ ) -+ -+ -+def apply_prefill_timing_payload(time_stats, timing) -> None: -+ for value, (attr_name, caster) in zip( -+ timing[: len(PREFILL_TIMING_DEST_ATTRS)].tolist(), -+ PREFILL_TIMING_DEST_ATTRS, -+ ): -+ if value > 0: -+ setattr(time_stats, attr_name, caster(value)) -+ -+ -+def iter_aux_transfer_specs( -+ aux_buffer_names: list[str], -+ prefill_aux_ptrs: list[int], -+ prefill_aux_item_lens: list[int], -+ dst_aux_ptrs: list[int], -+ prefill_aux_index: int, -+ dst_aux_index: int, -+): -+ profiling_enabled = is_slime_profiling_enabled() -+ for i, (buffer_name, dst_aux_ptr) in enumerate(zip(aux_buffer_names, dst_aux_ptrs)): -+ if not profiling_enabled and buffer_name == PREFILL_TIMING_AUX_BUFFER_NAME: -+ continue -+ length = prefill_aux_item_lens[i] -+ if length <= 0: -+ continue -+ src_addr = prefill_aux_ptrs[i] + length * prefill_aux_index -+ dst_addr = dst_aux_ptr + length * dst_aux_index -+ yield i, src_addr, dst_addr, length - - - ######################### -diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py -index 0ed5a1b44b..67e33c650d 100644 ---- a/python/sglang/srt/entrypoints/engine.py -+++ b/python/sglang/srt/entrypoints/engine.py -@@ -52,6 +52,7 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqInput, - MultimodalDataInputFormat, - OpenSessionReqInput, -+ PostProcessWeightsReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, - RpcReqInput, -@@ -641,6 +642,24 @@ class Engine(EngineBase): - self.tokenizer_manager.update_weights_from_ipc(obj, None) - ) - -+ def post_process_weights( -+ self, -+ restore_weights_before_load: bool = False, -+ post_process_quantization: bool = False, -+ ): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ Should be called after weight update is finished. -+ """ -+ obj = PostProcessWeightsReqInput( -+ restore_weights_before_load=restore_weights_before_load, -+ post_process_quantization=post_process_quantization, -+ ) -+ -+ return self.loop.run_until_complete( -+ self.tokenizer_manager.post_process_weights(obj, None) -+ ) -+ - def get_weights_by_name(self, name: str, truncate_size: int = 100): - """Get weights by parameter name.""" - obj = GetWeightsByNameReqInput(name=name, truncate_size=truncate_size) -diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py -index 1d6816c010..402b42e05b 100644 ---- a/python/sglang/srt/entrypoints/http_server.py -+++ b/python/sglang/srt/entrypoints/http_server.py -@@ -115,6 +115,7 @@ from sglang.srt.managers.io_struct import ( - OpenSessionReqInput, - ParseFunctionCallReq, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReqInput, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -574,10 +575,8 @@ async def model_info(): - @app.get("/weight_version") - async def weight_version(): - """Get the current weight version.""" -- raise HTTPException( -- status_code=404, -- detail="Endpoint '/get_weight_version' or '/weight_version' is deprecated. Please use '/model_info' instead.", -- ) -+ result = await model_info() -+ return {"weight_version": result.get("weight_version", None)} - - - @app.get("/get_server_info") -@@ -594,9 +593,19 @@ async def get_server_info(): - async def server_info(): - """Get the server information.""" - # Returns internal states per DP. -- internal_states: List[Dict[Any, Any]] = ( -- await _global_state.tokenizer_manager.get_internal_state() -- ) -+ # In large/disaggregated deployments this can occasionally block; keep endpoint responsive. -+ server_info_timeout = float(os.environ.get("SGLANG_SERVER_INFO_TIMEOUT", "2")) -+ try: -+ internal_states: List[Dict[Any, Any]] = await asyncio.wait_for( -+ _global_state.tokenizer_manager.get_internal_state(), -+ timeout=server_info_timeout, -+ ) -+ except asyncio.TimeoutError: -+ logger.warning( -+ "Timed out getting internal state for /server_info after %.1fs; returning empty internal_states", -+ server_info_timeout, -+ ) -+ internal_states = [] - - # This field is not serializable. - if hasattr(_global_state.tokenizer_manager.server_args, "model_config"): -@@ -1084,6 +1093,23 @@ async def update_weights_from_ipc(obj: UpdateWeightsFromIPCReqInput, request: Re - return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST) - - -+@app.post("/post_process_weights") -+@auth_level(AuthLevel.ADMIN_OPTIONAL) -+async def post_process_weights(req: PostProcessWeightsReqInput, request: Request): -+ """ -+ Optional post-processing for updated weights (e.g., Marlin conversion). -+ This should be called selectively after `update_weights_from_distributed/update_weights_from_tensor`. -+ """ -+ success, message = await _global_state.tokenizer_manager.post_process_weights( -+ req, request -+ ) -+ -+ content = {"success": success, "message": message} -+ return ORJSONResponse( -+ content, status_code=200 if success else HTTPStatus.BAD_REQUEST -+ ) -+ -+ - @app.post("/update_weight_version") - @auth_level(AuthLevel.ADMIN_OPTIONAL) - async def update_weight_version(obj: UpdateWeightVersionReqInput, request: Request): -diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py -index 8293796a2e..bff34e4221 100644 ---- a/python/sglang/srt/environ.py -+++ b/python/sglang/srt/environ.py -@@ -244,6 +244,7 @@ class Envs: - SGLANG_DISAGGREGATION_HEARTBEAT_MAX_FAILURE = EnvInt(2) - SGLANG_DISAGGREGATION_WAITING_TIMEOUT = EnvInt(300) - SGLANG_DISAGGREGATION_NIXL_BACKEND = EnvStr("UCX") -+ SLIME_ENABLE_PROFILING = EnvBool(False) - - # Scheduler: others: - SGLANG_EMPTY_CACHE_INTERVAL = EnvFloat(-1) # in seconds. Set if you observe high memory accumulation over a long serving period. -diff --git a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py -index 1cdf65b91c..4783cd18fb 100644 ---- a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py -+++ b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py -@@ -630,7 +630,6 @@ def _get_k_and_s_triton( - page_indices, - k_out, - s_out, -- seq_len, - page_size, - buf_numel_per_page, - index_head_dim, -@@ -647,7 +646,6 @@ def _get_k_and_s_triton_kernel( - page_indices_ptr, - k_out_ptr, - s_out_ptr, -- seq_len: tl.constexpr, - page_size: tl.constexpr, - buf_numel_per_page: tl.constexpr, - index_head_dim: tl.constexpr, -diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -index ca54a931b7..3540f77bae 100644 ---- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -+++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py -@@ -1,6 +1,7 @@ - from __future__ import annotations - - import contextlib -+import os - from abc import ABC, abstractmethod - from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple - -@@ -201,14 +202,31 @@ class Indexer(MultiPlatformOp): - prefix=add_prefix("weights_proj", prefix), - ) - self.k_norm = LayerNorm(self.head_dim, dtype=torch.float32) -+ server_args = get_global_server_args() -+ disable_flag = server_args.disable_indexer_rope_neox_style -+ env_raw = os.environ.get("INDEXER_ROPE_NEOX_STYLE", None) -+ if env_raw is not None: -+ env_value = env_raw == "1" -+ if disable_flag and env_value: -+ raise ValueError( -+ "Conflict: --disable-indexer-rope-neox-style is set but " -+ "INDEXER_ROPE_NEOX_STYLE='1'. " -+ "Please remove one or make them consistent." -+ ) -+ resolved_neox_style = env_value -+ elif disable_flag: -+ resolved_neox_style = False -+ else: -+ resolved_neox_style = is_neox_style -+ - self.rotary_emb = get_rope_wrapper( - rope_head_dim, - rotary_dim=rope_head_dim, - max_position=max_position_embeddings, - base=rope_theta, # type: ignore - rope_scaling=rope_scaling, -- is_neox_style=is_neox_style, -- device=get_global_server_args().device, -+ is_neox_style=resolved_neox_style, -+ device=server_args.device, - ) - self.block_size = block_size - self.scale_fmt = scale_fmt -@@ -244,6 +262,11 @@ class Indexer(MultiPlatformOp): - x = x.to(self.weights_proj.weight.dtype) - weights, _ = self.weights_proj(x) - weights = weights.float() -+ if weights.shape[1] < q_scale.shape[1]: -+ assert q_scale.shape[1] % weights.shape[1] == 0 -+ weights = weights.repeat_interleave( -+ q_scale.shape[1] // weights.shape[1], dim=1 -+ ) - weights = weights * self.n_heads**-0.5 - weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale - return weights -@@ -982,15 +1005,26 @@ class Indexer(MultiPlatformOp): - query, key = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch - ) -+ if query.shape[1] < 32: -+ assert 32 % query.shape[1] == 0 -+ query = query.repeat_interleave(32 // query.shape[1], dim=1) - q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) - with torch.cuda.stream(self.alt_stream): - k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) - current_stream.wait_stream(self.alt_stream) -+ if weights.shape[1] < q_scale.shape[1]: -+ assert q_scale.shape[1] % weights.shape[1] == 0 -+ weights = weights.repeat_interleave( -+ q_scale.shape[1] // weights.shape[1], dim=1 -+ ) - weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale - else: - query, key = self._get_q_k_bf16( - q_lora, x, positions, enable_dual_stream, forward_batch=forward_batch - ) -+ if query.shape[1] < 32: -+ assert 32 % query.shape[1] == 0 -+ query = query.repeat_interleave(32 // query.shape[1], dim=1) - - if enable_dual_stream: - current_stream = torch.cuda.current_stream() -diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -index de8a07ab30..5c9f4813a6 100644 ---- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -+++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py -@@ -697,6 +697,7 @@ class FusedMoE(torch.nn.Module): - "CompressedTensorsWNA16TritonMoE", - ] - ) -+ and "zero" not in weight_name - else loaded_weight - ) - -@@ -916,6 +917,7 @@ class FusedMoE(torch.nn.Module): - "CompressedTensorsWNA16TritonMoE", - ] - ) -+ and "zero" not in weight_name - else loaded_weight - ) - -diff --git a/python/sglang/srt/layers/moe/routed_experts_capturer.py b/python/sglang/srt/layers/moe/routed_experts_capturer.py -index 00bd687555..12d5577af2 100644 ---- a/python/sglang/srt/layers/moe/routed_experts_capturer.py -+++ b/python/sglang/srt/layers/moe/routed_experts_capturer.py -@@ -8,10 +8,15 @@ import torch - - from sglang.srt.configs.model_config import ModelConfig - from sglang.srt.layers.dp_attention import ( -+ attn_tp_all_gather_into_tensor, - get_attention_dp_rank, -+ get_attention_tp_size, - get_dp_local_info, - is_dp_attention_enabled, - ) -+from sglang.srt.layers.moe import ( -+ get_moe_a2a_backend, -+) - from sglang.srt.mem_cache.memory_pool import ReqToTokenPool - from sglang.srt.model_executor.forward_batch_info import ForwardBatch - from sglang.srt.server_args import get_global_server_args -@@ -181,13 +186,26 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - device=device, - ) - -+ if get_moe_a2a_backend().is_deepep(): -+ attn_tp_size = get_attention_tp_size() if is_dp_attention_enabled() else 1 -+ self.gather_buffer = torch.empty( -+ ( -+ self.device_cache.buffer.shape[0] * attn_tp_size, -+ self.device_cache.buffer.shape[2], -+ ), -+ dtype=torch.int32, -+ device=device, -+ ) -+ - def _sync_fwd_experts_buffer_DtoH( - self, - forward_batch: ForwardBatch, - can_run_graph: bool, - cuda_graph_batch: int, - ): -- if is_dp_attention_enabled(): -+ # When DeepEP is enabled, capture() already does all_gather, so device_cache.buffer -+ # contains data from all DP ranks. We should not slice by DP rank in this case. -+ if is_dp_attention_enabled() and not get_moe_a2a_backend().is_deepep(): - local_start_pos, local_num_tokens = get_dp_local_info(forward_batch) - # handle with cuda graph padding - if can_run_graph: -@@ -206,6 +224,12 @@ class _RoutedExpertsCapturerReal(RoutedExpertsCapturer): - ].cpu() - - def capture(self, layer_id: int, topk_ids: torch.Tensor): -+ if get_moe_a2a_backend().is_deepep(): -+ local_topk_ids = topk_ids -+ topk_ids = self.gather_buffer[ -+ : local_topk_ids.size(0) * get_attention_tp_size() -+ ] -+ attn_tp_all_gather_into_tensor(topk_ids, local_topk_ids) - self.device_cache.capture_fwd_routed_experts(layer_id, topk_ids) - - def get_routed_experts( -diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -index 4cbfed6f90..88b4527443 100644 ---- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -+++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py -@@ -499,7 +499,7 @@ class CompressedTensorsConfig(QuantizationConfig): - ) - is_static = not weight_quant.dynamic - -- return is_channel_group and input_quant_none and is_symmetric and is_static -+ return is_channel_group and input_quant_none and is_static - - def _is_mxint4a16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool: - input_quant_none = input_quant is None -@@ -968,6 +968,9 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase): - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - layer.scheme.process_weights_after_loading(layer) - -+ def restore_weights_before_loading(self, layer: torch.nn.Module) -> None: -+ layer.scheme.restore_weights_before_loading(layer) -+ - def create_weights( - self, - layer: torch.nn.Module, -diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -index 6264f36d04..f0310e305e 100644 ---- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -+++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py -@@ -17,7 +17,10 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import ( - CompressedTensorsMoEScheme, - ) - from sglang.srt.layers.quantization.gptq import gptq_marlin_moe_repack --from sglang.srt.layers.quantization.marlin_utils import marlin_moe_permute_scales -+from sglang.srt.layers.quantization.marlin_utils import ( -+ marlin_moe_permute_scales, -+ moe_awq_to_marlin_zero_points, -+) - from sglang.srt.layers.quantization.utils import replace_parameter - from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs - -@@ -64,7 +67,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - self.strategy = config.strategy - self.group_size = config.group_size - self.actorder = config.actorder -- assert config.symmetric, "Only symmetric quantization is supported for MoE" -+ self.sym = config.symmetric - - if not ( - self.quant_config.quant_format == CompressionFormat.pack_quantized.value -@@ -124,7 +127,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - - # In the case where we have actorder/g_idx, - # we do not partition the w2 scales -- load_full_w2 = self.actorder and self.group_size != -1 -+ load_full_w2 = (self.actorder != "static") and self.group_size != -1 - - if load_full_w2: - w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size -@@ -172,6 +175,32 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - layer.register_parameter("w13_weight_shape", w13_weight_shape) - set_weight_attrs(w13_weight_shape, extra_weight_attrs) - -+ # add zero param -+ if not self.sym: -+ w13_qzeros = torch.nn.Parameter( -+ torch.empty( -+ num_experts, -+ num_groups_w13, -+ 2 * intermediate_size_per_partition // self.packed_factor, -+ dtype=torch.int32, -+ ), -+ requires_grad=False, -+ ) -+ layer.register_parameter("w13_weight_zero_point", w13_qzeros) -+ set_weight_attrs(w13_qzeros, extra_weight_attrs) -+ -+ w2_qzeros = torch.nn.Parameter( -+ torch.empty( -+ num_experts, -+ num_groups_w2, -+ hidden_size // self.packed_factor, -+ dtype=torch.int32, -+ ), -+ requires_grad=False, -+ ) -+ layer.register_parameter("w2_weight_zero_point", w2_qzeros) -+ set_weight_attrs(w2_qzeros, extra_weight_attrs) -+ - w13_g_idx = torch.nn.Parameter( - torch.empty( - num_experts, -@@ -225,11 +254,14 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - - # Force record: these are the target GPTQ shapes for rollback. - layer._original_shapes["w13_weight_packed"] = tuple(w13_weight.shape) -- layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) -+ layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) -+ if not self.sym: -+ layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape - -- # Also record the shapes of the scales. -+ layer._original_shapes["w2_weight_packed"] = tuple(w2_weight.shape) - layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape) -- layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape) -+ if not self.sym: -+ layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape) - - def process_weights_after_loading(self, layer: torch.nn.Module) -> None: - -@@ -334,6 +366,24 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - ) - replace_tensor("w2_weight_scale", marlin_w2_scales) - -+ # Repack zero -+ if not self.sym: -+ marlin_w13_zp = moe_awq_to_marlin_zero_points( -+ layer.w13_weight_zero_point, -+ size_k=layer.w13_weight_zero_point.shape[1], -+ size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor, -+ num_bits=self.num_bits, -+ ) -+ replace_tensor("w13_weight_zero_point", marlin_w13_zp) -+ -+ marlin_w2_zp = moe_awq_to_marlin_zero_points( -+ layer.w2_weight_zero_point, -+ size_k=layer.w2_weight_zero_point.shape[1], -+ size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor, -+ num_bits=self.num_bits, -+ ) -+ replace_tensor("w2_weight_zero_point", marlin_w2_zp) -+ - layer.is_marlin_converted = True - - def restore_weights_before_loading(self, layer: torch.nn.Module): -@@ -399,6 +449,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme): - g_idx2=layer.w2_weight_g_idx, - sort_indices1=layer.w13_g_idx_sort_indices, - sort_indices2=layer.w2_g_idx_sort_indices, -+ w1_zeros=layer.w13_weight_zero_point if not self.sym else None, -+ w2_zeros=layer.w2_weight_zero_point if not self.sym else None, - num_bits=self.num_bits, - is_k_full=self.is_k_full, - routed_scaling_factor=self.moe_runner_config.routed_scaling_factor, -diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py -index 6522278603..7d3a5d0c4c 100644 ---- a/python/sglang/srt/managers/detokenizer_manager.py -+++ b/python/sglang/srt/managers/detokenizer_manager.py -@@ -405,6 +405,17 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): - prefill_launch_delay=recv_obj.prefill_launch_delay, - prefill_launch_latency=recv_obj.prefill_launch_latency, - prefill_finished_ts=recv_obj.prefill_finished_ts, -+ pd_prefill_bootstrap_queue_duration=recv_obj.pd_prefill_bootstrap_queue_duration, -+ pd_prefill_forward_duration=recv_obj.pd_prefill_forward_duration, -+ pd_prefill_transfer_queue_duration=recv_obj.pd_prefill_transfer_queue_duration, -+ pd_decode_prealloc_duration=recv_obj.pd_decode_prealloc_duration, -+ pd_decode_transfer_duration=recv_obj.pd_decode_transfer_duration, -+ pd_decode_forward_duration=recv_obj.pd_decode_forward_duration, -+ pd_bootstrap_duration=recv_obj.pd_bootstrap_duration, -+ pd_alloc_waiting_duration=recv_obj.pd_alloc_waiting_duration, -+ pd_transfer_speed_gb_s=recv_obj.pd_transfer_speed_gb_s, -+ pd_transfer_total_mb=recv_obj.pd_transfer_total_mb, -+ pd_prefill_retry_count=recv_obj.pd_prefill_retry_count, - ) - - def handle_multimodal_decode_req(self, recv_obj: BatchMultimodalDecodeReq): -diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py -index ff17745673..f947e71d7d 100644 ---- a/python/sglang/srt/managers/io_struct.py -+++ b/python/sglang/srt/managers/io_struct.py -@@ -101,6 +101,42 @@ class RequestTimingMetricsMixin: - # This marks when the prefill computation finishes. - prefill_finished_ts: Optional[List[Optional[float]]] - -+ # --- PD disaggregation timing fields --- -+ # All fields are None when profiling is disabled or not in PD disaggregation mode. -+ -+ # P instance: duration spent in bootstrap queue before entering the wait queue. -+ pd_prefill_bootstrap_queue_duration: Optional[List[Optional[float]]] -+ -+ # P instance: duration for the actual prefill forward computation. -+ pd_prefill_forward_duration: Optional[List[Optional[float]]] -+ -+ # P instance: duration spent in the KV transfer queue. -+ pd_prefill_transfer_queue_duration: Optional[List[Optional[float]]] -+ -+ # D instance: duration waiting for KV cache slot pre-allocation. -+ pd_decode_prealloc_duration: Optional[List[Optional[float]]] -+ -+ # D instance: duration waiting for the KV cache transfer to complete. -+ pd_decode_transfer_duration: Optional[List[Optional[float]]] -+ -+ # D instance: duration for the actual decode forward computation. -+ pd_decode_forward_duration: Optional[List[Optional[float]]] -+ -+ # Bootstrap handshake duration (P and D instances). -+ pd_bootstrap_duration: Optional[List[Optional[float]]] -+ -+ # KV cache allocation waiting duration (P and D instances). -+ pd_alloc_waiting_duration: Optional[List[Optional[float]]] -+ -+ # KV cache transfer speed in GB/s. -+ pd_transfer_speed_gb_s: Optional[List[Optional[float]]] -+ -+ # Total KV cache transferred in MB. -+ pd_transfer_total_mb: Optional[List[Optional[float]]] -+ -+ # Number of prefill retries (P instance only). -+ pd_prefill_retry_count: Optional[List[Optional[int]]] -+ - - @dataclass - class SpeculativeDecodingMetricsMixin: -@@ -1403,6 +1439,20 @@ class UpdateWeightsFromIPCReqOutput(BaseReq): - message: str - - -+@dataclass -+class PostProcessWeightsReqInput(BaseReq): -+ # Whether to restore weights before loading new weights -+ restore_weights_before_load: bool = False -+ # Whether to enable quantization post-processing -+ post_process_quantization: bool = False -+ -+ -+@dataclass -+class PostProcessWeightsReqOutput(BaseReq): -+ success: bool -+ message: str -+ -+ - @dataclass - class InitWeightsSendGroupForRemoteInstanceReqOutput(BaseReq): - success: bool -@@ -1802,6 +1852,10 @@ class GetLoadReqOutput(BaseReq): - num_waiting_reqs: int - num_tokens: int - ts_tic: float -+ # Per-queue breakdown: list of {name, num_reqs, num_tokens, reqs: [{rid, seqlen, input_len, output_len}]} -+ queue_details: Optional[List[Dict[str, Any]]] = None -+ # Running batch info -+ running_details: Optional[Dict[str, Any]] = None - - - @dataclass -diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py -index e1236aa0f3..daa598a1f6 100644 ---- a/python/sglang/srt/managers/multi_tokenizer_mixin.py -+++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py -@@ -142,6 +142,39 @@ def _handle_output_by_index(output, i): - prefill_finished_ts=_extract_field_by_index( - output, "prefill_finished_ts", i - ), -+ pd_prefill_bootstrap_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_bootstrap_queue_duration", i -+ ), -+ pd_prefill_forward_duration=_extract_field_by_index( -+ output, "pd_prefill_forward_duration", i -+ ), -+ pd_prefill_transfer_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_transfer_queue_duration", i -+ ), -+ pd_decode_prealloc_duration=_extract_field_by_index( -+ output, "pd_decode_prealloc_duration", i -+ ), -+ pd_decode_transfer_duration=_extract_field_by_index( -+ output, "pd_decode_transfer_duration", i -+ ), -+ pd_decode_forward_duration=_extract_field_by_index( -+ output, "pd_decode_forward_duration", i -+ ), -+ pd_bootstrap_duration=_extract_field_by_index( -+ output, "pd_bootstrap_duration", i -+ ), -+ pd_alloc_waiting_duration=_extract_field_by_index( -+ output, "pd_alloc_waiting_duration", i -+ ), -+ pd_transfer_speed_gb_s=_extract_field_by_index( -+ output, "pd_transfer_speed_gb_s", i -+ ), -+ pd_transfer_total_mb=_extract_field_by_index( -+ output, "pd_transfer_total_mb", i -+ ), -+ pd_prefill_retry_count=_extract_field_by_index( -+ output, "pd_prefill_retry_count", i -+ ), - finished_reasons=_extract_field_by_index(output, "finished_reasons", i), - decoded_texts=_extract_field_by_index(output, "decoded_texts", i), - decode_ids=_extract_field_by_index(output, "decode_ids", i), -@@ -211,6 +244,50 @@ def _handle_output_by_index(output, i): - elif isinstance(output, BatchEmbeddingOutput): - new_output = BatchEmbeddingOutput( - rids=[output.rids[i]], -+ queue_time=_extract_field_by_index(output, "queue_time", i), -+ forward_entry_time=_extract_field_by_index(output, "forward_entry_time", i), -+ prefill_launch_delay=_extract_field_by_index( -+ output, "prefill_launch_delay", i -+ ), -+ prefill_launch_latency=_extract_field_by_index( -+ output, "prefill_launch_latency", i -+ ), -+ prefill_finished_ts=_extract_field_by_index( -+ output, "prefill_finished_ts", i -+ ), -+ pd_prefill_bootstrap_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_bootstrap_queue_duration", i -+ ), -+ pd_prefill_forward_duration=_extract_field_by_index( -+ output, "pd_prefill_forward_duration", i -+ ), -+ pd_prefill_transfer_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_transfer_queue_duration", i -+ ), -+ pd_decode_prealloc_duration=_extract_field_by_index( -+ output, "pd_decode_prealloc_duration", i -+ ), -+ pd_decode_transfer_duration=_extract_field_by_index( -+ output, "pd_decode_transfer_duration", i -+ ), -+ pd_decode_forward_duration=_extract_field_by_index( -+ output, "pd_decode_forward_duration", i -+ ), -+ pd_bootstrap_duration=_extract_field_by_index( -+ output, "pd_bootstrap_duration", i -+ ), -+ pd_alloc_waiting_duration=_extract_field_by_index( -+ output, "pd_alloc_waiting_duration", i -+ ), -+ pd_transfer_speed_gb_s=_extract_field_by_index( -+ output, "pd_transfer_speed_gb_s", i -+ ), -+ pd_transfer_total_mb=_extract_field_by_index( -+ output, "pd_transfer_total_mb", i -+ ), -+ pd_prefill_retry_count=_extract_field_by_index( -+ output, "pd_prefill_retry_count", i -+ ), - finished_reasons=_extract_field_by_index(output, "finished_reasons", i), - embeddings=_extract_field_by_index(output, "embeddings", i), - prompt_tokens=_extract_field_by_index(output, "prompt_tokens", i), -@@ -239,6 +316,39 @@ def _handle_output_by_index(output, i): - prefill_finished_ts=_extract_field_by_index( - output, "prefill_finished_ts", i - ), -+ pd_prefill_bootstrap_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_bootstrap_queue_duration", i -+ ), -+ pd_prefill_forward_duration=_extract_field_by_index( -+ output, "pd_prefill_forward_duration", i -+ ), -+ pd_prefill_transfer_queue_duration=_extract_field_by_index( -+ output, "pd_prefill_transfer_queue_duration", i -+ ), -+ pd_decode_prealloc_duration=_extract_field_by_index( -+ output, "pd_decode_prealloc_duration", i -+ ), -+ pd_decode_transfer_duration=_extract_field_by_index( -+ output, "pd_decode_transfer_duration", i -+ ), -+ pd_decode_forward_duration=_extract_field_by_index( -+ output, "pd_decode_forward_duration", i -+ ), -+ pd_bootstrap_duration=_extract_field_by_index( -+ output, "pd_bootstrap_duration", i -+ ), -+ pd_alloc_waiting_duration=_extract_field_by_index( -+ output, "pd_alloc_waiting_duration", i -+ ), -+ pd_transfer_speed_gb_s=_extract_field_by_index( -+ output, "pd_transfer_speed_gb_s", i -+ ), -+ pd_transfer_total_mb=_extract_field_by_index( -+ output, "pd_transfer_total_mb", i -+ ), -+ pd_prefill_retry_count=_extract_field_by_index( -+ output, "pd_prefill_retry_count", i -+ ), - finished_reasons=_extract_field_by_index(output, "finished_reasons", i), - output_strs=_extract_field_by_index(output, "output_strs", i), - output_ids=_extract_field_by_index(output, "output_ids", i), -@@ -524,6 +634,60 @@ def monkey_patch_uvicorn_multiprocessing(timeout: float = 10): - "uvicorn.supervisors.multiprocess not found, skipping monkey patch" - ) - -+ # Fix stdin fd issue when running under Ray (or other managed -+ # environments where stdin may not be a real terminal): -+ # -+ # Uvicorn's get_subprocess() captures sys.stdin.fileno() in the parent -+ # and passes it to spawn'd children, which call os.fdopen(stdin_fileno) -+ # to re-attach stdin. This is intended for interactive debugging (e.g. -+ # pdb attach to a child worker). -+ # -+ # In Ray Actors, sys.stdin.fileno() succeeds in the parent (returns a -+ # valid fd number), but the fd is not inheritable across spawn. The -+ # child's os.fdopen() then crashes with OSError: [Errno 9] Bad file -+ # descriptor, killing every tokenizer worker. -+ # -+ # Instead of unconditionally disabling stdin passthrough, we probe -+ # whether the fd is truly usable by dup'ing it. If os.dup() fails, -+ # the fd won't survive spawn either, so we fall back to None. In a -+ # normal terminal environment os.dup() succeeds and debugging ability -+ # is preserved. -+ try: -+ import uvicorn._subprocess as _uv_sub -+ import uvicorn.supervisors.multiprocess as _uv_mp -+ -+ def _safe_get_stdin_fileno(): -+ """Return stdin fileno only if it is genuinely usable.""" -+ try: -+ fileno = sys.stdin.fileno() -+ # Verify the fd is valid and duplicable — if it isn't, -+ # spawn'd children won't be able to reopen it either. -+ dup_fd = os.dup(fileno) -+ os.close(dup_fd) -+ return fileno -+ except (AttributeError, OSError): -+ return None -+ -+ def _patched_get_subprocess(config, target, sockets): -+ stdin_fileno = _safe_get_stdin_fileno() -+ kwargs = { -+ "config": config, -+ "target": target, -+ "sockets": sockets, -+ "stdin_fileno": stdin_fileno, -+ } -+ return _uv_sub.spawn.Process( -+ target=_uv_sub.subprocess_started, kwargs=kwargs -+ ) -+ -+ # Must patch both: the supervisor module caches its own reference -+ # to get_subprocess at import time via -+ # ``from uvicorn._subprocess import get_subprocess``. -+ _uv_sub.get_subprocess = _patched_get_subprocess -+ _uv_mp.get_subprocess = _patched_get_subprocess -+ except Exception: -+ pass -+ - - class SenderWrapper: - def __init__(self, port_args: PortArgs, send_to_scheduler: zmq.Socket): -diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py -index c079957980..dd8ca7167d 100644 ---- a/python/sglang/srt/managers/schedule_batch.py -+++ b/python/sglang/srt/managers/schedule_batch.py -@@ -1869,7 +1869,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): - while first_iter or ( - not self.check_decode_mem(selected_indices=sorted_indices) - ): -- if len(sorted_indices) == 1: -+ # We should allow all requests to be retracted in decode disaggregation mode -+ # because there call be prealloc prefill requests. -+ num_minimum_reqs = 0 if server_args.disaggregation_mode == "decode" else 1 -+ if len(sorted_indices) == num_minimum_reqs: - # Always keep at least one request - break - -diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py -index a9ff0ac94b..a50dd5122b 100644 ---- a/python/sglang/srt/managers/scheduler.py -+++ b/python/sglang/srt/managers/scheduler.py -@@ -114,6 +114,7 @@ from sglang.srt.managers.io_struct import ( - OpenSessionReqInput, - OpenSessionReqOutput, - PauseGenerationReqInput, -+ PostProcessWeightsReqInput, - ProfileReq, - ReleaseMemoryOccupationReqInput, - ResumeMemoryOccupationReqInput, -@@ -1063,6 +1064,7 @@ class Scheduler( - ), - (UpdateWeightsFromTensorReqInput, self.update_weights_from_tensor), - (UpdateWeightsFromIPCReqInput, self.update_weights_from_ipc), -+ (PostProcessWeightsReqInput, self.post_process_weights), - (GetWeightsByNameReqInput, self.get_weights_by_name), - (ReleaseMemoryOccupationReqInput, self.release_memory_occupation), - (ResumeMemoryOccupationReqInput, self.resume_memory_occupation), -diff --git a/python/sglang/srt/managers/scheduler_metrics_mixin.py b/python/sglang/srt/managers/scheduler_metrics_mixin.py -index 30b2732b9f..68090b1617 100644 ---- a/python/sglang/srt/managers/scheduler_metrics_mixin.py -+++ b/python/sglang/srt/managers/scheduler_metrics_mixin.py -@@ -609,12 +609,54 @@ class SchedulerMetricsMixin: - num_tokens += sum(req.seqlen for queue in waiting_queues for req in queue) - num_waiting_reqs = sum(len(queue) for queue in waiting_queues) - -+ # Collect per-queue details -+ queue_names = ["waiting_queue"] -+ if self.disaggregation_mode == DisaggregationMode.PREFILL: -+ queue_names.append("bootstrap_queue") -+ elif self.disaggregation_mode == DisaggregationMode.DECODE: -+ queue_names.append("prealloc_queue") -+ queue_names.append("transfer_queue") -+ queue_names.append("retracted_queue") -+ -+ queue_details = [] -+ for name, queue in zip(queue_names, waiting_queues): -+ reqs_info = [] -+ for req in queue: -+ reqs_info.append( -+ { -+ "seqlen": req.seqlen, -+ } -+ ) -+ queue_details.append( -+ { -+ "name": name, -+ "num_reqs": len(queue), -+ "num_tokens": sum(r["seqlen"] for r in reqs_info), -+ "reqs": reqs_info, -+ } -+ ) -+ -+ # Collect running batch details -+ running_reqs_info = [] -+ for req in self.running_batch.reqs: -+ running_reqs_info.append( -+ { -+ "seqlen": req.seqlen, -+ } -+ ) -+ running_details = { -+ "num_reqs": len(self.running_batch.reqs), -+ "reqs": running_reqs_info, -+ } -+ - return GetLoadReqOutput( - dp_rank=self.dp_rank, - num_reqs=len(self.running_batch.reqs) + num_waiting_reqs, - num_waiting_reqs=num_waiting_reqs, - num_tokens=num_tokens, - ts_tic=time.perf_counter(), -+ queue_details=queue_details, -+ running_details=running_details, - ) - - def get_loads(self: Scheduler, req: GetLoadsReqInput = None) -> GetLoadsReqOutput: -diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -index 482bc6ca66..fbc4864176 100644 ---- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py -+++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py -@@ -922,6 +922,18 @@ class SchedulerOutputProcessorMixin: - prefill_launch_delays = [] - prefill_launch_latencies = [] - prefill_finished_timestamps = [] -+ profiling_enabled = envs.SLIME_ENABLE_PROFILING.get() -+ pd_prefill_bootstrap_queue_durations = [] if profiling_enabled else None -+ pd_prefill_forward_durations = [] if profiling_enabled else None -+ pd_prefill_transfer_queue_durations = [] if profiling_enabled else None -+ pd_decode_prealloc_durations = [] if profiling_enabled else None -+ pd_decode_transfer_durations = [] if profiling_enabled else None -+ pd_decode_forward_durations = [] if profiling_enabled else None -+ pd_bootstrap_durations = [] if profiling_enabled else None -+ pd_alloc_waiting_durations = [] if profiling_enabled else None -+ pd_transfer_speeds_gb_s = [] if profiling_enabled else None -+ pd_transfer_totals_mb = [] if profiling_enabled else None -+ pd_prefill_retry_counts = [] if profiling_enabled else None - - if return_logprob: - input_token_logprobs_val = [] -@@ -1037,6 +1049,40 @@ class SchedulerOutputProcessorMixin: - prefill_finished_timestamps.append( - req.time_stats.get_prefill_finished_ts() - ) -+ if profiling_enabled: -+ pd_prefill_bootstrap_queue_durations.append( -+ req.time_stats.get_pd_prefill_bootstrap_queue_duration() -+ ) -+ pd_prefill_forward_durations.append( -+ req.time_stats.get_pd_prefill_forward_duration() -+ ) -+ pd_prefill_transfer_queue_durations.append( -+ req.time_stats.get_pd_prefill_transfer_queue_duration() -+ ) -+ pd_decode_prealloc_durations.append( -+ req.time_stats.get_pd_decode_prealloc_duration() -+ ) -+ pd_decode_transfer_durations.append( -+ req.time_stats.get_pd_decode_transfer_duration() -+ ) -+ pd_decode_forward_durations.append( -+ req.time_stats.get_pd_decode_forward_duration() -+ ) -+ pd_bootstrap_durations.append( -+ req.time_stats.get_pd_bootstrap_duration() -+ ) -+ pd_alloc_waiting_durations.append( -+ req.time_stats.get_pd_alloc_waiting_duration() -+ ) -+ pd_transfer_speeds_gb_s.append( -+ req.time_stats.get_pd_transfer_speed_gb_s() -+ ) -+ pd_transfer_totals_mb.append( -+ req.time_stats.get_pd_transfer_total_mb() -+ ) -+ pd_prefill_retry_counts.append( -+ req.time_stats.get_pd_prefill_retry_count() -+ ) - - if not self.spec_algorithm.is_none(): - spec_verify_ct.append(req.spec_verify_ct) -@@ -1134,7 +1180,7 @@ class SchedulerOutputProcessorMixin: - req.log_time_stats() - - # Send to detokenizer -- if reqs or is_idle_batch: -+ if rids or is_idle_batch: - if self.model_config.is_multimodal_gen: - return - self.send_to_detokenizer.send_output( -@@ -1149,6 +1195,17 @@ class SchedulerOutputProcessorMixin: - prefill_launch_delay=prefill_launch_delays, - prefill_launch_latency=prefill_launch_latencies, - prefill_finished_ts=prefill_finished_timestamps, -+ pd_prefill_bootstrap_queue_duration=pd_prefill_bootstrap_queue_durations, -+ pd_prefill_forward_duration=pd_prefill_forward_durations, -+ pd_prefill_transfer_queue_duration=pd_prefill_transfer_queue_durations, -+ pd_decode_prealloc_duration=pd_decode_prealloc_durations, -+ pd_decode_transfer_duration=pd_decode_transfer_durations, -+ pd_decode_forward_duration=pd_decode_forward_durations, -+ pd_bootstrap_duration=pd_bootstrap_durations, -+ pd_alloc_waiting_duration=pd_alloc_waiting_durations, -+ pd_transfer_speed_gb_s=pd_transfer_speeds_gb_s, -+ pd_transfer_total_mb=pd_transfer_totals_mb, -+ pd_prefill_retry_count=pd_prefill_retry_counts, - finished_reasons=finished_reasons, - decoded_texts=decoded_texts, - decode_ids=decode_ids_list, -@@ -1198,6 +1255,18 @@ class SchedulerOutputProcessorMixin: - prefill_launch_delays = [] - prefill_launch_latencies = [] - prefill_finished_timestamps = [] -+ profiling_enabled = envs.SLIME_ENABLE_PROFILING.get() -+ pd_prefill_bootstrap_queue_durations = [] if profiling_enabled else None -+ pd_prefill_forward_durations = [] if profiling_enabled else None -+ pd_prefill_transfer_queue_durations = [] if profiling_enabled else None -+ pd_decode_prealloc_durations = [] if profiling_enabled else None -+ pd_decode_transfer_durations = [] if profiling_enabled else None -+ pd_decode_forward_durations = [] if profiling_enabled else None -+ pd_bootstrap_durations = [] if profiling_enabled else None -+ pd_alloc_waiting_durations = [] if profiling_enabled else None -+ pd_transfer_speeds_gb_s = [] if profiling_enabled else None -+ pd_transfer_totals_mb = [] if profiling_enabled else None -+ pd_prefill_retry_counts = [] if profiling_enabled else None - retraction_counts = [] - for req in reqs: - if req.finished(): -@@ -1221,6 +1290,40 @@ class SchedulerOutputProcessorMixin: - prefill_finished_timestamps.append( - req.time_stats.get_prefill_finished_ts() - ) -+ if profiling_enabled: -+ pd_prefill_bootstrap_queue_durations.append( -+ req.time_stats.get_pd_prefill_bootstrap_queue_duration() -+ ) -+ pd_prefill_forward_durations.append( -+ req.time_stats.get_pd_prefill_forward_duration() -+ ) -+ pd_prefill_transfer_queue_durations.append( -+ req.time_stats.get_pd_prefill_transfer_queue_duration() -+ ) -+ pd_decode_prealloc_durations.append( -+ req.time_stats.get_pd_decode_prealloc_duration() -+ ) -+ pd_decode_transfer_durations.append( -+ req.time_stats.get_pd_decode_transfer_duration() -+ ) -+ pd_decode_forward_durations.append( -+ req.time_stats.get_pd_decode_forward_duration() -+ ) -+ pd_bootstrap_durations.append( -+ req.time_stats.get_pd_bootstrap_duration() -+ ) -+ pd_alloc_waiting_durations.append( -+ req.time_stats.get_pd_alloc_waiting_duration() -+ ) -+ pd_transfer_speeds_gb_s.append( -+ req.time_stats.get_pd_transfer_speed_gb_s() -+ ) -+ pd_transfer_totals_mb.append( -+ req.time_stats.get_pd_transfer_total_mb() -+ ) -+ pd_prefill_retry_counts.append( -+ req.time_stats.get_pd_prefill_retry_count() -+ ) - retraction_counts.append(req.retraction_count) - self.send_to_detokenizer.send_output( - BatchEmbeddingOutput( -@@ -1231,6 +1334,17 @@ class SchedulerOutputProcessorMixin: - prefill_launch_delay=prefill_launch_delays, - prefill_launch_latency=prefill_launch_latencies, - prefill_finished_ts=prefill_finished_timestamps, -+ pd_prefill_bootstrap_queue_duration=pd_prefill_bootstrap_queue_durations, -+ pd_prefill_forward_duration=pd_prefill_forward_durations, -+ pd_prefill_transfer_queue_duration=pd_prefill_transfer_queue_durations, -+ pd_decode_prealloc_duration=pd_decode_prealloc_durations, -+ pd_decode_transfer_duration=pd_decode_transfer_durations, -+ pd_decode_forward_duration=pd_decode_forward_durations, -+ pd_bootstrap_duration=pd_bootstrap_durations, -+ pd_alloc_waiting_duration=pd_alloc_waiting_durations, -+ pd_transfer_speed_gb_s=pd_transfer_speeds_gb_s, -+ pd_transfer_total_mb=pd_transfer_totals_mb, -+ pd_prefill_retry_count=pd_prefill_retry_counts, - finished_reasons=finished_reasons, - embeddings=embeddings, - prompt_tokens=prompt_tokens, -diff --git a/python/sglang/srt/managers/scheduler_profiler_mixin.py b/python/sglang/srt/managers/scheduler_profiler_mixin.py -index 7d08f12b35..afc045da20 100644 ---- a/python/sglang/srt/managers/scheduler_profiler_mixin.py -+++ b/python/sglang/srt/managers/scheduler_profiler_mixin.py -@@ -347,7 +347,7 @@ class SchedulerProfilerMixin: - if self.profiler_prefill_ct > self.profiler_target_prefill_ct: - if self.profile_in_progress: - self.stop_profile(stage=ForwardMode.EXTEND) -- elif batch.forward_mode.is_decode(): -+ elif batch.forward_mode.is_decode() or batch.forward_mode.is_prebuilt(): - if self.profiler_decode_ct == 0: - if self.profile_in_progress: - # force trace flush -diff --git a/python/sglang/srt/managers/scheduler_update_weights_mixin.py b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -index 293a843508..244ea4eb1b 100644 ---- a/python/sglang/srt/managers/scheduler_update_weights_mixin.py -+++ b/python/sglang/srt/managers/scheduler_update_weights_mixin.py -@@ -12,6 +12,7 @@ from sglang.srt.constants import ( - GPU_MEMORY_TYPE_KV_CACHE, - GPU_MEMORY_TYPE_WEIGHTS, - ) -+from sglang.srt.disaggregation.utils import DisaggregationMode - from sglang.srt.managers.io_struct import ( - CheckWeightsReqInput, - CheckWeightsReqOutput, -@@ -21,6 +22,8 @@ from sglang.srt.managers.io_struct import ( - GetWeightsByNameReqOutput, - InitWeightsUpdateGroupReqInput, - InitWeightsUpdateGroupReqOutput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ReleaseMemoryOccupationReqInput, - ReleaseMemoryOccupationReqOutput, - ResumeMemoryOccupationReqInput, -@@ -114,6 +117,11 @@ class SchedulerUpdateWeightsMixin: - torch.distributed.barrier(group=self.tp_cpu_group) - return UpdateWeightsFromIPCReqOutput(success, message) - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Optional post-processing for updated weights (e.g., Marlin conversion).""" -+ success, message = self.tp_worker.post_process_weights(recv_req) -+ return PostProcessWeightsReqOutput(success, message) -+ - def get_weights_by_name(self: Scheduler, recv_req: GetWeightsByNameReqInput): - parameter = self.tp_worker.get_weights_by_name(recv_req) - return GetWeightsByNameReqOutput(parameter) -@@ -137,6 +145,15 @@ class SchedulerUpdateWeightsMixin: - self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE) - self.flush_cache() - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_transfer_queue"): -+ self.disagg_decode_transfer_queue.release_memory_occupation() -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.release_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.release_memory_occupation() -+ - if GPU_MEMORY_TYPE_WEIGHTS in tags: - self.stashed_model_static_state = _export_static_state( - self.tp_worker.model_runner.model -@@ -177,6 +194,15 @@ class SchedulerUpdateWeightsMixin: - if GPU_MEMORY_TYPE_KV_CACHE in tags: - self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) - -+ if self.disaggregation_mode == DisaggregationMode.DECODE: -+ if hasattr(self, "disagg_decode_transfer_queue"): -+ self.disagg_decode_transfer_queue.resume_memory_occupation() -+ if hasattr(self, "disagg_decode_prealloc_queue"): -+ self.disagg_decode_prealloc_queue.resume_memory_occupation() -+ elif self.disaggregation_mode == DisaggregationMode.PREFILL: -+ if hasattr(self, "disagg_prefill_bootstrap_queue"): -+ self.disagg_prefill_bootstrap_queue.resume_memory_occupation() -+ - return ResumeMemoryOccupationReqOutput() - - def check_weights(self: Scheduler, recv_req: CheckWeightsReqInput): -diff --git a/python/sglang/srt/managers/tokenizer_communicator_mixin.py b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -index f2ffa9909d..6e4d1d460b 100644 ---- a/python/sglang/srt/managers/tokenizer_communicator_mixin.py -+++ b/python/sglang/srt/managers/tokenizer_communicator_mixin.py -@@ -59,6 +59,8 @@ from sglang.srt.managers.io_struct import ( - LoadLoRAAdapterReqOutput, - LoRAUpdateOutput, - OpenSessionReqInput, -+ PostProcessWeightsReqInput, -+ PostProcessWeightsReqOutput, - ProfileReq, - ProfileReqOutput, - ProfileReqType, -@@ -187,6 +189,9 @@ class TokenizerCommunicatorMixin: - self.update_weights_from_ipc_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -+ self.post_process_weights_communicator = _Communicator( -+ self.send_to_scheduler, server_args.dp_size -+ ) - self.get_weights_by_name_communicator = _Communicator( - self.send_to_scheduler, server_args.dp_size - ) -@@ -272,6 +277,10 @@ class TokenizerCommunicatorMixin: - UpdateWeightsFromIPCReqOutput, - self.update_weights_from_ipc_communicator.handle_recv, - ), -+ ( -+ PostProcessWeightsReqOutput, -+ self.post_process_weights_communicator.handle_recv, -+ ), - ( - GetWeightsByNameReqOutput, - self.get_weights_by_name_communicator.handle_recv, -@@ -522,6 +531,17 @@ class TokenizerCommunicatorMixin: - - return success, message - -+ async def post_process_weights( -+ self: TokenizerManager, -+ obj: PostProcessWeightsReqInput, -+ request: Optional[fastapi.Request] = None, -+ ) -> Tuple[bool, str]: -+ """Trigger post-processing hooks for weights after loading (e.g., Marlin conversion).""" -+ self.auto_create_handle_loop() -+ async with self.model_update_lock.writer_lock: -+ results = await self.post_process_weights_communicator(obj) -+ return _Communicator.merge_results(results) -+ - async def init_weights_send_group_for_remote_instance( - self, - obj: InitWeightsSendGroupForRemoteInstanceReqInput, -diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py -index 0914a5230b..9114e3e713 100644 ---- a/python/sglang/srt/managers/tokenizer_manager.py -+++ b/python/sglang/srt/managers/tokenizer_manager.py -@@ -1327,7 +1327,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi - async with self.is_pause_cond: - self.is_pause = True - if obj.mode != "abort": -- await self.send_to_scheduler.send_pyobj(obj) -+ self.send_to_scheduler.send_pyobj(obj) - else: - # we are using the model_update_lock to check if there is still on-going requests. - while True: -@@ -1341,7 +1341,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi - async def continue_generation(self, obj: ContinueGenerationReqInput): - async with self.is_pause_cond: - self.is_pause = False -- await self.send_to_scheduler.send_pyobj(obj) -+ self.send_to_scheduler.send_pyobj(obj) - self.is_pause_cond.notify_all() - - async def update_weights_from_disk( -@@ -1510,6 +1510,40 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi - self._add_metric_if_present( - recv_obj, "prefill_finished_ts", meta_info, i - ) -+ # PD disaggregation timing -+ self._add_metric_if_present( -+ recv_obj, "pd_prefill_bootstrap_queue_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_prefill_forward_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_prefill_transfer_queue_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_decode_prealloc_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_decode_transfer_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_decode_forward_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_bootstrap_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_alloc_waiting_duration", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_transfer_speed_gb_s", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_transfer_total_mb", meta_info, i -+ ) -+ self._add_metric_if_present( -+ recv_obj, "pd_prefill_retry_count", meta_info, i -+ ) - - if getattr(state.obj, "return_logprob", False): - self.convert_logprob_style( -@@ -1955,19 +1989,17 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi - if custom_labels - else self.metrics_collector.labels - ) -- if ( -- state.first_token_time == 0.0 -- and self.disaggregation_mode != DisaggregationMode.PREFILL -- ): -+ if state.first_token_time == 0.0: - state.first_token_time = state.last_time = time.time() - state.first_token_time_perf = time.perf_counter() - state.last_completion_tokens = completion_tokens -- self.metrics_collector.observe_time_to_first_token( -- labels, state.first_token_time - state.created_time -- ) -+ if self.disaggregation_mode != DisaggregationMode.PREFILL: -+ self.metrics_collector.observe_time_to_first_token( -+ labels, state.first_token_time - state.created_time -+ ) - else: - num_new_tokens = completion_tokens - state.last_completion_tokens -- if num_new_tokens: -+ if num_new_tokens > 0: - new_time = time.time() - interval = new_time - state.last_time - self.metrics_collector.observe_inter_token_latency( -@@ -1976,7 +2008,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi - num_new_tokens, - ) - state.last_time = new_time -- state.last_completion_tokens = completion_tokens -+ state.last_completion_tokens = completion_tokens - - if state.finished: - retraction_count = ( -diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py -index 86b009df4e..16ebd52ae3 100644 ---- a/python/sglang/srt/managers/tp_worker.py -+++ b/python/sglang/srt/managers/tp_worker.py -@@ -29,6 +29,7 @@ from sglang.srt.managers.io_struct import ( - InitWeightsUpdateGroupReqInput, - LoadLoRAAdapterFromTensorsReqInput, - LoadLoRAAdapterReqInput, -+ PostProcessWeightsReqInput, - SendWeightsToRemoteInstanceReqInput, - UnloadLoRAAdapterReqInput, - UpdateWeightFromDiskReqInput, -@@ -168,6 +169,11 @@ class BaseTpWorker(ABC): - success, message = self.model_runner.update_weights_from_ipc(recv_req) - return success, message - -+ def post_process_weights(self, recv_req: PostProcessWeightsReqInput): -+ """Perform optional post-processing on the updated model weights (e.g., Marlin conversion).""" -+ success, message = self.model_runner.post_process_weights(recv_req) -+ return success, message -+ - def get_weights_by_name(self, recv_req: GetWeightsByNameReqInput): - parameter = self.model_runner.get_weights_by_name( - recv_req.name, recv_req.truncate_size -diff --git a/python/sglang/srt/mem_cache/allocator.py b/python/sglang/srt/mem_cache/allocator.py -index fa08bb66a4..22c1c2a127 100644 ---- a/python/sglang/srt/mem_cache/allocator.py -+++ b/python/sglang/srt/mem_cache/allocator.py -@@ -411,7 +411,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): - - self.seen_max_num_extend_tokens_next_power_of_2 = max( - self.seen_max_num_extend_tokens_next_power_of_2, -- min(tl.core.TRITON_MAX_TENSOR_NUMEL, next_power_of_2(extend_num_tokens)), -+ min(65536, next_power_of_2(extend_num_tokens)), - ) - - bs = len(prefix_lens) -@@ -424,7 +424,7 @@ class PagedTokenToKVPoolAllocator(BaseTokenToKVPoolAllocator): - (extend_num_tokens,), dtype=torch.int64, device=self.device - ) - -- if extend_num_tokens < tl.core.TRITON_MAX_TENSOR_NUMEL: -+ if extend_num_tokens < 65536: - alloc_extend_kernel[(bs,)]( - prefix_lens, - seq_lens, -diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py -index d7cd472a98..9cf1185cbb 100644 ---- a/python/sglang/srt/mem_cache/hiradix_cache.py -+++ b/python/sglang/srt/mem_cache/hiradix_cache.py -@@ -750,9 +750,8 @@ class HiRadixCache(RadixCache): - self._update_leaf_status(node) - self._update_host_leaf_status(node) - if node.parent is None: -- assert ( -- node is self.root_node -- ), f"This request holds the node from another tree" -+ # Node belongs to a stale (flushed) tree — stop traversal gracefully. -+ break - node = node.parent - return delta - -@@ -827,6 +826,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(node) - # update leaf status for the parent because the node is evicted - self._update_leaf_status(node.parent) -+ self._update_host_leaf_status(node.parent) - return num_evicted - - def _evict_regular(self, node: TreeNode): -@@ -1330,6 +1330,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(node) - # update parent status as a new leaf is added into device - self._update_leaf_status(node.parent) -+ self._update_host_leaf_status(node.parent) - else: - self._inc_hit_count(node, chunked) - total_prefix_length += prefix_len -@@ -1345,6 +1346,7 @@ class HiRadixCache(RadixCache): - self._update_host_leaf_status(new_node) - # update parent status as a new leaf is added into device - self._update_leaf_status(new_node.parent) -+ self._update_host_leaf_status(new_node.parent) - else: - self._inc_hit_count(new_node, chunked) - total_prefix_length += prefix_len -diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py -index 1d917137c6..669e5c5181 100644 ---- a/python/sglang/srt/mem_cache/memory_pool.py -+++ b/python/sglang/srt/mem_cache/memory_pool.py -@@ -1777,9 +1777,12 @@ class NSATokenToKVPool(MLATokenToKVPool): - else: - assert self.page_size == 64 - with ( -- torch.cuda.use_mem_pool(self.custom_mem_pool) -- if self.custom_mem_pool -- else nullcontext() -+ ( -+ torch.cuda.use_mem_pool(self.custom_mem_pool) -+ if self.custom_mem_pool -+ else nullcontext() -+ ), -+ self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE), - ): - self.index_k_with_scale_buffer = [ - torch.zeros( -@@ -1801,6 +1804,11 @@ class NSATokenToKVPool(MLATokenToKVPool): - ) - for _ in range(layer_num) - ] -+ self.index_k_with_scale_buffer_ptrs = torch.tensor( -+ [x.data_ptr() for x in self.index_k_with_scale_buffer], -+ dtype=torch.uint64, -+ device=self.device, -+ ) - self._finalize_allocation_log(size) - - def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: -@@ -1876,6 +1884,50 @@ class NSATokenToKVPool(MLATokenToKVPool): - ] - return data_ptrs, data_lens, item_lens - -+ def get_cpu_copy(self, indices): -+ # First, save the kv_buffer (inherited from MLATokenToKVPool) -+ kv_cache_cpu = super().get_cpu_copy(indices) -+ -+ # Additionally, save the index_k_with_scale_buffer (page-indexed) -+ page_indices = indices[:: self.page_size] // self.page_size -+ torch.cuda.synchronize() -+ index_k_cpu = [] -+ chunk_size = self.cpu_offloading_chunk_size -+ # Convert chunk_size from token-level to page-level -+ page_chunk_size = max(1, chunk_size // self.page_size) -+ for layer_id in range(self.layer_num): -+ index_k_cpu.append([]) -+ for i in range(0, len(page_indices), page_chunk_size): -+ chunk_page_indices = page_indices[i : i + page_chunk_size] -+ idx_cpu = self.index_k_with_scale_buffer[layer_id][ -+ chunk_page_indices -+ ].to("cpu", non_blocking=True) -+ index_k_cpu[-1].append(idx_cpu) -+ torch.cuda.synchronize() -+ -+ return {"kv": kv_cache_cpu, "index_k": index_k_cpu} -+ -+ def load_cpu_copy(self, kv_cache_cpu_dict, indices): -+ # Restore the kv_buffer (inherited from MLATokenToKVPool) -+ super().load_cpu_copy(kv_cache_cpu_dict["kv"], indices) -+ -+ # Restore the index_k_with_scale_buffer (page-indexed) -+ page_indices = indices[:: self.page_size] // self.page_size -+ index_k_cpu = kv_cache_cpu_dict["index_k"] -+ torch.cuda.synchronize() -+ chunk_size = self.cpu_offloading_chunk_size -+ page_chunk_size = max(1, chunk_size // self.page_size) -+ for layer_id in range(self.layer_num): -+ for i in range(0, len(page_indices), page_chunk_size): -+ chunk_page_indices = page_indices[i : i + page_chunk_size] -+ idx_cpu = index_k_cpu[layer_id][i // page_chunk_size] -+ assert idx_cpu.shape[0] == len(chunk_page_indices) -+ idx_chunk = idx_cpu.to( -+ self.index_k_with_scale_buffer[0].device, non_blocking=True -+ ) -+ self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk -+ torch.cuda.synchronize() -+ - def get_kv_size_bytes(self): - kv_size_bytes = super().get_kv_size_bytes() - for index_k_cache in self.index_k_with_scale_buffer: -diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py -index 42b169728a..8e799196a4 100644 ---- a/python/sglang/srt/mem_cache/radix_cache.py -+++ b/python/sglang/srt/mem_cache/radix_cache.py -@@ -495,7 +495,17 @@ class RadixCache(BasePrefixCache): - if self.disable: - return - -- token_ids = req.fill_ids -+ # Limit to kv_committed_len to avoid including tokens (e.g., the just-generated -+ # token in disagg prefill) that don't have computed KV yet. If fill_ids is longer -+ # than kv_committed_len, the extra tokens would produce stale values (0 from -+ # req_to_token_pool initialization), leading to spurious tree nodes and memory -+ # leak when page-aligned token counts happen to cross a page boundary. -+ kv_committed_len = req.kv_committed_len -+ token_ids = ( -+ req.fill_ids[:kv_committed_len] -+ if kv_committed_len < len(req.fill_ids) -+ else req.fill_ids -+ ) - kv_indices = self.req_to_token_pool.req_to_token[ - req.req_pool_idx, : len(token_ids) - ] -@@ -619,9 +629,8 @@ class RadixCache(BasePrefixCache): - node.lock_ref -= 1 - self._update_leaf_status(node) - if node.parent is None: -- assert ( -- node is self.root_node -- ), f"This request holds the node from another tree" -+ # Node belongs to a stale (flushed) tree — stop traversal gracefully. -+ break - node = node.parent - return delta - -diff --git a/python/sglang/srt/metrics/collector.py b/python/sglang/srt/metrics/collector.py -index 255d41ccc0..f93bedb4dc 100644 ---- a/python/sglang/srt/metrics/collector.py -+++ b/python/sglang/srt/metrics/collector.py -@@ -20,7 +20,10 @@ import time - from dataclasses import dataclass, field - from typing import Any, Dict, List, Optional, Union - --from sglang.srt.disaggregation.utils import DisaggregationMode -+from sglang.srt.disaggregation.utils import ( -+ DisaggregationMode, -+ is_slime_profiling_enabled, -+) - from sglang.srt.environ import envs - from sglang.srt.metrics.utils import exponential_buckets, generate_buckets - from sglang.srt.model_executor.forward_batch_info import ForwardMode -@@ -77,6 +80,17 @@ class TimeStats: - # Number of prefill retries for this request - prefill_retry_count: int = 0 - -+ # Prefill-side durations forwarded via metadata transfer from P to D instance. -+ # Set on the decode instance after KV cache transfer completes. -+ fwd_prefill_bootstrap_queue_duration: Optional[float] = None -+ fwd_prefill_forward_duration: Optional[float] = None -+ fwd_prefill_transfer_queue_duration: Optional[float] = None -+ fwd_bootstrap_duration: Optional[float] = None -+ fwd_alloc_waiting_duration: Optional[float] = None -+ fwd_transfer_speed_gb_s: Optional[float] = None -+ fwd_transfer_total_mb: Optional[float] = None -+ fwd_prefill_retry_count: Optional[int] = None -+ - # Timestamp when prefill phase finishes, obtained from `time.time()`. - # Note that this differs from the other `_time` fields tracked by the - # `TimeStats` class, which are obtained from `time.perf_counter()`. -@@ -102,6 +116,148 @@ class TimeStats: - return self.prefill_finished_ts - return None - -+ # --- PD disaggregation timing getters --- -+ -+ def get_pd_prefill_bootstrap_queue_duration(self) -> Optional[float]: -+ """P instance: time spent in bootstrap queue before entering the wait queue.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_bootstrap_queue_duration is not None: -+ return self.fwd_prefill_bootstrap_queue_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.prefill_bootstrap_queue_entry_time > 0.0 -+ and self.wait_queue_entry_time > 0.0 -+ ): -+ return self.wait_queue_entry_time - self.prefill_bootstrap_queue_entry_time -+ return None -+ -+ def get_pd_prefill_forward_duration(self) -> Optional[float]: -+ """P instance: time for the actual prefill forward computation.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_forward_duration is not None: -+ return self.fwd_prefill_forward_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.forward_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.forward_entry_time -+ return None -+ -+ def get_pd_prefill_transfer_queue_duration(self) -> Optional[float]: -+ """P instance: time spent in the transfer queue (KV cache send).""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_transfer_queue_duration is not None: -+ return self.fwd_prefill_transfer_queue_duration -+ if ( -+ self.disagg_mode == DisaggregationMode.PREFILL -+ and self.prefill_transfer_queue_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.prefill_transfer_queue_entry_time -+ return None -+ -+ def get_pd_decode_prealloc_duration(self) -> Optional[float]: -+ """D instance: time spent in the pre-alloc queue (waiting for KV cache slot allocation).""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_prealloc_queue_entry_time > 0.0 -+ and self.decode_transfer_queue_entry_time > 0.0 -+ ): -+ return ( -+ self.decode_transfer_queue_entry_time -+ - self.decode_prealloc_queue_entry_time -+ ) -+ return None -+ -+ def get_pd_decode_transfer_duration(self) -> Optional[float]: -+ """D instance: time spent waiting for KV cache transfer to complete.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.decode_transfer_queue_entry_time > 0.0 -+ and self.wait_queue_entry_time > 0.0 -+ ): -+ return self.wait_queue_entry_time - self.decode_transfer_queue_entry_time -+ return None -+ -+ def get_pd_decode_forward_duration(self) -> Optional[float]: -+ """D instance: time for the actual decode forward computation.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if ( -+ self.disagg_mode == DisaggregationMode.DECODE -+ and self.forward_entry_time > 0.0 -+ and self.completion_time > 0.0 -+ ): -+ return self.completion_time - self.forward_entry_time -+ return None -+ -+ def get_pd_bootstrap_duration(self) -> Optional[float]: -+ """Bootstrap handshake duration (both P and D instances).""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_bootstrap_duration is not None: -+ return self.fwd_bootstrap_duration -+ if ( -+ self.disagg_mode != DisaggregationMode.NULL -+ and self.bootstrap_duration > 0.0 -+ ): -+ return self.bootstrap_duration -+ return None -+ -+ def get_pd_alloc_waiting_duration(self) -> Optional[float]: -+ """KV cache allocation waiting duration (both P and D instances).""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_alloc_waiting_duration is not None: -+ return self.fwd_alloc_waiting_duration -+ if ( -+ self.disagg_mode != DisaggregationMode.NULL -+ and self.alloc_waiting_duration > 0.0 -+ ): -+ return self.alloc_waiting_duration -+ return None -+ -+ def get_pd_transfer_speed_gb_s(self) -> Optional[float]: -+ """KV cache transfer speed in GB/s.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_transfer_speed_gb_s is not None: -+ return self.fwd_transfer_speed_gb_s -+ if ( -+ self.disagg_mode != DisaggregationMode.NULL -+ and self.transfer_speed_gb_s > 0.0 -+ ): -+ return self.transfer_speed_gb_s -+ return None -+ -+ def get_pd_transfer_total_mb(self) -> Optional[float]: -+ """Total KV cache transferred in MB.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_transfer_total_mb is not None: -+ return self.fwd_transfer_total_mb -+ if self.disagg_mode != DisaggregationMode.NULL and self.transfer_total_mb > 0.0: -+ return self.transfer_total_mb -+ return None -+ -+ def get_pd_prefill_retry_count(self) -> Optional[int]: -+ """Number of prefill retries for this request.""" -+ if not is_slime_profiling_enabled(): -+ return None -+ if self.fwd_prefill_retry_count is not None: -+ return self.fwd_prefill_retry_count -+ if self.disagg_mode == DisaggregationMode.PREFILL: -+ return self.prefill_retry_count -+ return None -+ - def convert_to_duration(self) -> str: - if self.disagg_mode == DisaggregationMode.NULL: - queue_duration = self.forward_entry_time - self.wait_queue_entry_time -diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py -index 275775a73d..e4e2fdc398 100644 ---- a/python/sglang/srt/model_executor/model_runner.py -+++ b/python/sglang/srt/model_executor/model_runner.py -@@ -395,7 +395,12 @@ class ModelRunner(ModelRunnerKVCacheMixin): - self.forward_stream = torch.get_device_module(self.device).Stream() - - # CPU offload -- set_offloader(create_offloader_from_server_args(server_args, dp_rank=dp_rank)) -+ # For draft worker (e.g., MTP), do not set offloader to avoid overriding -+ # the main model's offloader. Draft worker uses NoopOffloader instead. -+ if not is_draft_worker: -+ set_offloader( -+ create_offloader_from_server_args(server_args, dp_rank=dp_rank) -+ ) - - self._weight_checker = WeightChecker(model_runner=self) - -@@ -600,7 +605,8 @@ class ModelRunner(ModelRunnerKVCacheMixin): - ) - - # Init routed experts capturer -- self.init_routed_experts_capturer() -+ if not self.is_draft_worker: -+ self.init_routed_experts_capturer() - - if self.device == "cuda" or self.device == "musa": - self.init_cublas() -@@ -2429,11 +2435,19 @@ class ModelRunner(ModelRunnerKVCacheMixin): - output.expert_distribution_metrics = recorder_outputs.get("metrics") - - # Copy cached routing experts' buffers back to CPU cache -- get_global_experts_capturer().on_forward_end( -- forward_batch=forward_batch, -- can_run_graph=output.can_run_graph, -- cuda_graph_batch=getattr(self.graph_runner, "bs", None), -- ) -+ if not self.is_draft_worker: -+ # In speculative decoding, num_tokens_per_bs > 1, so we need to pass -+ # the actual number of tokens per dp rank in cuda graph, not batch size. -+ cuda_graph_num_tokens = None -+ if getattr(self.graph_runner, "bs", None): -+ cuda_graph_num_tokens = ( -+ self.graph_runner.bs * self.graph_runner.num_tokens_per_bs -+ ) -+ get_global_experts_capturer().on_forward_end( -+ forward_batch=forward_batch, -+ can_run_graph=output.can_run_graph, -+ cuda_graph_batch=cuda_graph_num_tokens, -+ ) - - if self.eplb_manager is not None: - self.eplb_manager.on_forward_pass_end() -@@ -2664,6 +2678,42 @@ class ModelRunner(ModelRunnerKVCacheMixin): - device=self.device, - ) - -+ def post_process_weights(self, recv_req): -+ """ -+ Execute post-processing logic for model weights, such as Marlin quantization format conversion. -+ """ -+ from sglang.srt.model_loader.loader import device_loading_context -+ -+ target_device = torch.device("cuda", torch.cuda.current_device()) -+ -+ if recv_req.restore_weights_before_load: -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports restoring weights -+ if quant_method is not None and hasattr( -+ quant_method, "restore_weights_before_loading" -+ ): -+ -+ with device_loading_context(module, target_device): -+ quant_method.restore_weights_before_loading(module) -+ -+ if recv_req.post_process_quantization: -+ # Iterate through all modules to apply specific post-loading processing -+ for _, module in self.model.named_modules(): -+ quant_method = getattr(module, "quant_method", None) -+ -+ # Check if the module supports quantization post-processing -+ if quant_method is not None and hasattr( -+ quant_method, "process_weights_after_loading" -+ ): -+ -+ # Apply the post-processing (e.g., repacking weights for Marlin kernel) -+ with device_loading_context(module, target_device): -+ quant_method.process_weights_after_loading(module) -+ -+ return True, "Success" -+ - - def _model_load_weights_direct(model, named_tensors: List[Tuple[str, torch.Tensor]]): - params_dict = dict(model.named_parameters()) -diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py -index cc673a9cac..06c430d2c4 100644 ---- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py -+++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py -@@ -1,4 +1,5 @@ - from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph -+from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend - from sglang.srt.layers.attention.tbo_backend import TboAttnBackend - from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( - AttnForwardMethod, -@@ -150,6 +151,8 @@ def handle_attention_nsa(attn, forward_batch): - backend = forward_batch.attn_backend - if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend - backend = backend.primary -+ if isinstance(backend, HybridAttnBackend): -+ backend = backend._select_backend(forward_batch.forward_mode) - if hasattr(backend, "use_mha") and backend.use_mha: - return AttnForwardMethod.MHA_ONE_SHOT - return AttnForwardMethod.MLA -diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py -index cb13a7c676..d9669ce086 100644 ---- a/python/sglang/srt/models/deepseek_nextn.py -+++ b/python/sglang/srt/models/deepseek_nextn.py -@@ -29,6 +29,7 @@ from sglang.srt.layers.attention.nsa.utils import ( - can_cp_split, - cp_all_gather_rerange_output, - cp_split_and_rebuild_data, -+ cp_split_and_rebuild_position, - is_nsa_enable_prefill_cp, - nsa_use_prefill_cp, - prepare_input_dp_with_cp_dsa, -@@ -160,15 +161,17 @@ class DeepseekModelNextN(nn.Module): - - if nsa_use_prefill_cp(forward_batch, self.nsa_enable_prefill_cp): - hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) -+ positions = cp_split_and_rebuild_position(forward_batch, positions) - residual = None - with get_global_expert_distribution_recorder().disable_this_region(): -- hidden_states, residual = self.decoder( -+ hidden_states, residual, *rest = self.decoder( - positions, - hidden_states, - forward_batch, - residual, - zero_allocator, - ) -+ topk_indices = rest[0] if rest else None - - if not forward_batch.forward_mode.is_idle(): - if residual is not None: -diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py -index 1583dd7880..a35c00f96c 100644 ---- a/python/sglang/srt/models/deepseek_v2.py -+++ b/python/sglang/srt/models/deepseek_v2.py -@@ -1085,6 +1085,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - prefix: str = "", - alt_stream: Optional[torch.cuda.Stream] = None, - skip_rope: bool = False, -+ is_nextn: bool = False, - ) -> None: - super().__init__() - self.layer_id = layer_id -@@ -1154,6 +1155,8 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - prefix=add_prefix("kv_a_proj_with_mqa", prefix), - ) - -+ self.skip_topk = False -+ self.next_skip_topk = False - if self.use_nsa: - is_neox_style = not getattr(config, "indexer_rope_interleave", False) - self.indexer = Indexer( -@@ -1174,6 +1177,31 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - layer_id=layer_id, - alt_stream=alt_stream, - ) -+ if not is_nextn: -+ self.index_topk_freq = getattr(config, "index_topk_freq", 1) -+ self.index_topk_pattern = getattr(config, "index_topk_pattern", None) -+ self.index_skip_topk_offset = getattr( -+ config, "index_skip_topk_offset", 2 -+ ) -+ if self.index_topk_pattern is None: -+ self.skip_topk = ( -+ max(layer_id - self.index_skip_topk_offset + 1, 0) -+ % self.index_topk_freq -+ != 0 -+ ) -+ self.next_skip_topk = ( -+ max(layer_id - self.index_skip_topk_offset + 2, 0) -+ % self.index_topk_freq -+ != 0 -+ ) -+ else: -+ self.skip_topk = self.index_topk_pattern[layer_id] == "S" -+ if layer_id < len(self.index_topk_pattern) - 1: -+ self.next_skip_topk = ( -+ self.index_topk_pattern[layer_id + 1] == "S" -+ ) -+ else: -+ self.next_skip_topk = False - - self.kv_b_proj = ColumnParallelLinear( - self.kv_lora_rank, -@@ -1362,6 +1390,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - llama_4_scaling: Optional[torch.Tensor] = None, -+ prev_topk_indices: Optional[torch.Tensor] = None, - ): - s = self.forward_prepare( - positions=positions, -@@ -1369,6 +1398,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - forward_batch=forward_batch, - zero_allocator=zero_allocator, - llama_4_scaling=llama_4_scaling, -+ prev_topk_indices=prev_topk_indices, - ) - return self.forward_core(s) - -@@ -1379,6 +1409,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - llama_4_scaling: Optional[torch.Tensor] = None, -+ prev_topk_indices: Optional[torch.Tensor] = None, - ): - if self.attn_mha.kv_b_proj is None: - self.attn_mha.kv_b_proj = self.kv_b_proj -@@ -1418,7 +1449,12 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - ) - elif attn_forward_method == AttnForwardMethod.MLA: - inner_state = self.forward_absorb_prepare( -- positions, hidden_states, forward_batch, zero_allocator, llama_4_scaling -+ positions, -+ hidden_states, -+ forward_batch, -+ zero_allocator, -+ llama_4_scaling, -+ prev_topk_indices, - ) - elif attn_forward_method == AttnForwardMethod.MLA_FUSED_ROPE: - inner_state = self.forward_absorb_fused_mla_rope_prepare( -@@ -1529,6 +1565,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - forward_batch: ForwardBatch, - zero_allocator: BumpAllocator, - llama_4_scaling: Optional[torch.Tensor] = None, -+ prev_topk_indices: Optional[torch.Tensor] = None, - ): - from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode - -@@ -1620,18 +1657,7 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - q = self.q_b_proj(q)[0].view( - -1, self.num_local_heads, self.qk_head_dim - ) -- topk_indices = self.indexer( -- x=hidden_states, -- q_lora=q_lora, -- positions=positions, -- forward_batch=forward_batch, -- layer_id=self.layer_id, -- ) -- current_stream.wait_stream(self.alt_stream) -- else: -- k_nope = k_nope.unsqueeze(1) -- q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) -- if q_lora is not None: -+ if not self.skip_topk: - topk_indices = self.indexer( - x=hidden_states, - q_lora=q_lora, -@@ -1639,6 +1665,23 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - forward_batch=forward_batch, - layer_id=self.layer_id, - ) -+ else: -+ topk_indices = prev_topk_indices -+ current_stream.wait_stream(self.alt_stream) -+ else: -+ k_nope = k_nope.unsqueeze(1) -+ q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim) -+ if q_lora is not None: -+ if not self.skip_topk: -+ topk_indices = self.indexer( -+ x=hidden_states, -+ q_lora=q_lora, -+ positions=positions, -+ forward_batch=forward_batch, -+ layer_id=self.layer_id, -+ ) -+ else: -+ topk_indices = prev_topk_indices - else: - q = self.q_proj(hidden_states)[0].view( - -1, self.num_local_heads, self.qk_head_dim -@@ -1929,8 +1972,10 @@ class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin): - ).transpose(0, 1), - ) - output, _ = self.o_proj(attn_bmm_output) -- -- return output -+ if not self.next_skip_topk: -+ return output, None -+ else: -+ return output, topk_indices - - def forward_absorb_fused_mla_rope_prepare( - self, -@@ -2275,6 +2320,7 @@ class DeepseekV2DecoderLayer(nn.Module): - reduce_results=False, - prefix=add_prefix("self_attn", prefix), - alt_stream=alt_stream, -+ is_nextn=is_nextn, - ) - - self.is_layer_sparse = self._is_layer_sparse(layer_id, is_nextn=is_nextn) -@@ -2357,6 +2403,7 @@ class DeepseekV2DecoderLayer(nn.Module): - zero_allocator: BumpAllocator, - gemm_output_zero_allocator: BumpAllocator = None, - llama_4_scaling: Optional[torch.Tensor] = None, -+ prev_topk_indices: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - quant_format = ( - "mxfp4" -@@ -2398,7 +2445,12 @@ class DeepseekV2DecoderLayer(nn.Module): - forward_batch=forward_batch, - zero_allocator=zero_allocator, - llama_4_scaling=llama_4_scaling, -+ prev_topk_indices=prev_topk_indices, - ) -+ if isinstance(hidden_states, tuple): -+ hidden_states, topk_indices = hidden_states -+ else: -+ topk_indices = None - - hidden_states, residual = self.layer_communicator.prepare_mlp( - hidden_states, residual, forward_batch -@@ -2434,7 +2486,7 @@ class DeepseekV2DecoderLayer(nn.Module): - hidden_states, residual, forward_batch - ) - -- return hidden_states, residual -+ return hidden_states, residual, topk_indices - - def op_comm_prepare_attn( - self, -@@ -2710,6 +2762,7 @@ class DeepseekV2Model(nn.Module): - elif self.first_k_dense_replace < normal_start_layer: - normal_end_layer = normal_start_layer = 0 - aux_hidden_states = [] -+ topk_indices = None - for i in range(normal_start_layer, normal_end_layer): - # NOTE: torch dynamo does not support graph break in context manager - ctx = ( -@@ -2727,7 +2780,7 @@ class DeepseekV2Model(nn.Module): - else: - aux_hidden_states.append(hidden_states + residual) - layer = self.layers[i] -- hidden_states, residual = layer( -+ hidden_states, residual, *rest = layer( - positions, - hidden_states, - forward_batch, -@@ -2735,7 +2788,9 @@ class DeepseekV2Model(nn.Module): - zero_allocator, - gemm_output_zero_allocator, - llama_4_scaling, -+ prev_topk_indices=topk_indices, - ) -+ topk_indices = rest[0] if rest else None - - if normal_end_layer != self.end_layer: - hidden_states, residual = model_forward_maybe_tbo( -diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py -index db8c1c7ce7..53ffadf6d0 100644 ---- a/python/sglang/srt/models/glm4_moe.py -+++ b/python/sglang/srt/models/glm4_moe.py -@@ -678,8 +678,13 @@ class Glm4MoeDecoderLayer(nn.Module): - nn.Module.__init__(self) - self.hidden_size = config.hidden_size - self.config = config -- rope_theta = getattr(config, "rope_theta", 10000) -- rope_scaling = getattr(config, "rope_scaling", None) -+ # rope_theta may be stored in rope_parameters dict (e.g. GLM-4.6V) -+ _rope_params = getattr(config, "rope_parameters", None) -+ if isinstance(_rope_params, dict) and "rope_theta" in _rope_params: -+ rope_theta = _rope_params["rope_theta"] -+ else: -+ rope_theta = getattr(config, "rope_theta", 10000) -+ rope_scaling = getattr(config, "rope_scaling", None) or _rope_params - partial_rotary_factor = getattr( - getattr(config, "rope_parameters", None), "partial_rotary_factor", None - ) or getattr(config, "partial_rotary_factor", 0.5) -@@ -773,6 +778,7 @@ class Glm4MoeDecoderLayer(nn.Module): - hidden_states: torch.Tensor, - forward_batch: ForwardBatch, - residual: Optional[torch.Tensor], -+ **kwargs, - ) -> torch.Tensor: - - hidden_states, residual = self.layer_communicator.prepare_attn( -diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py -index 1f6e753646..546cce4ab5 100644 ---- a/python/sglang/srt/models/glm4_moe_nextn.py -+++ b/python/sglang/srt/models/glm4_moe_nextn.py -@@ -103,7 +103,7 @@ class Glm4MoeModelNextN(nn.Module): - - residual = None - with get_global_expert_distribution_recorder().disable_this_region(): -- hidden_states, residual = self.decoder( -+ hidden_states, residual, *rest = self.decoder( - positions, hidden_states, forward_batch, residual - ) - -diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py -index 324de18b49..fc72faa031 100644 ---- a/python/sglang/srt/models/glm4v_moe.py -+++ b/python/sglang/srt/models/glm4v_moe.py -@@ -52,11 +52,31 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - self.num_fused_shared_experts = 0 - self.determine_num_fused_shared_experts() - -- self.model = Glm4MoeModel( -- config, -- quant_config, -- prefix=add_prefix("language_model", prefix), -- ) -+ if not self.config.encoder_only: -+ self.model = Glm4MoeModel( -+ config, -+ quant_config, -+ prefix=add_prefix("language_model", prefix), -+ ) -+ -+ if self.pp_group.is_last_rank: -+ if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: -+ self.lm_head = self.model.embed_tokens -+ else: -+ self.lm_head = ParallelLMHead( -+ config.vocab_size, -+ config.hidden_size, -+ quant_config=quant_config, -+ prefix=add_prefix("lm_head", prefix), -+ use_attn_tp_group=get_global_server_args().enable_dp_lm_head, -+ ) -+ else: -+ # ranks other than the last rank will have a placeholder layer -+ self.lm_head = PPMissingLayer() -+ else: -+ # encoder_only mode: no language model, so no lm_head needed -+ self.lm_head = None -+ - self.visual = Glm4vVisionModel( - config.vision_config, - quant_config=quant_config, -@@ -64,24 +84,14 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - use_data_parallel=self.use_data_parallel, - ) - -- if self.pp_group.is_last_rank: -- if self.pp_group.world_size == 1 and self.config.tie_word_embeddings: -- self.lm_head = self.model.embed_tokens -- else: -- self.lm_head = ParallelLMHead( -- config.vocab_size, -- config.hidden_size, -- quant_config=quant_config, -- prefix=add_prefix("lm_head", prefix), -- use_attn_tp_group=get_global_server_args().enable_dp_lm_head, -- ) -- else: -- # ranks other than the last rank will have a placeholder layer -- self.lm_head = PPMissingLayer() -- - self.logits_processor = LogitsProcessor(config) - self.pooler = Pooler(pooling_type=PoolingType.LAST, normalize=True) -- self.is_mrope_enabled = "mrope_section" in self.config.rope_scaling -+ _rope_cfg = ( -+ getattr(self.config, "rope_scaling", None) -+ or getattr(self.config, "rope_parameters", None) -+ or {} -+ ) -+ self.is_mrope_enabled = "mrope_section" in _rope_cfg - - # For EAGLE3 support - self.capture_aux_hidden_states = False -@@ -219,6 +229,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue -+ # Skip loading visual/language model weights -+ if ( -+ self.config.encoder_only or self.config.language_only -+ ) and name not in params_dict: -+ continue - if name not in params_dict: - continue - -@@ -234,6 +249,8 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - param_name, weight_name, expert_id, shard_id = mapping - if weight_name not in name: - continue -+ if "visual" in name or self.config.encoder_only: -+ continue - - # Mark as expert weight regardless of whether we can process it - is_expert_weight = True -@@ -265,6 +282,11 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): - # Skip loading extra bias for GPTQ models. - if name.endswith(".bias") and name not in params_dict: - continue -+ # Skip loading mm/language parameters -+ if ( -+ self.config.encoder_only or self.config.language_only -+ ) and name not in params_dict: -+ continue - if name not in params_dict: - continue - -diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py -index f01225487b..1dad8bb8e5 100644 ---- a/python/sglang/srt/models/qwen3_5.py -+++ b/python/sglang/srt/models/qwen3_5.py -@@ -372,6 +372,7 @@ class Qwen3_5LinearDecoderLayer(nn.Module): - input_layernorm=self.input_layernorm, - post_attention_layernorm=self.post_attention_layernorm, - allow_reduce_scatter=True, -+ is_last_layer=(layer_id == config.num_hidden_layers - 1), - ) - - def forward( -@@ -400,11 +401,24 @@ class Qwen3_5LinearDecoderLayer(nn.Module): - use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( - forward_batch - ) -- hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) - -- hidden_states, residual = self.layer_communicator.postprocess_layer( -- hidden_states, residual, forward_batch -+ should_allreduce_fusion = ( -+ self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( -+ forward_batch -+ ) - ) -+ if isinstance(self.mlp, Qwen2MoeSparseMoeBlock): -+ hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) -+ else: -+ hidden_states = self.mlp( -+ hidden_states, should_allreduce_fusion, use_reduce_scatter -+ ) -+ if should_allreduce_fusion: -+ hidden_states._sglang_needs_allreduce_fusion = True -+ else: -+ hidden_states, residual = self.layer_communicator.postprocess_layer( -+ hidden_states, residual, forward_batch -+ ) - - return hidden_states, residual - -@@ -549,6 +563,7 @@ class Qwen3_5AttentionDecoderLayer(nn.Module): - input_layernorm=self.input_layernorm, - post_attention_layernorm=self.post_attention_layernorm, - allow_reduce_scatter=True, -+ is_last_layer=(layer_id == config.num_hidden_layers - 1), - ) - - self.alt_stream = alt_stream -@@ -633,11 +648,24 @@ class Qwen3_5AttentionDecoderLayer(nn.Module): - use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( - forward_batch - ) -- hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) - -- hidden_states, residual = self.layer_communicator.postprocess_layer( -- hidden_states, residual, forward_batch -+ should_allreduce_fusion = ( -+ self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( -+ forward_batch -+ ) - ) -+ if isinstance(self.mlp, Qwen2MoeSparseMoeBlock): -+ hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) -+ else: -+ hidden_states = self.mlp( -+ hidden_states, should_allreduce_fusion, use_reduce_scatter -+ ) -+ if should_allreduce_fusion: -+ hidden_states._sglang_needs_allreduce_fusion = True -+ else: -+ hidden_states, residual = self.layer_communicator.postprocess_layer( -+ hidden_states, residual, forward_batch -+ ) - - return hidden_states, residual - -diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py -index d641826e33..3abc39ef32 100644 ---- a/python/sglang/srt/models/qwen3_vl.py -+++ b/python/sglang/srt/models/qwen3_vl.py -@@ -711,14 +711,19 @@ class Qwen3LLMModel(Qwen3Model): - hidden_states + residual if residual is not None else hidden_states - ) - -+ deepstack_embeds = None -+ if input_deepstack_embeds is not None: -+ prev_layer_idx = layer_idx - 1 -+ if prev_layer_idx in self.deepstack_embed_to_decoder_layer: -+ sep = self.hidden_size * prev_layer_idx -+ deepstack_embeds = input_deepstack_embeds[ -+ :, sep : sep + self.hidden_size -+ ] -+ - # SGLang applies residual at the START of the next layer, not at the END like HuggingFace. - # See: https://github.com/huggingface/transformers/blob/v5.0.0rc0/src/transformers/models/qwen3_vl/modeling_qwen3_vl.py#L549 - # To match HF behavior, deepstack must be added AFTER residual: (hidden_states + residual) + deepstack - # The order matters because addition with different tensors is not associative in practice. -- # Deepstack for prev_layer is applied at the start of current layer via post_residual_addition. -- deepstack_embeds = self.get_deepstack_embeds( -- layer_idx - 1, input_deepstack_embeds -- ) - hidden_states, residual = layer( - positions, - hidden_states, -diff --git a/python/sglang/srt/multimodal/processors/glm4v.py b/python/sglang/srt/multimodal/processors/glm4v.py -index 33cce6fe25..0970c4550d 100644 ---- a/python/sglang/srt/multimodal/processors/glm4v.py -+++ b/python/sglang/srt/multimodal/processors/glm4v.py -@@ -1,6 +1,9 @@ - from typing import List, Union - -+import torch -+ - from sglang.srt.layers.rotary_embedding import MRotaryEmbedding -+from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem - from sglang.srt.models.glm4v import Glm4vForConditionalGeneration - from sglang.srt.models.glm4v_moe import Glm4vMoeForConditionalGeneration - from sglang.srt.multimodal.processors.base_processor import ( -@@ -45,6 +48,8 @@ class Glm4vImageProcessor(SGLangBaseProcessor): - self.IMAGE_END_TOKEN_ID = hf_config.image_end_token_id - self.VIDEO_START_TOKEN_ID = hf_config.video_start_token_id - self.VIDEO_END_TOKEN_ID = hf_config.video_end_token_id -+ self.IM_START_TOKEN_ID = self.IMAGE_START_TOKEN_ID -+ self.IM_END_TOKEN_ID = self.IMAGE_END_TOKEN_ID - - # Vision config - self.IMAGE_FACTOR = 28 -@@ -59,6 +64,36 @@ class Glm4vImageProcessor(SGLangBaseProcessor): - video_token_id=self.IM_TOKEN_ID, - ).build(_processor) - -+ def get_mm_data(self, prompt, embeddings, img_grid_thw): -+ input_ids, offsets = self.build_input_ids(prompt, img_grid_thw) -+ mm_items = [ -+ MultimodalDataItem( -+ modality=Modality.IMAGE, -+ offsets=offsets, -+ precomputed_embeddings=embeddings, -+ ) -+ ] -+ -+ input_ids_tensor = torch.tensor(input_ids) -+ mrope_positions, mrope_position_delta = MRotaryEmbedding.get_rope_index_glm4v( -+ input_ids=input_ids_tensor.unsqueeze(0), -+ hf_config=self.hf_config, -+ image_grid_thw=img_grid_thw, -+ video_grid_thw=None, -+ attention_mask=None, -+ ) -+ mrope_positions = mrope_positions.squeeze(1) -+ -+ return { -+ "input_ids": input_ids, -+ "mm_items": mm_items, -+ "im_start_id": self.IM_START_TOKEN_ID, -+ "im_end_id": self.IM_END_TOKEN_ID, -+ "im_token_id": self.IM_TOKEN_ID, -+ "mrope_positions": mrope_positions, -+ "mrope_position_delta": mrope_position_delta, -+ } -+ - async def process_mm_data_async( - self, - image_data: List[Union[str, bytes]], -diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py -index 4395654e4e..f9b5ea4abb 100644 ---- a/python/sglang/srt/multimodal/processors/qwen_vl.py -+++ b/python/sglang/srt/multimodal/processors/qwen_vl.py -@@ -317,7 +317,7 @@ class QwenVLImageProcessor(SGLangBaseProcessor): - **kwargs, - ): - entry_time = time.perf_counter() -- base_output = self.load_mm_data( -+ base_output = self.legacy_load_mm_data( - prompt=input_text, - image_data=image_data, - video_data=request_obj.video_data, -diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py -index b080aeb168..5b29ebf566 100644 ---- a/python/sglang/srt/server_args.py -+++ b/python/sglang/srt/server_args.py -@@ -635,6 +635,7 @@ class ServerArgs: - # Context parallelism used in the long sequence prefill phase of DeepSeek v3.2 - enable_nsa_prefill_context_parallel: bool = False - nsa_prefill_cp_mode: str = "round-robin-split" -+ disable_indexer_rope_neox_style: bool = False - enable_fused_qk_norm_rope: bool = False - enable_precise_embedding_interpolation: bool = False - -@@ -4781,6 +4782,12 @@ class ServerArgs: - help="Token splitting mode for the prefill phase of DeepSeek v3.2 under context parallelism. Optional values: 'round-robin-split'(default), 'in-seq-split' " - "'round-robin-split' distributes tokens across ranks based on token_idx %% cp_size. It supports multi-batch prefill, fused MoE, and FP8 KV cache.", - ) -+ parser.add_argument( -+ "--disable-indexer-rope-neox-style", -+ action="store_true", -+ help="Disable NSA indexer RoPE neox style (equivalent to INDEXER_ROPE_NEOX_STYLE=0). " -+ "If the environment variable INDEXER_ROPE_NEOX_STYLE is also set and conflicts, an error is raised.", -+ ) - parser.add_argument( - "--enable-fused-qk-norm-rope", - action="store_true", -diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -index 5fe45086ca..b283d2e9bd 100644 ---- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -+++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py -@@ -341,7 +341,10 @@ class EAGLEDraftCudaGraphRunner: - self.seq_lens.fill_(self.seq_len_fill_value) - self.out_cache_loc.zero_() - self.positions.zero_() -- -+ self.topk_p.zero_() -+ self.topk_index.zero_() -+ self.hidden_states.zero_() -+ self.req_pool_indices.zero_() - num_tokens = bs * self.num_tokens_per_bs - - # Common inputs -@@ -350,8 +353,12 @@ class EAGLEDraftCudaGraphRunner: - forward_batch.out_cache_loc - ) - self.positions[:raw_num_token].copy_(forward_batch.positions) -- self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) -- self.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) -+ self.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p.clamp(0, 1)) -+ self.topk_index[:raw_bs].copy_( -+ forward_batch.spec_info.topk_index.clamp( -+ 0, self.model_runner.model_config.vocab_size - 1 -+ ) -+ ) - self.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) - self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) - -diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py -index ac629c7ee5..c039d23508 100644 ---- a/python/sglang/srt/speculative/eagle_info.py -+++ b/python/sglang/srt/speculative/eagle_info.py -@@ -774,6 +774,10 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] - self.verified_id = self.verified_id[: len(new_indices)] -+ if self.accept_length is not None: -+ self.accept_length = self.accept_length[: len(new_indices)] -+ if self.accept_length_cpu is not None: -+ self.accept_length_cpu = self.accept_length_cpu[: len(new_indices)] - else: - # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` - self.topk_p = self.topk_p[new_indices] -@@ -805,6 +809,27 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): - self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0) - self.topk_p = torch.cat([self.topk_p, spec_info.topk_p]) - self.topk_index = torch.cat([self.topk_index, spec_info.topk_index]) -+ if self.accept_length is not None and spec_info.accept_length is not None: -+ self.accept_length = torch.cat( -+ [self.accept_length, spec_info.accept_length] -+ ) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif self.accept_length is not None: -+ zeros = torch.zeros( -+ [spec_info.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([self.accept_length, zeros]) -+ self.accept_length_cpu = self.accept_length.tolist() -+ elif spec_info.accept_length is not None: -+ zeros = torch.zeros( -+ [self.verified_id.shape[0]], -+ dtype=self.accept_length.dtype, -+ device=self.accept_length.device, -+ ) -+ self.accept_length = torch.cat([zeros, spec_info.accept_length]) -+ self.accept_length_cpu = self.accept_length.tolist() - - - @dataclass -diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py -index 4636128fa7..a9b61df393 100644 ---- a/python/sglang/srt/utils/common.py -+++ b/python/sglang/srt/utils/common.py -@@ -2359,6 +2359,8 @@ class SafeUnpickler(pickle.Unpickler): - "sglang.srt.model_executor.model_runner.", - "sglang.srt.layers.", - "sglang.srt.utils.", -+ # --- slime --- -+ "slime.", - } - - DENY_CLASSES = { -diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py -index 3be16446e0..1b2371c839 100644 ---- a/python/sglang/srt/utils/weight_checker.py -+++ b/python/sglang/srt/utils/weight_checker.py -@@ -69,6 +69,9 @@ def _check_tensors( - actual_should_compare, - actual, - ) in zip(expect_tensors, actual_tensors, strict=True): -+ if ".cos_sin_cache" in expect_name: -+ # skip cos/sin cache which is deterministic from shape and dtype and may have different shapes due to different implementations. -+ continue - assert expect_name == actual_name, f"{expect_name=} {actual_name=}" - assert ( - expect_should_compare == actual_should_compare diff --git a/docs/conf.py b/docs/conf.py index e9a4215e8..929cbe0d6 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -122,8 +122,8 @@ html_context = { "display_github": True, - "github_user": "sgl-project", - "github_repo": "sgl-project.github.io", + "github_user": "vllm-project", + "github_repo": "vime", "github_version": "main", "conf_py_path": "/docs/", } diff --git a/docs/en/advanced/speculative-decoding.md b/docs/en/advanced/speculative-decoding.md index 9029edcc4..a98d01d1a 100644 --- a/docs/en/advanced/speculative-decoding.md +++ b/docs/en/advanced/speculative-decoding.md @@ -9,7 +9,7 @@ which slime forwards via `--vllm-speculative-config`. For models with MTP layers (e.g., GLM-4.7, DeepSeek-V3/R1), pass: ```bash ---vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' +--vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ``` To use a separately trained draft model, set `model` (and optionally `draft_tensor_parallel_size`) diff --git a/docs/en/examples/glm4.7-30B-A3B.md b/docs/en/examples/glm4.7-30B-A3B.md index ff8a549b6..8e10043b5 100644 --- a/docs/en/examples/glm4.7-30B-A3B.md +++ b/docs/en/examples/glm4.7-30B-A3B.md @@ -89,7 +89,7 @@ GLM-4.7-Flash includes 1 MTP (Multi-Token Prediction) layer, which can be used f VLLM_ARGS=( ... # MTP speculative decoding (EAGLE) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) ``` diff --git a/docs/en/examples/glm4.7-355B-A32B.md b/docs/en/examples/glm4.7-355B-A32B.md index 18d74b691..3ae3dc497 100644 --- a/docs/en/examples/glm4.7-355B-A32B.md +++ b/docs/en/examples/glm4.7-355B-A32B.md @@ -105,7 +105,7 @@ GLM-4.7 includes MTP (Multi-Token Prediction) layers that can be used for specul VLLM_ARGS=( ... # MTP speculative decoding (EAGLE) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) ``` @@ -176,7 +176,7 @@ VLLM_ARGS=( --vllm-enable-expert-parallel --vllm-cudagraph-capture-sizes 1 2 4 8 $(seq 16 8 128) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' --vllm-all2all-backend deepep_high_throughput ) diff --git a/docs/zh/advanced/speculative-decoding.md b/docs/zh/advanced/speculative-decoding.md index ad5723040..f65fcc6af 100644 --- a/docs/zh/advanced/speculative-decoding.md +++ b/docs/zh/advanced/speculative-decoding.md @@ -8,7 +8,7 @@ vLLM 把投机采样的所有配置收敛到一个 JSON(`SpeculativeConfig`) `--vllm-speculative-config` 透传。对于有 MTP 层的模型(例如 GLM-4.7、DeepSeek-V3/R1),传入: ```bash ---vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' +--vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ``` 如果要使用单独训练的 draft model,在同一个 JSON 里加上 `model`(可选还可加 diff --git a/docs/zh/examples/glm4.7-30B-A3B.md b/docs/zh/examples/glm4.7-30B-A3B.md index 290880910..0e0b005ad 100644 --- a/docs/zh/examples/glm4.7-30B-A3B.md +++ b/docs/zh/examples/glm4.7-30B-A3B.md @@ -87,7 +87,7 @@ GLM-4.7-Flash 包含 1 层 MTP(Multi-Token Prediction)层,可用于推理 VLLM_ARGS=( ... # MTP 投机解码 (EAGLE) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) ``` diff --git a/docs/zh/examples/glm4.7-355B-A32B.md b/docs/zh/examples/glm4.7-355B-A32B.md index 798b3c60a..56519fca2 100644 --- a/docs/zh/examples/glm4.7-355B-A32B.md +++ b/docs/zh/examples/glm4.7-355B-A32B.md @@ -105,7 +105,7 @@ GLM-4.7 包含 MTP(Multi-Token Prediction)层,可以在推理阶段用于 VLLM_ARGS=( ... # MTP 投机解码 (EAGLE) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) ``` @@ -176,7 +176,7 @@ VLLM_ARGS=( --vllm-enable-expert-parallel --vllm-cudagraph-capture-sizes 1 2 4 8 $(seq 16 8 128) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' --vllm-all2all-backend deepep_high_throughput ) diff --git a/examples/geo3k_vlm/run_geo3k_qwen35.sh b/examples/geo3k_vlm/run_geo3k_qwen35.sh index 31bd4bf96..2813898f1 100644 --- a/examples/geo3k_vlm/run_geo3k_qwen35.sh +++ b/examples/geo3k_vlm/run_geo3k_qwen35.sh @@ -122,7 +122,7 @@ VLLM_ARGS=( --vllm-cudagraph-capture-sizes 1 2 4 8 $(seq 16 8 256) # MTP speculative decoding - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' --vllm-max-num-seqs 512 ) diff --git a/requirements.txt b/requirements.txt index 2cdaa711c..7db56918c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,7 +13,6 @@ pyyaml qwen_vl_utils # for VLM ray[default] ring_flash_attn -sglang-router>=0.2.3 tensorboard transformers vllm-router>=0.1.14 diff --git a/scripts/run-deepseek-r1.sh b/scripts/run-deepseek-r1.sh index 821ade380..464c45567 100644 --- a/scripts/run-deepseek-r1.sh +++ b/scripts/run-deepseek-r1.sh @@ -123,7 +123,7 @@ VLLM_ARGS=( --vllm-server-concurrency 1024 --vllm-enable-expert-parallel --vllm-all2all-backend deepep_high_throughput - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) MISC_ARGS=( diff --git a/scripts/run-glm4.7-30B-A3B.sh b/scripts/run-glm4.7-30B-A3B.sh index 07d5bdbab..9cbf1dd31 100644 --- a/scripts/run-glm4.7-30B-A3B.sh +++ b/scripts/run-glm4.7-30B-A3B.sh @@ -121,7 +121,7 @@ VLLM_ARGS=( --vllm-max-num-seqs 512 - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' --vllm-max-cudagraph-capture-size 64 ) diff --git a/scripts/run-glm4.7-355B-A32B.sh b/scripts/run-glm4.7-355B-A32B.sh index 081f3c796..3e7673dd3 100644 --- a/scripts/run-glm4.7-355B-A32B.sh +++ b/scripts/run-glm4.7-355B-A32B.sh @@ -117,7 +117,7 @@ VLLM_ARGS=( # mtp --vllm-enable-expert-parallel - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) MISC_ARGS=( diff --git a/scripts/run-glm5-744B-A40B.sh b/scripts/run-glm5-744B-A40B.sh index 0071e89db..7949b0d67 100644 --- a/scripts/run-glm5-744B-A40B.sh +++ b/scripts/run-glm5-744B-A40B.sh @@ -119,7 +119,7 @@ VLLM_ARGS=( --vllm-max-num-seqs 512 --vllm-enable-expert-parallel --vllm-all2all-backend deepep_high_throughput - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' --vllm-block-size 64 --vllm-max-cudagraph-capture-size 8 --vllm-max-num-batched-tokens 131072 diff --git a/scripts/run-mimo-7B-rl-eagle.sh b/scripts/run-mimo-7B-rl-eagle.sh index b14d490ac..80c02631b 100644 --- a/scripts/run-mimo-7B-rl-eagle.sh +++ b/scripts/run-mimo-7B-rl-eagle.sh @@ -112,7 +112,7 @@ VLLM_ARGS=( # sometimes flashinfer has IMA bugs. Use fa3 as instead --vllm-attention-backend FLASH_ATTN - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) MISC_ARGS=( diff --git a/scripts/run-qwen3-next-80B-A3B.sh b/scripts/run-qwen3-next-80B-A3B.sh index 10d811cc5..74cc77854 100644 --- a/scripts/run-qwen3-next-80B-A3B.sh +++ b/scripts/run-qwen3-next-80B-A3B.sh @@ -132,7 +132,7 @@ VLLM_ARGS=( --vllm-max-num-seqs 256 --vllm-enable-expert-parallel --vllm-cudagraph-capture-sizes 1 2 4 8 $(seq 16 8 128) - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) MISC_ARGS=( diff --git a/scripts/run-qwen3.5-27B.sh b/scripts/run-qwen3.5-27B.sh index a9be251da..ed368630a 100755 --- a/scripts/run-qwen3.5-27B.sh +++ b/scripts/run-qwen3.5-27B.sh @@ -126,7 +126,7 @@ WANDB_ARGS=( VLLM_ARGS=( --rollout-num-gpus-per-engine 2 --vllm-gpu-memory-utilization 0.75 - --vllm-speculative-config '{"method":"eagle","num_speculative_tokens":3}' + --vllm-speculative-config '{"method":"mtp","num_speculative_tokens":3}' ) MISC_ARGS=( diff --git a/slime/backends/megatron_utils/fp8_helpers.py b/slime/backends/megatron_utils/fp8_helpers.py new file mode 100644 index 000000000..14014753c --- /dev/null +++ b/slime/backends/megatron_utils/fp8_helpers.py @@ -0,0 +1,74 @@ +"""FP8 / UE8M0 quantization helpers for megatron → vLLM weight transfer. + +All symbols fall back to ``None`` when vLLM's deep_gemm helpers are not +available, which disables the UE8M0 requantization path. +""" + +import torch + +# deep_gemm is a third-party DeepSeek MoE GEMM library; not all images ship it. +# When missing, the UE8M0 requantization path falls back to None so callers can +# skip it gracefully (matches the module docstring contract). +try: + import deep_gemm.utils.layout as _deep_gemm_layout + from vllm.utils.deep_gemm import get_tma_aligned_size as _get_tma_aligned_size + _HAS_DEEP_GEMM = True +except ImportError: + _deep_gemm_layout = None + _get_tma_aligned_size = None + _HAS_DEEP_GEMM = False + +try: + from vllm.utils.deep_gemm import is_deep_gemm_e8m0_used as _vllm_is_e8m0 + from vllm.utils.deep_gemm import per_block_cast_to_fp8 as _vllm_per_block_cast +except ImportError: + _vllm_is_e8m0 = lambda: False # noqa: E731 + _vllm_per_block_cast = None + + +def should_deepgemm_weight_requant_ue8m0(weight_block_size) -> bool: + return weight_block_size is not None and _vllm_is_e8m0() + + +def quant_weight_ue8m0( + weight_dequant: torch.Tensor, + weight_block_size: list[int], +): + assert weight_block_size == [128, 128] + assert weight_dequant.dtype == torch.bfloat16, f"{weight_dequant.dtype=} {weight_dequant.shape=}" + *batch_dims, n, k = weight_dequant.shape + flat = weight_dequant.view(-1, k) + out_w_flat, out_s_flat = _vllm_per_block_cast(flat, block_size=[128, 128], use_ue8m0=True) + out_w = out_w_flat.view(*batch_dims, n, k) + from math import ceil + + out_s = out_s_flat.view( + *batch_dims, + ceil(n / weight_block_size[0]), + ceil(k / weight_block_size[1]), + ) + return out_w, out_s + + +def transform_scale_ue8m0(sf: torch.Tensor, mn: int, use_torch_impl: bool = False): + if _deep_gemm_layout is None: + raise RuntimeError("deep_gemm not installed; UE8M0 requantization unavailable.") + get_fn = _deep_gemm_layout.get_mn_major_tma_aligned_packed_ue8m0_tensor + sf = sf.index_select(-2, torch.arange(mn, device=sf.device) // 128) + sf = get_fn(sf) + if sf.shape[-1] == 1: + get_tma_aligned_size = _get_tma_aligned_size # pre-imported with fallback + + aligned_mn = get_tma_aligned_size(sf.shape[-2], sf.element_size()) + if sf.stride(-1) != aligned_mn: + new_stride = list(sf.stride()) + new_stride[-1] = aligned_mn + sf = sf.as_strided(sf.shape, tuple(new_stride)) + return sf + + +__all__ = [ + "quant_weight_ue8m0", + "transform_scale_ue8m0", + "should_deepgemm_weight_requant_ue8m0", +] diff --git a/slime/backends/megatron_utils/megatron_to_hf/__init__.py b/slime/backends/megatron_utils/megatron_to_hf/__init__.py index d6cccc23f..b51cb1b57 100644 --- a/slime/backends/megatron_utils/megatron_to_hf/__init__.py +++ b/slime/backends/megatron_utils/megatron_to_hf/__init__.py @@ -28,10 +28,6 @@ def convert_to_hf(args, model_name, name, param, quantization_config=None): return quantize_params(args, name, converted_named_tensors, quantization_config) -# TODO optimize -_cached_tensors = {} - - # TODO optimize code details def _convert_to_hf_core(args, model_name, name, param): if "glm4moelite" in model_name or "deepseekv3" in model_name: @@ -59,31 +55,4 @@ def _convert_to_hf_core(args, model_name, name, param): else: raise ValueError(f"Unsupported model: {model_name}") - # to compatible with sglang implementation - if args.q_lora_rank is not None: - old_converted_named_tensors = converted_named_tensors - converted_named_tensors = [] - for converted_name, converted_param in old_converted_named_tensors: - if "q_a_proj" in converted_name: - pair_name = converted_name.replace("q_a_proj", "kv_a_proj_with_mqa") - if pair_name in _cached_tensors: - converted_named_tensors += [ - (converted_name, converted_param), - (pair_name, _cached_tensors[pair_name]), - ] - del _cached_tensors[pair_name] - else: - _cached_tensors[converted_name] = converted_param - elif "kv_a_proj_with_mqa" in converted_name: - pair_name = converted_name.replace("kv_a_proj_with_mqa", "q_a_proj") - if pair_name in _cached_tensors: - converted_named_tensors += [ - (converted_name, converted_param), - (pair_name, _cached_tensors[pair_name]), - ] - del _cached_tensors[pair_name] - else: - _cached_tensors[converted_name] = converted_param - else: - converted_named_tensors.append((converted_name, converted_param)) return converted_named_tensors diff --git a/slime/backends/megatron_utils/megatron_to_hf/gpt_oss.py b/slime/backends/megatron_utils/megatron_to_hf/gpt_oss.py index b90507f04..0db77f912 100644 --- a/slime/backends/megatron_utils/megatron_to_hf/gpt_oss.py +++ b/slime/backends/megatron_utils/megatron_to_hf/gpt_oss.py @@ -4,7 +4,7 @@ def convert_gpt_oss_to_hf(args, name, param): - """Convert Megatron GPT-OSS parameter names to HF format for weight update to SGLang.""" + """Convert Megatron GPT-OSS parameter names to HF format for weight update to vLLM.""" if name == "module.module.embedding.word_embeddings.weight": return [("model.embed_tokens.weight", param)] diff --git a/slime/backends/megatron_utils/megatron_to_hf/processors/quantizer_fp8.py b/slime/backends/megatron_utils/megatron_to_hf/processors/quantizer_fp8.py index f97cfd693..36d26b705 100644 --- a/slime/backends/megatron_utils/megatron_to_hf/processors/quantizer_fp8.py +++ b/slime/backends/megatron_utils/megatron_to_hf/processors/quantizer_fp8.py @@ -4,7 +4,7 @@ from slime.backends.megatron_utils.kernels.fp8_kernel import blockwise_cast_to_fp8_triton -from ...sglang import quant_weight_ue8m0, should_deepgemm_weight_requant_ue8m0, transform_scale_ue8m0 +from ...fp8_helpers import quant_weight_ue8m0, should_deepgemm_weight_requant_ue8m0, transform_scale_ue8m0 def quantize_params_fp8(args, megatron_name, converted_named_params, quantization_config): diff --git a/slime/backends/megatron_utils/sglang.py b/slime/backends/megatron_utils/sglang.py deleted file mode 100644 index 44cd66fe0..000000000 --- a/slime/backends/megatron_utils/sglang.py +++ /dev/null @@ -1,24 +0,0 @@ -# the file to manage all sglang deps in the megatron actor -try: - from sglang.srt.layers.quantization.fp8_utils import quant_weight_ue8m0, transform_scale_ue8m0 - from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0 -except ImportError: - quant_weight_ue8m0 = None - transform_scale_ue8m0 = None - should_deepgemm_weight_requant_ue8m0 = None - -from sglang.srt.utils import MultiprocessingSerializer - - -try: - from sglang.srt.weight_sync.tensor_bucket import FlattenedTensorBucket # type: ignore[import] -except ImportError: - from sglang.srt.model_executor.model_runner import FlattenedTensorBucket # type: ignore[import] - -__all__ = [ - "quant_weight_ue8m0", - "transform_scale_ue8m0", - "should_deepgemm_weight_requant_ue8m0", - "MultiprocessingSerializer", - "FlattenedTensorBucket", -] diff --git a/slime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py b/slime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py index 9ea728697..fa6f464e2 100644 --- a/slime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py +++ b/slime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py @@ -14,7 +14,7 @@ def _patch_bridge_expert_cache_to_cpu(): """Monkey-patch GPTOSSBridge class to cache expert weights on CPU. This avoids GPU OOM when torch.cat merges all experts, especially in - colocated mode where SGLang and Megatron share the same GPU. + colocated mode where vLLM and Megatron share the same GPU. """ try: from megatron.bridge.models.gpt_oss.gpt_oss_bridge import GPTOSSBridge diff --git a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py index 150ea1573..b3a230863 100644 --- a/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py +++ b/slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py @@ -80,7 +80,7 @@ def connect_rollout_engines( # For TP: # 1. AllGather parameters to rank 0 - # 2. Broadcast parameters from rank 0 to all sglang engines + # 2. Broadcast parameters from rank 0 to all vLLM engines self._is_pp_src_rank = ( mpu.get_data_parallel_rank(with_context_parallel=True) == 0 and mpu.get_tensor_model_parallel_rank() == 0 ) diff --git a/slime/backends/sglang_utils/__init__.py b/slime/backends/sglang_utils/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/slime/backends/sglang_utils/arguments.py b/slime/backends/sglang_utils/arguments.py deleted file mode 100644 index 0a4801743..000000000 --- a/slime/backends/sglang_utils/arguments.py +++ /dev/null @@ -1,197 +0,0 @@ -import argparse - -from sglang.srt.server_args import ServerArgs -from slime.utils.http_utils import _wrap_ipv6 - - -# TODO: use all sglang router arguments with `--sglang-router` prefix -def add_sglang_router_arguments(parser): - """ - Add arguments to the parser for the SGLang router. - """ - parser.add_argument( - "--sglang-router-ip", - type=str, - default=None, - help="IP address of the SGLang router", - ) - parser.add_argument( - "--sglang-router-port", - type=int, - default=None, - help="Port of the SGLang router", - ) - parser.add_argument( - "--sglang-router-request-timeout-secs", - type=int, - default=14400, - help="Timeout for requests to the SGLang router in seconds", - ) - return parser - - -def add_sglang_arguments(parser): - """ - Add arguments to the parser for the SGLang server. - """ - parser = add_sglang_router_arguments(parser) - parser.set_defaults(router_balance_abs_threshold=10, router_balance_rel_threshold=1.2) - parser.add_argument("--sglang-server-concurrency", type=int, default=512) - - old_add_argument = parser.add_argument - - skipped_args = [ - "model_path", - "config", - "trust_remote_code", - "random_seed", - # memory - "enable_memory_saver", - # distributed - "tp_size", - "port", - "nnodes", - "node_rank", - "dist_init_addr", - "gpu_id_step", - "base_gpu_id", - "nccl_port", - "skip_server_warmup", - "enable_return_routed_experts", - ] - - def new_add_argument_wrapper(*name_or_flags, **kwargs): - """ - Add arguments to the parser, ensuring that the server arguments are prefixed and skippable. - """ - # Determine the canonical name for skip check (e.g., "model_path") - canonical_name_for_skip_check = None - if "dest" in kwargs: - canonical_name_for_skip_check = kwargs["dest"] - else: - for flag_name_candidate in name_or_flags: - if isinstance(flag_name_candidate, str) and flag_name_candidate.startswith("--"): - # Derive from first long flag: --foo-bar -> foo_bar - stem = flag_name_candidate[2:] - canonical_name_for_skip_check = stem.replace("-", "_") - break - # If no long flag and no dest, skip logic might not catch it unless short flags imply a dest. - - if canonical_name_for_skip_check and canonical_name_for_skip_check in skipped_args: - return # Skip this entire argument definition - - # If not skipped, proceed to prefix flags and dest - new_name_or_flags_list = [] - for item_flag in name_or_flags: - if isinstance(item_flag, str) and item_flag.startswith("-"): - original_flag_stem = item_flag.lstrip("-") # "foo-bar" from "--foo-bar", or "f" from "-f" - prefixed_item = f"--sglang-{original_flag_stem}" - new_name_or_flags_list.append(prefixed_item) - else: - # Positional arguments or non-string items - new_name_or_flags_list.append(item_flag) - - # Prepare kwargs for the actual add_argument call. - # Make a copy to avoid modifying the original kwargs dict. - final_kwargs = kwargs.copy() - - # If 'dest' is explicitly provided and is a string, prefix it. - # This ensures the attribute on the args namespace becomes, e.g., args.sglang_dest_name. - if "dest" in final_kwargs and isinstance(final_kwargs["dest"], str): - original_dest = final_kwargs["dest"] - # Avoid double prefixing if dest somehow already starts with sglang_ - if not original_dest.startswith("sglang_"): - final_kwargs["dest"] = f"sglang_{original_dest}" - # If 'dest' is not explicitly provided (or is None/not a string), - # argparse will derive 'dest' from the (now prefixed) flag names. - # E.g., if the first flag is "--sglang-foo-bar", argparse sets dest to "sglang_foo_bar". - - old_add_argument(*new_name_or_flags_list, **final_kwargs) - - parser.add_argument = new_add_argument_wrapper - ServerArgs.add_cli_args(parser) - parser.add_argument = old_add_argument - - # PD disaggregation / multi-group config - parser.add_argument( - "--prefill-num-servers", - type=int, - default=None, - help="Number of prefill servers for disaggregation.", - ) - parser.add_argument( - "--sglang-config", - type=str, - default=None, - help=( - "Path to a YAML config for SGLang engine deployment. " - "Defines server_groups with worker_type (regular/prefill/decode/placeholder), " - "num_gpus per group, and optional per-group 'overrides' dict of " - "ServerArgs field names that override the base --sglang-* CLI args. " - "Placeholder groups reserve GPU slots without creating engines. " - "Mutually exclusive with --prefill-num-servers." - ), - ) - - return parser - - -def validate_args(args): - args.sglang_dp_size = args.sglang_data_parallel_size - args.sglang_pp_size = args.sglang_pipeline_parallel_size - args.sglang_ep_size = args.sglang_expert_parallel_size - - # Compute effective TP size considering PP size - if args.sglang_pp_size > 1: - assert args.rollout_num_gpus_per_engine % args.sglang_pp_size == 0, ( - f"rollout_num_gpus_per_engine ({args.rollout_num_gpus_per_engine}) must be divisible by " - f"sglang_pipeline_parallel_size ({args.sglang_pp_size})" - ) - args.sglang_tp_size = args.rollout_num_gpus_per_engine // args.sglang_pp_size - else: - args.sglang_tp_size = args.rollout_num_gpus_per_engine - - if args.sglang_dp_size > 1: - assert args.sglang_enable_dp_attention - - if getattr(args, "sglang_router_ip", None): - args.sglang_router_ip = _wrap_ipv6(args.sglang_router_ip) - - # Mutual-exclusion checks for PD disaggregation / sglang-config. - assert not ( - getattr(args, "prefill_num_servers", None) is not None and args.rollout_external - ), "prefill_num_servers cannot be set when rollout_external is set." - - assert not ( - getattr(args, "sglang_config", None) is not None and args.rollout_external - ), "sglang_config cannot be set when rollout_external is set." - - assert not ( - getattr(args, "sglang_config", None) is not None and getattr(args, "prefill_num_servers", None) is not None - ), "sglang_config and prefill_num_servers are mutually exclusive. Use server_groups in the YAML config instead." - - -def sglang_parse_args(): - """ - Parse sglang server arguments independently using a separate ArgumentParser. - Uses parse_known_args() to only consume sglang-related arguments from sys.argv, - allowing the remaining arguments to be parsed by megatron separately. - - Returns: - argparse.Namespace: Parsed sglang arguments (all attributes prefixed with sglang_). - """ - parser = argparse.ArgumentParser(add_help=False) - add_sglang_arguments(parser) - - # Compute default sglang_tensor_parallel_size from CLI args - temp_parser = argparse.ArgumentParser(add_help=False) - temp_parser.add_argument("--rollout-num-gpus-per-engine", type=int, default=1) - temp_parser.add_argument("--sglang-pp-size", type=int, default=1) - temp_parser.add_argument("--sglang-pipeline-parallel-size", type=int, default=1) - temp_args, _ = temp_parser.parse_known_args() - pp_size = temp_args.sglang_pp_size if temp_args.sglang_pp_size != 1 else temp_args.sglang_pipeline_parallel_size - sglang_tp_size = temp_args.rollout_num_gpus_per_engine // pp_size - parser.set_defaults(sglang_tensor_parallel_size=sglang_tp_size) - - args, _ = parser.parse_known_args() - return args diff --git a/slime/backends/sglang_utils/sglang_engine.py b/slime/backends/sglang_utils/sglang_engine.py deleted file mode 100644 index c4465b3c9..000000000 --- a/slime/backends/sglang_utils/sglang_engine.py +++ /dev/null @@ -1,621 +0,0 @@ -import dataclasses -import ipaddress -import logging -import multiprocessing -import os -import time -from urllib.parse import quote - -import requests -import sglang_router -from packaging.version import parse -from sglang.srt.server_args import ServerArgs -from sglang.srt.utils import kill_process_tree -from urllib3.exceptions import NewConnectionError - -from slime.ray.ray_actor import RayActor -from slime.utils.http_utils import get_host_info - -logger = logging.getLogger(__name__) - - -def get_base_gpu_id(args, rank): - num_gpus = min(args.num_gpus_per_node, args.rollout_num_gpus_per_engine) - if args.colocate: - start_index = (rank * num_gpus) % args.num_gpus_per_node - else: - num_actor_gpus = 0 if args.debug_rollout_only else args.actor_num_gpus_per_node * args.actor_num_nodes - start_index = (num_actor_gpus + rank * num_gpus) % args.num_gpus_per_node - return start_index - - -def _to_local_gpu_id(physical_gpu_id: int) -> int: - cvd = os.environ.get("CUDA_VISIBLE_DEVICES") - if not cvd: - return physical_gpu_id # no remapping - # CUDA_VISIBLE_DEVICES can be like "4,5,6,7" - visible = [int(x) for x in cvd.split(",") if x.strip() != ""] - # In a remapped process, valid torch device indices are 0..len(visible)-1 - if physical_gpu_id in visible: - return visible.index(physical_gpu_id) - # If we're already getting local IDs, allow them - if 0 <= physical_gpu_id < len(visible): - return physical_gpu_id - raise RuntimeError( - f"GPU id {physical_gpu_id} is not valid under CUDA_VISIBLE_DEVICES={cvd}. " - f"Expected one of {visible} (physical) or 0..{len(visible)-1} (local)." - ) - - -def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process: - if getattr(server_args, "encoder_only", False): - from sglang.srt.disaggregation.encode_server import launch_server_process as sglang_launch_server_process - - return sglang_launch_server_process( - server_args, - start_method="spawn", - wait_for_server=True, - ) - - from sglang.srt.entrypoints.http_server import launch_server - - multiprocessing.set_start_method("spawn", force=True) - server_args.host = server_args.host.strip("[]") - p = multiprocessing.Process(target=launch_server, args=(server_args,)) - p.start() - - if getattr(server_args, "node_rank", 0) != 0: - return p - - _wait_server_healthy( - base_url=server_args.url(), - api_key=server_args.api_key, - is_process_alive=lambda: p.is_alive(), - ) - - return p - - -def _wait_server_healthy(base_url, api_key, is_process_alive): - headers = { - "Content-Type": "application/json; charset=utf-8", - "Authorization": f"Bearer {api_key}", - } - - with requests.Session() as session: - while True: - try: - response = session.get(f"{base_url}/health_generate", headers=headers) - if response.status_code == 200: - break - except requests.RequestException: - pass - - if not is_process_alive(): - raise Exception("Server process terminated unexpectedly.") - - time.sleep(2) - - -class SGLangEngine(RayActor): - def __init__( - self, - args, - rank: int, - worker_type: str = "regular", - base_gpu_id: int | None = None, - sglang_overrides: dict | None = None, - num_gpus_per_engine: int | None = None, - ): - self.args = args - self.rank = rank - self.worker_type = worker_type - self.base_gpu_id = base_gpu_id - self.sglang_overrides = sglang_overrides or {} - self.num_gpus_per_engine = num_gpus_per_engine - - def init( - self, - dist_init_addr, - port, - nccl_port, - host=None, - disaggregation_bootstrap_port=None, - router_ip=None, - router_port=None, - ): - self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip - self.router_port = router_port if router_port is not None else self.args.sglang_router_port - - host = host or get_host_info()[1] - - def _format_v6_uri(addr): - if not addr or addr.startswith("["): - return addr - try: - if ipaddress.ip_address(addr).version == 6: - return f"[{addr}]" - except ValueError: - pass - return addr - - host = _format_v6_uri(host) - ip_part, port_part = dist_init_addr.rsplit(":", 1) - dist_init_addr = f"{_format_v6_uri(ip_part)}:{port_part}" - - server_args_dict, external_engine_need_check_fields = _compute_server_args( - self.args, - self.rank, - dist_init_addr, - nccl_port, - host, - port, - self.worker_type, - disaggregation_bootstrap_port, - base_gpu_id=self.base_gpu_id, - sglang_overrides=self.sglang_overrides, - num_gpus_per_engine=self.num_gpus_per_engine, - ) - - self.node_rank = server_args_dict["node_rank"] - self.server_host = server_args_dict["host"] # with [] if ipv6 - self.server_port = server_args_dict["port"] - - if self.args.rollout_external: - self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields) - else: - self._init_normal(server_args_dict) - - def _init_external(self, expect_server_args, external_engine_need_check_fields): - logger.info(f"Use external SGLang engine (rank={self.rank}, expect_server_args={expect_server_args})") - - def _get_actual_server_args(): - response = requests.get(f"http://{self.server_host}:{self.server_port}/get_server_info") - response.raise_for_status() - return response.json() - - def _sanity_check_server_args(actual_server_args, expect_server_args): - for name in external_engine_need_check_fields: - expect_value = expect_server_args.get(name) - actual_value = actual_server_args.get(name) - assert ( - actual_value == expect_value - ), f"{name=} {expect_value=} {actual_value=} {expect_server_args=} {actual_server_args=}" - - _wait_server_healthy( - base_url=f"http://{self.server_host}:{self.server_port}", - api_key=None, - is_process_alive=lambda: True, - ) - actual_server_args = _get_actual_server_args() - _sanity_check_server_args(actual_server_args, expect_server_args) - - def _init_normal(self, server_args_dict): - logger.info(f"Launch HttpServerEngineAdapter at: {self.server_host}:{self.server_port}") - self.process = launch_server_process(ServerArgs(**server_args_dict)) - - if self.worker_type == "encoder": - return - - if self.node_rank == 0 and self.router_ip and self.router_port: - if parse(sglang_router.__version__) <= parse("0.2.1"): - assert self.worker_type == "regular", "pd disaggregation is not supported in old router." - response = requests.post( - f"http://{self.router_ip}:{self.router_port}/add_worker?url=http://{self.server_host}:{self.server_port}", - ) - else: - payload = { - "url": f"http://{self.server_host}:{self.server_port}", - "worker_type": self.worker_type, - } - if self.worker_type == "prefill": - payload["bootstrap_port"] = server_args_dict["disaggregation_bootstrap_port"] - response = requests.post( - f"http://{self.router_ip}:{self.router_port}/workers", - json=payload, - ) - response.raise_for_status() - - def _make_request(self, endpoint: str, payload: dict | None = None): - """Make a POST request to the specified endpoint with the given payload. - - Args: - endpoint: The API endpoint to call - payload: The JSON payload to send (default: empty dict) - - Returns: - The JSON response from the server - """ - if self.node_rank != 0: - return - - url = f"http://{self.server_host}:{self.server_port}/{endpoint}" - response = requests.post(url, json=payload or {}) - try: - response.raise_for_status() - except requests.exceptions.HTTPError as e: - e.add_note(f"{response.text=}") - raise - return response.json() - - def health_generate(self, timeout: float = 5.0) -> bool: - """Run /health_generate on the underlying SGLang HTTP server. - - Args: - timeout: Timeout for the health request in seconds. - - Returns: - True if the server responds with HTTP 200. - - Raises: - requests.RequestException: If the request fails for any reason, including timeout. - """ - if self.node_rank != 0: - return True - - response = requests.get( - f"http://{self.server_host}:{self.server_port}/health_generate", - timeout=timeout, - ) - response.raise_for_status() - return True - - def update_weights_from_tensor( - self, - serialized_named_tensors: list[str], - load_format: str | None = None, - flush_cache: bool = False, - weight_version: str | None = None, - ): - """ - Update model weights from tensor data. The HTTP server will only post meta data, and the real weights will be copied directly from GPUs. - - Note: The model should be on GPUs rather than CPU for this functionality to work properly. - If you encounter issues, ensure your model is loaded on GPU devices rather than CPU. - """ - payload = { - "serialized_named_tensors": serialized_named_tensors, - "load_format": load_format, - "flush_cache": flush_cache, - } - if weight_version is not None: - payload["weight_version"] = weight_version - return self._make_request( - "update_weights_from_tensor", - payload, - ) - - def flush_cache(self): - """Flush the cache of the server.""" - if self.node_rank != 0: - return - # flush cache will not return status_code 200 when there are pending requests - for _ in range(60): - try: - response = requests.get(f"http://{self.server_host}:{self.server_port}/flush_cache") - if response.status_code == 200: - break - except NewConnectionError as e: - raise e - except Exception as e: - logger.info(f"Error flushing cache: {e}") - time.sleep(1) - continue - else: - raise TimeoutError("Timeout while flushing cache.") - - def get_url(self): - if self.node_rank != 0: - return None - return f"http://{self.server_host}:{self.server_port}" - - def shutdown(self): - if self.args.rollout_external: - return - - logger.info(f"Shutdown engine {self.server_host}:{self.server_port}...") - if self.worker_type != "encoder" and self.node_rank == 0: - worker_url = f"http://{self.server_host}:{self.server_port}" - response = None - if parse(sglang_router.__version__) <= parse("0.2.1"): - response = requests.post( - f"http://{self.router_ip}:{self.router_port}/remove_worker?url=http://{self.server_host}:{self.server_port}" - ) - elif parse(sglang_router.__version__) < parse("0.3.0"): - worker_url = quote(worker_url, safe="") - response = requests.delete(f"http://{self.router_ip}:{self.router_port}/workers/{worker_url}") - else: - try: - all_workers = requests.get(f"http://{self.router_ip}:{self.router_port}/workers").json()["workers"] - for worker in all_workers: - if worker["url"] == worker_url: - worker_id = worker["id"] - response = requests.delete( - f"http://{self.router_ip}:{self.router_port}/workers/{worker_id}" - ) - break - else: - logger.warning(f"Worker {worker_url} not found in router during shutdown.") - except Exception as e: - logger.warning(f"Failed to fetch workers list or remove worker: {e}") - - if response is not None: - response.raise_for_status() - kill_process_tree(self.process.pid) - - def get_weight_version(self): - if self.node_rank != 0: - return - url = f"http://{self.server_host}:{self.server_port}/get_weight_version" - response = requests.get(url) - response.raise_for_status() - return response.json()["weight_version"] - - def release_memory_occupation(self): - self.flush_cache() - return self._make_request("release_memory_occupation") - - def resume_memory_occupation(self, tags: list[str] = None): - """ - Available tags for multi-stage resume: weights, kv_cache - """ - return self._make_request( - "resume_memory_occupation", - {"tags": tags}, - ) - - def check_weights(self, action: str): - return self._make_request("weights_checker", {"action": action}) - - def update_weights_from_disk(self, model_path: str, load_format: str | None = None): - """Reload weights from *model_path* without restarting the engine. - - Used for non-updatable (frozen) models that overlap with megatron: - after offload, weights are restored from disk instead of CPU cache. - """ - payload = {"model_path": model_path} - if load_format is not None: - payload["load_format"] = load_format - return self._make_request("update_weights_from_disk", payload) - - def init_weights_update_group(self, master_address, master_port, rank_offset, world_size, group_name, backend): - return self._make_request( - "init_weights_update_group", - { - "master_address": master_address, - "master_port": master_port, - "rank_offset": rank_offset, - "world_size": world_size, - "group_name": group_name, - "backend": backend, - }, - ) - - def destroy_weights_update_group(self, group_name): - try: - return self._make_request( - "destroy_weights_update_group", - { - "group_name": group_name, - }, - ) - except requests.exceptions.RequestException: - # catch the case there the engine is just created and does not have the group. - pass - - def update_weights_from_distributed( - self, - names, - dtypes, - shapes, - group_name, - flush_cache=False, - weight_version: str | None = None, - packed: bool = False, - ): - del packed - payload = { - "names": names, - "dtypes": [str(dtype).replace("torch.", "") for dtype in dtypes], - "shapes": shapes, - "group_name": group_name, - "flush_cache": flush_cache, - } - if weight_version is not None: - payload["weight_version"] = weight_version - return self._make_request( - "update_weights_from_distributed", - payload, - ) - - def pause_generation(self): - response = requests.post(f"http://{self.server_host}:{self.server_port}/pause_generation", json={}) - response.raise_for_status() - return response - - def continue_generation(self): - response = requests.post(f"http://{self.server_host}:{self.server_port}/continue_generation", json={}) - response.raise_for_status() - return response - - def post_process_weights( - self, - restore_weights_before_load: bool = False, - post_process_quantization: bool = False, - ): - """ - Update model weights from tensor data. The HTTP server will only post meta data, and the real weights will be copied directly from GPUs. - Note: The model should be on GPUs rather than CPU for this functionality to work properly. - If you encounter issues, ensure your model is loaded on GPU devices rather than CPU. - """ - - return self._make_request( - "post_process_weights", - { - "restore_weights_before_load": restore_weights_before_load, - "post_process_quantization": post_process_quantization, - }, - ) - - def start_profile( - self, - # The output directory - output_dir: str | None = None, - # If set, it profile as many as this number of steps. - # If it is set, profiling is automatically stopped after this step, and - # the caller doesn't need to run stop_profile. - start_step: int | None = None, - num_steps: int | None = None, - activities: list[str] | None = None, - profile_by_stage: bool = False, - with_stack: bool | None = None, - record_shapes: bool | None = None, - ): - response = requests.post( - f"http://{self.server_host}:{self.server_port}/start_profile", - json={ - "output_dir": output_dir, - "start_step": start_step, - "num_steps": num_steps, - "activities": activities, - "profile_by_stage": profile_by_stage, - "with_stack": with_stack, - "record_shapes": record_shapes, - }, - ) - response.raise_for_status() - return response - - def stop_profile(self): - response = requests.post(f"http://{self.server_host}:{self.server_port}/stop_profile", json={}) - response.raise_for_status() - return response - - def simulate_crash(self): - if self.args.rollout_external or not getattr(self, "process", None): - logger.info( - "simulate_crash called but no local engine process exists (rollout_external=%s); skip kill", - self.args.rollout_external, - ) - return - - logger.info(f"Simulating crash on engine {self.server_host}:{self.server_port}...") - self.shutdown() - - -def _compute_server_args( - args, - rank, - dist_init_addr, - nccl_port, - host, - port, - worker_type: str = "regular", - disaggregation_bootstrap_port: int | None = None, - base_gpu_id: int | None = None, - sglang_overrides: dict | None = None, - num_gpus_per_engine: int | None = None, -): - _gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine - nnodes = max(1, _gpus_per_engine // args.num_gpus_per_node) - node_rank = rank % nnodes - base = base_gpu_id if base_gpu_id is not None else get_base_gpu_id(args, rank) - base = _to_local_gpu_id(base) - kwargs = { - "model_path": args.hf_checkpoint, - "trust_remote_code": True, - "random_seed": args.seed + rank, - # memory - "enable_memory_saver": args.offload_rollout, - # distributed - "host": host, - "port": port, - "nccl_port": nccl_port, - "nnodes": nnodes, - "node_rank": node_rank, - "dist_init_addr": dist_init_addr, - "gpu_id_step": 1, - "base_gpu_id": base, - # parallel - "tp_size": _gpus_per_engine // args.sglang_pp_size, - "dp_size": args.sglang_dp_size, - "pp_size": args.sglang_pp_size, - "ep_size": args.sglang_ep_size, - # always skip warmup to prevent warmup timeout. - "skip_server_warmup": True, - # always enable draft weights cpu backup so that we run training without mtp weights. - "enable_draft_weights_cpu_backup": True, - # Always enable Prometheus metrics so the /engine_metrics endpoint is - # available for W&B scraping (regardless of --sglang-enable-metrics). - "enable_metrics": True, - } - - if worker_type == "prefill": - kwargs["disaggregation_mode"] = "prefill" - kwargs["load_balance_method"] = "follow_bootstrap_room" - assert ( - disaggregation_bootstrap_port is not None - ), "disaggregation_bootstrap_port must be set for prefill worker" - kwargs["disaggregation_bootstrap_port"] = disaggregation_bootstrap_port - elif worker_type == "decode": - kwargs["disaggregation_mode"] = "decode" - kwargs["prefill_round_robin_balance"] = True - elif worker_type == "encoder": - kwargs["encoder_only"] = True - - if args.use_rollout_routing_replay: - kwargs["enable_return_routed_experts"] = True - if args.fp16: - kwargs["dtype"] = "float16" - external_engine_need_check_fields = [k for k in kwargs.keys() if k not in _EXTERNAL_ENGINE_SKIP_CHECK_FIELDS] - - server_arg_fields = dataclasses.fields(ServerArgs) - server_arg_field_names = {attr.name for attr in server_arg_fields} - unused_keys = set(kwargs.keys()) - for attr in server_arg_fields: - if worker_type == "decode" and attr.name == "enable_hierarchical_cache": - continue - if hasattr(args, f"sglang_{attr.name}") and attr.name not in kwargs: - kwargs[attr.name] = getattr(args, f"sglang_{attr.name}") - unused_keys.discard(attr.name) - - # Per-server-group overrides from --sglang-config YAML. - # Applied after base args so they take highest priority. - if sglang_overrides: - for key, value in sglang_overrides.items(): - normalized_key = key.replace("-", "_") - if normalized_key != key: - logger.warning( - f"sglang_overrides key '{key}' normalized to '{normalized_key}' (rank={rank}). " - "Please use underscore style in YAML overrides." - ) - if normalized_key in kwargs: - logger.info( - f"sglang_overrides: overriding {normalized_key}={kwargs[normalized_key]} -> {value} (rank={rank})" - ) - kwargs[normalized_key] = value - if normalized_key in server_arg_field_names: - unused_keys.discard(normalized_key) - else: - unused_keys.add(normalized_key) - - # for compatibility with old args - if len(unused_keys) > 0: - logger.info(f"Warning: The following arguments is not supported in the current sglang: {unused_keys}.") - for key in unused_keys: - kwargs.pop(key) - - return kwargs, external_engine_need_check_fields - - -_EXTERNAL_ENGINE_SKIP_CHECK_FIELDS = [ - "model_path", - "trust_remote_code", - "random_seed", - "nccl_port", - "dist_init_addr", - "skip_server_warmup", - "enable_draft_weights_cpu_backup", - "enable_metrics", - "mem_fraction_static", -] diff --git a/slime/backends/vllm_utils/arguments.py b/slime/backends/vllm_utils/arguments.py index 68290e782..c65f76927 100644 --- a/slime/backends/vllm_utils/arguments.py +++ b/slime/backends/vllm_utils/arguments.py @@ -1,6 +1,5 @@ """vLLM rollout backend argument definitions. -Mirrors slime/backends/sglang_utils/arguments.py: - Wholesale-imports ``AsyncEngineArgs.add_cli_args(parser)`` with a wrapper that prefixes every flag with ``--vllm-`` and every dest with ``vllm_``. - Adds a small set of vime-specific orchestration extras (router endpoint, @@ -77,8 +76,7 @@ def _detect_user_provided_dests(parser, argv: list[str]) -> tuple[set[str], dict # Dests already managed at vime / megatron level (orchestrator decides them) -# or non-applicable to subprocess `vllm serve` mode. Same intent as -# sglang_utils/arguments.py:43-61 `skipped_args`. +# or non-applicable to subprocess `vllm serve` mode. SKIPPED_DESTS = [ # model identity: hf_checkpoint owns this "model", @@ -97,7 +95,7 @@ def _detect_user_provided_dests(parser, argv: list[str]) -> tuple[set[str], dict # tp_size is fully owned by the orchestrator (rollout_num_gpus_per_engine # // pp_size) — see validate_args. pipeline_parallel_size and # data_parallel_size remain user-controllable and auto-forward to the vllm - # subprocess when set; mirrors slime sglang_utils which exposes both via CLI. + # subprocess when set. "tensor_parallel_size", # network: engine launcher decides per-engine port/host "port", @@ -108,30 +106,47 @@ def _detect_user_provided_dests(parser, argv: list[str]) -> tuple[set[str], dict def add_vllm_router_arguments(parser): - """vime's router orchestration flags (where to reach the router; not in vllm-router's CLI). - - Named without a ``--vllm-`` backend prefix so they sit alongside - ``vllm_router.RouterArgs`` knobs (``--router-policy``, ``--router-cache-threshold``, - …) under a single ``--router-*`` namespace. PR #10 hardcodes a single vllm backend, - so backend-disambiguation prefix would only leak implementation. - """ + """vime's vllm-router orchestration flags (where to reach the router; not in vllm-router's CLI).""" parser.add_argument( - "--router-ip", + "--vllm-router-ip", type=str, default=None, - help="IP address of the router (where vime connects to send rollout requests).", + help="IP address of the vllm router (where vime connects to send rollout requests).", ) parser.add_argument( - "--router-port", + "--vllm-router-port", type=int, default=None, - help="Port of the router.", + help="Port of the vllm router.", ) + # Bare ``--router-request-timeout-secs`` (dest ``router_request_timeout_secs``): + # this is a genuine vllm-router knob (a RouterArgs field), so it shares the + # ``--router-*`` namespace with policy / cache_threshold / retries / … rather + # than vime's ``--vllm-router-*`` endpoint flags. Only --vllm-router-ip and + # --vllm-router-port keep the vllm_ prefix (RouterArgs excludes host/port from + # its CLI via exclude_host_port=True, so vime owns those two outright). parser.add_argument( "--router-request-timeout-secs", type=int, default=14400, - help="Timeout (seconds) for HTTP requests vime makes to the router.", + help="Timeout (seconds) for HTTP requests vime makes to the vllm router.", + ) + # dest is ``router_policy`` (NOT ``vllm_router_policy``): this is a real + # vllm-router knob, so it must flow into ``RouterArgs.from_cli_args(args, + # use_router_prefix=True)`` (which reads ``args.router_policy``) AND is read + # by ``vllm_rollout.generate`` to decide whether to send the ``x-session-id`` + # header for consistent-hash session-affinity routing (routing replay). + parser.add_argument( + "--vllm-router-policy", + type=str, + default="cache_aware", + dest="router_policy", + choices=["random", "round_robin", "cache_aware", "power_of_two", "consistent_hash"], + help=( + "vllm-router load-balancing policy. Use 'consistent_hash' to enable " + "session-affinity routing replay (routes a sample's requests to the same " + "engine via the x-session-id header)." + ), ) return parser @@ -200,6 +215,18 @@ def add_vllm_arguments(parser): default=512, help="Max concurrent inference requests sent to each vLLM server worker.", ) + parser.add_argument( + "--vllm-enable-deterministic-inference", + action="store_true", + default=False, + help=( + "Make rollout sampling deterministic. Forwards a per-sample ``seed`` " + "(derived from ``--rollout-seed`` and the sample's index in the group) " + "AND exports ``VLLM_BATCH_INVARIANT=1`` to the vLLM subprocess so attention " + "/ comm / MM kernels pick batch-invariant variants. Both are required for " + "true determinism — seed alone does not control kernel selection." + ), + ) parser.add_argument( "--vllm-weight-transfer-timeout-sec", type=float, @@ -252,10 +279,9 @@ def patched_add_argument_group(*g_args, **g_kwargs): # weight_transfer_config based on colocate) are applied explicitly in # ``vllm_engine.launch_server_process``. - # PD disaggregation / multi-group config — mirrors slime sglang_utils/arguments.py. - # vllm-side PD plumbing is not yet wired; the CLI surface is reserved so that - # rollout.py's `args.prefill_num_servers is not None` check is well-defined and - # the eventual vllm PD wiring can light up without further arg-layer changes. + # PD disaggregation / multi-group config. + # The CLI surface is reserved so that rollout.py's + # `args.prefill_num_servers is not None` check is well-defined. parser.add_argument( "--prefill-num-servers", type=int, @@ -263,6 +289,18 @@ def patched_add_argument_group(*g_args, **g_kwargs): help="Number of prefill servers for PD disaggregation.", ) + parser.add_argument( + "--vllm-config", + type=str, + default=None, + dest="vllm_config", + help=( + "Path to a YAML config file for fine-grained vLLM rollout engine deployment. " + "Enables multi-model serving, PD disaggregation, and heterogeneous server groups. " + "Mutually exclusive with --prefill-num-servers and --rollout-external." + ), + ) + return parser @@ -281,15 +319,13 @@ def validate_args(args): else: args.vllm_tp_size = args.rollout_num_gpus_per_engine - if getattr(args, "router_ip", None): - args.router_ip = _wrap_ipv6(args.router_ip) + if getattr(args, "vllm_router_ip", None): + args.vllm_router_ip = _wrap_ipv6(args.vllm_router_ip) def vllm_parse_args(): """Parse vllm flags via an independent ArgumentParser + parse_known_args. - Mirrors ``sglang_parse_args()`` so the merge flow in - ``slime/utils/arguments.py`` is symmetric. Returns an ``argparse.Namespace`` with all attrs prefixed ``vllm_``, plus: - ``_vllm_user_provided``: set of dests the user named on argv - ``_vllm_raw_values``: per-dest mapping to the user's literal CLI string @@ -320,12 +356,22 @@ def vllm_parse_args(): # try to forward them as command-line flags to the subprocess. _VIME_ORCHESTRATION_DESTS = frozenset( { - "router_ip", - "router_port", + "vllm_router_ip", + "vllm_router_port", + # bare dest: shares the --router-* namespace with the other vllm-router knobs. "router_request_timeout_secs", + # vllm-router routing policy: consumed by RouterArgs.from_cli_args when + # launching the router; never a `vllm serve` flag. + "router_policy", "vllm_server_concurrency", + "vllm_enable_deterministic_inference", "vllm_weight_transfer_timeout_sec", "vllm_weight_sync_packed", + # vime-only flags for fine-grained deployment; consumed in slime/ray/rollout.py + # (start_rollout_servers / _resolve_vllm_config) and must NOT be forwarded to + # the per-engine "vllm serve" subprocess. + "vllm_config", + "prefill_num_servers", } ) diff --git a/slime/backends/sglang_utils/sglang_config.py b/slime/backends/vllm_utils/vllm_config.py similarity index 73% rename from slime/backends/sglang_utils/sglang_config.py rename to slime/backends/vllm_utils/vllm_config.py index 827b47a1b..300fccdb4 100644 --- a/slime/backends/sglang_utils/sglang_config.py +++ b/slime/backends/vllm_utils/vllm_config.py @@ -1,4 +1,24 @@ -"""Configuration dataclasses for SGLang engine deployment.""" +"""Deployment configuration dataclasses for the vLLM rollout engine. + +YAML format:: + + vllm: + - name: actor + model_path: /path/to/actor + update_weights: true + num_gpus_per_engine: 2 + server_groups: + - worker_type: prefill + num_gpus: 4 + - worker_type: decode + num_gpus: 8 + - name: ref + model_path: /path/to/ref + update_weights: false + server_groups: + - worker_type: regular + num_gpus: 4 +""" import dataclasses import logging @@ -23,9 +43,7 @@ class ServerGroupConfig: num_gpus: Total number of GPUs for this group. num_gpus_per_engine: GPUs per engine for this group. Overrides the model-level or global ``--rollout-num-gpus-per-engine``. - overrides: Optional dict of SGLang ``ServerArgs`` field overrides. - These are applied on top of the base CLI ``--sglang-*`` - arguments in ``_compute_server_args``. + overrides: Optional dict of vLLM engine argument field overrides. """ worker_type: str @@ -72,11 +90,9 @@ def resolve(self, args) -> None: for g in self.server_groups: if g.num_gpus_per_engine is None: g.num_gpus_per_engine = default_gpus_per_engine - # Inject model_path into overrides so _compute_server_args picks it up. if "model_path" not in g.overrides: g.overrides["model_path"] = default_model_path - # Validate: all server groups within a model must share the same model_path. if self.server_groups: model_paths = {g.overrides["model_path"] for g in self.server_groups} assert len(model_paths) == 1, ( @@ -87,7 +103,6 @@ def resolve(self, args) -> None: else: effective_model_path = default_model_path - # Auto-infer update_weights when not explicitly set. if self.update_weights is None: if effective_model_path != args.hf_checkpoint: logger.warning( @@ -113,58 +128,29 @@ def total_num_gpus(self) -> int: @dataclasses.dataclass -class SglangConfig: - """Configuration for SGLang engine deployment. - - Loaded from ``--sglang-config`` YAML file. - - **Config format**:: - - sglang: - - name: actor - model_path: /path/to/actor - update_weights: true # receives training weight updates (default) - num_gpus_per_engine: 2 - server_groups: - - worker_type: prefill - num_gpus: 4 - num_gpus_per_engine: 2 - - worker_type: decode - num_gpus: 8 - num_gpus_per_engine: 4 - - name: ref - model_path: /path/to/ref - update_weights: false # frozen, no weight updates - server_groups: - - worker_type: regular - num_gpus: 4 - - Each model gets its own router. ``placeholder`` groups reserve GPU - slots without creating engines. ``overrides`` are ``ServerArgs`` - field names applied on top of the base ``--sglang-*`` CLI args. - - Set ``update_weights: false`` for frozen models (reference, reward, - etc.) that should not receive weight updates from training. - - .. note:: - - ``engine_groups`` is accepted as a backward-compatible alias for - ``server_groups`` in the YAML config. +class VllmConfig: + """Configuration for vLLM rollout engine deployment. + + Loaded from ``--vllm-config`` YAML file. Supports multi-model + serving, PD disaggregation, and heterogeneous server groups. + + See module docstring for the YAML format. """ models: list[ModelConfig] @staticmethod - def from_yaml(path: str) -> "SglangConfig": + def from_yaml(path: str) -> "VllmConfig": with open(path) as f: data = yaml.safe_load(f) - assert "sglang" in data, ( - f"sglang config must have a 'sglang' key, got {list(data.keys())}. " - f"Wrap your server_groups inside a model entry under 'sglang'." - ) + if "vllm" not in data: + raise ValueError( + f"vllm config must have a 'vllm' key, got {list(data.keys())}. " + "Wrap your server_groups inside a model entry under 'vllm'." + ) models = [] - for m in data["sglang"]: + for m in data["vllm"]: # Accept both "server_groups" and legacy "engine_groups". raw_groups = m.get("server_groups") or m.get("engine_groups") or [] groups = [ServerGroupConfig(**g) for g in raw_groups] @@ -177,16 +163,16 @@ def from_yaml(path: str) -> "SglangConfig": update_weights=m.get("update_weights"), ) ) - return SglangConfig(models=models) + return VllmConfig(models=models) @staticmethod - def from_prefill_num_servers(args) -> "SglangConfig": + def from_prefill_num_servers(args) -> "VllmConfig": """Build a config equivalent to the legacy --prefill-num-servers flag.""" total_gpus = args.rollout_num_gpus prefill_gpus = args.prefill_num_servers * args.rollout_num_gpus_per_engine decode_gpus = total_gpus - prefill_gpus assert decode_gpus > 0, f"No decode GPUs: total {total_gpus}, prefill {prefill_gpus}" - return SglangConfig( + return VllmConfig( models=[ ModelConfig( name="default", diff --git a/slime/backends/vllm_utils/vllm_engine.py b/slime/backends/vllm_utils/vllm_engine.py index 3b205dbb7..10a98582f 100644 --- a/slime/backends/vllm_utils/vllm_engine.py +++ b/slime/backends/vllm_utils/vllm_engine.py @@ -18,7 +18,7 @@ _spawn_ctx = multiprocessing.get_context("spawn") -# vLLM sleep/wake only supports these tags (SGLang also uses ``cuda_graph``, which must be dropped). +# vLLM sleep/wake only supports these tags. _VLLM_WAKE_TAGS = frozenset({"weights", "kv_cache"}) @@ -258,11 +258,10 @@ def launch_server_process( rank: int, visible_devices: str, model_path: str, + num_gpus_per_engine: int | None = None, ) -> multiprocessing.Process: """Spawn ``vllm serve`` (OpenAI API server) in a subprocess. - Contrasts with SGLang's launcher, which starts the HTTP server in-process from ``ServerArgs``. - Fixed flags (model identity, distributed topology, port/host, seed) are set by the orchestrator. Every other ``vllm serve`` flag is reachable via ``--vllm-`` and auto-forwarded by ``_forward_vllm_cli_args`` when the user overrides the vllm default. @@ -272,6 +271,20 @@ def launch_server_process( env.setdefault("NCCL_CUMEM_ENABLE", "0") env["CUDA_VISIBLE_DEVICES"] = visible_devices env.setdefault("VLLM_SERVER_DEV_MODE", "1") + # DeepGEMM fp8 kernels: vLLM enables them by default; set explicitly so the + # rollout engine is pinned regardless of the base image's defaults (replaces + # SGLang's SGLANG_JIT_DEEPGEMM_PRECOMPILE / _FAST_WARMUP). WARMUP="relax" JITs + # the required kernels before serving (no hot-path JIT) without the longer + # "full" sweep. Both overridable via the environment. + env.setdefault("VLLM_USE_DEEP_GEMM", "1") + env.setdefault("VLLM_DEEP_GEMM_WARMUP", "relax") + # Deterministic inference: VLLM_BATCH_INVARIANT=1 makes attention / comm / + # MM kernels pick batch-invariant variants so the same token sequence yields + # the same logits regardless of batch composition. Per-sample seed alone (see + # GenerateState / eval_rollout_single_dataset in vllm_rollout.py) is necessary + # but not sufficient for determinism. + if getattr(args, "vllm_enable_deterministic_inference", False): + env["VLLM_BATCH_INVARIANT"] = "1" # Colocate loads --worker-extension-cls from slime; vLLM subprocess must see the # same tree as the trainer (editable vime), not an older site-packages slime. if getattr(args, "colocate", False): @@ -288,7 +301,11 @@ def launch_server_process( host_for_subprocess = bind_host.strip("[]") model = model_path - tp = args.rollout_num_gpus_per_engine + # Per-engine TP: honor this engine's ServerGroup num_gpus_per_engine so a group + # configured with tp>1 launches tp>1, matching the weight-sync world_size + # accounting in update_weight_from_distributed (engine_gpu_counts). Falls back + # to the global flag for single-GPU-per-engine groups. + tp = num_gpus_per_engine or args.rollout_num_gpus_per_engine seed = getattr(args, "seed", 1234) + rank # Orchestrator-owned flags (correspond to SKIPPED_DESTS in vllm_utils/arguments.py). @@ -417,7 +434,7 @@ def _redact_cmd_for_log(cmd: list[str]) -> str: def _wait_server_healthy(base_url: str, process: multiprocessing.Process | None, timeout_s: float = 300.0) -> None: - """Wait until the vLLM server responds on ``GET /health`` (SGLang stacks typically use ``GET /health_generate``).""" + """Wait until the vLLM server responds on ``GET /health``.""" start = time.time() while True: try: @@ -444,7 +461,7 @@ def __init__( worker_type: str = "regular", base_gpu_id: int | None = None, model_path: str | None = None, - sglang_overrides: dict | None = None, + vllm_overrides: dict | None = None, num_gpus_per_engine: int | None = None, ): self.args = args @@ -453,7 +470,7 @@ def __init__( self.base_gpu_id = base_gpu_id self.model_path = model_path or args.hf_checkpoint # Uniform Ray ``start_engines`` kwargs; unused when launching vLLM over HTTP. - self.sglang_overrides = sglang_overrides or {} + self.vllm_overrides = vllm_overrides or {} self.num_gpus_per_engine = num_gpus_per_engine self.process: multiprocessing.Process | None = None self._weight_version: str | None = None @@ -478,8 +495,8 @@ def init( ): del dist_init_addr, nccl_port, disaggregation_bootstrap_port - self.router_ip = router_ip if router_ip is not None else self.args.router_ip - self.router_port = router_port if router_port is not None else self.args.router_port + self.router_ip = router_ip if router_ip is not None else self.args.vllm_router_ip + self.router_port = router_port if router_port is not None else self.args.vllm_router_port host = host or get_host_info()[1] self.server_host = _format_v6_uri(host) @@ -538,7 +555,6 @@ def _init_external(self) -> None: def _wait_external_config_ready(self) -> None: """External engine: best-effort ``GET /server_info`` TP check (non-fatal).""" try: - # SGLang external mode uses ``/get_server_info``; vLLM exposes ``/server_info``. actual = requests.get(f"{self._http_base()}/server_info", params={"config_format": "json"}, timeout=30) actual.raise_for_status() body = actual.json() @@ -558,7 +574,8 @@ def _wait_external_config_ready(self) -> None: def _init_normal(self) -> None: logger.info("Launch vLLM OpenAI api_server at: %s:%s", self.server_host, self.server_port) - num_gpus = min(self.args.num_gpus_per_node, self.args.rollout_num_gpus_per_engine) + gpus_per_engine = self.num_gpus_per_engine or self.args.rollout_num_gpus_per_engine + num_gpus = min(self.args.num_gpus_per_node, gpus_per_engine) base = self.base_gpu_id if self.base_gpu_id is not None else get_base_gpu_id(self.args, self.rank) base = _to_local_gpu_id(base) visible_devices = ",".join(str(base + i) for i in range(num_gpus)) @@ -571,6 +588,7 @@ def _init_normal(self) -> None: rank=self.rank, visible_devices=visible_devices, model_path=self.model_path, + num_gpus_per_engine=gpus_per_engine, ) _wait_server_healthy(self._http_base(), process=self.process) @@ -592,7 +610,7 @@ def _post_vllm_update_weights_http(self, update_info: dict) -> dict: return _response_json(response) def health_generate(self, timeout: float = 5.0) -> bool: - """Return True if ``GET /health`` succeeds (SGLang uses ``GET /health_generate`` for the same role).""" + """Return True if ``GET /health`` succeeds.""" if self.node_rank != 0: return True response = requests.get(f"{self._http_base()}/health", timeout=timeout) @@ -628,7 +646,7 @@ def update_weights_from_tensor( return response def flush_cache(self): - """Clear prefix cache via ``POST /reset_prefix_cache`` (SGLang uses ``GET /flush_cache``).""" + """Clear prefix cache via ``POST /reset_prefix_cache``.""" if self.node_rank != 0: return params = {"reset_running_requests": False, "reset_external": False} @@ -701,7 +719,7 @@ def get_weight_version(self) -> str | None: return self._weight_version def release_memory_occupation(self, level: int = 1): - """``POST /sleep?level={level}`` when sleep mode is enabled (SGLang: ``POST /release_memory_occupation``). + """``POST /sleep?level={level}`` when sleep mode is enabled. level=1 (default) releases KV cache only. level=0 releases both KV cache and model weights (required before IPC tensor injection). @@ -720,7 +738,7 @@ def release_memory_occupation(self, level: int = 1): return _response_json(response) def resume_memory_occupation(self, tags: list[str] | None = None): - """``POST /wake_up`` when sleep mode is on (SGLang: ``POST /resume_memory_occupation``); else a small placeholder dict.""" + """``POST /wake_up`` when sleep mode is on; else a small placeholder dict.""" if not getattr(self.args, "vllm_enable_sleep_mode", False): return {"ok": True, "sleep_mode": False} tags = _normalize_vllm_wake_tags(tags) @@ -773,7 +791,7 @@ def finish_weight_update(self) -> dict: return _response_json(response) def check_weights(self, action: str): - """No vLLM ``weights_checker`` route; return a placeholder (SGLang posts to ``/weights_checker``).""" + """No vLLM ``weights_checker`` route; return a placeholder.""" del action return {"ok": True, "supported": False, "note": "vLLM has no weights_checker endpoint."} @@ -806,7 +824,7 @@ def init_weights_update_group(self, master_address, master_port, rank_offset, wo raise RuntimeError(f"vLLM init_weight_transfer_engine failed: {last_error}") from last_error def destroy_weights_update_group(self, group_name): - """No vLLM destroy call; return ``None`` (SGLang may ``POST /destroy_weights_update_group`` and swallow errors).""" + """No vLLM destroy call; return ``None``.""" del group_name return None @@ -820,7 +838,7 @@ def update_weights_from_distributed( weight_version: str | None = None, packed: bool = True, ): - """NCCL path: POST ``/update_weights`` (SGLang: ``POST /update_weights_from_distributed``). + """NCCL path: POST ``/update_weights``. Payload matches vLLM NCCL weight transfer (see upstream rlhf_http_nccl example). """ @@ -839,7 +857,7 @@ def update_weights_from_distributed( return self._post_vllm_update_weights_http(update_info) def update_weights_from_disk(self, model_path: str, load_format: str | None = None): - """``POST /collective_rpc`` with ``reload_weights`` and ``weights_path`` (SGLang uses a dedicated disk API).""" + """``POST /collective_rpc`` with ``reload_weights`` and ``weights_path``.""" if self.node_rank != 0: return del load_format @@ -854,7 +872,7 @@ def update_weights_from_disk(self, model_path: str, load_format: str | None = No return _response_json(response) def pause_generation(self): - """``POST /pause`` with mode="keep" (SGLang: ``POST /pause_generation``); returns the ``requests.Response``.""" + """``POST /pause`` with mode="keep"; returns the ``requests.Response``.""" if self.node_rank != 0: return None response = requests.post( @@ -867,7 +885,7 @@ def pause_generation(self): return response def continue_generation(self): - """``POST /resume`` (SGLang: ``POST /continue_generation``).""" + """``POST /resume``.""" if self.node_rank != 0: return None response = requests.post(f"{self._http_base()}/resume", json={}, timeout=120) @@ -879,7 +897,7 @@ def post_process_weights( restore_weights_before_load: bool = False, post_process_quantization: bool = False, ): - """No vLLM HTTP hook (SGLang: ``POST /post_process_weights``); return a noop placeholder dict.""" + """No vLLM HTTP hook; return a noop placeholder dict.""" del restore_weights_before_load, post_process_quantization return {"ok": True, "noop": True, "note": "vLLM post_process is internal to load; no HTTP API."} diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py index c9ce21555..85ca748a6 100644 --- a/slime/ray/actor_group.py +++ b/slime/ray/actor_group.py @@ -51,8 +51,8 @@ def _allocate_gpus_for_actor(self, pg, num_gpus_per_actor): pg, reordered_bundle_indices, _reordered_gpu_ids = pg env_vars = { - # because sglang will always set NCCL_CUMEM_ENABLE to 0 - # we need also set it to 0 to prevent nccl error. + # Default NCCL_CUMEM_ENABLE to "0" to prevent intermittent NCCL + # init errors observed when the vLLM side disables CUMEM. "NCCL_CUMEM_ENABLE": os.environ.get("NCCL_CUMEM_ENABLE", "0"), "NVTE_FP8_BLOCK_SCALING_FP32_SCALES": os.environ.get("NVTE_FP8_BLOCK_SCALING_FP32_SCALES", "1"), **{name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST}, diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index 712e2d254..93c7bcb50 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -13,9 +13,12 @@ import ray import torch from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy -from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH, GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_WEIGHTS +from slime.backends.vllm_utils.vllm_config import ModelConfig, ServerGroupConfig, VllmConfig -from slime.backends.sglang_utils.sglang_config import ModelConfig, ServerGroupConfig, SglangConfig +# Memory-type tag strings shared with the vLLM engine's sleep/wake_up API. +GPU_MEMORY_TYPE_KV_CACHE = "kv_cache" +GPU_MEMORY_TYPE_WEIGHTS = "weights" +GPU_MEMORY_TYPE_CUDA_GRAPH = "cuda_graph" from slime.rollout.base_types import call_rollout_fn from slime.utils import logging_utils from slime.utils.health_monitor import RolloutHealthMonitor @@ -36,7 +39,7 @@ def _sanitize_vllm_router_args(ra: Any) -> Any: - """Replace negative int fields with dataclass defaults (sglang CLI may use -1; vllm-router rejects it).""" + """Replace negative int fields with dataclass defaults (legacy CLI may use -1; vllm-router rejects it).""" from vllm_router.router_args import RouterArgs as VR fixes: dict[str, Any] = {} @@ -62,7 +65,7 @@ def _vllm_router_args_from_cli(args: Namespace) -> Any: @dataclasses.dataclass class ServerGroup: - """A group of homogeneous SGLang engines with the same configuration. + """A group of homogeneous vLLM engines with the same configuration. All engines in a group share the same tp_size / nodes_per_engine / pg. A RolloutServer may contain multiple ServerGroups (e.g. prefill vs decode @@ -77,7 +80,7 @@ class ServerGroup: worker_type: str = "regular" # "regular", "prefill", "decode", or "placeholder" rank_offset: int = 0 # cumulative engine count before this group gpu_offset: int = 0 # cumulative GPU count before this group - sglang_overrides: dict = dataclasses.field(default_factory=dict) + vllm_overrides: dict = dataclasses.field(default_factory=dict) needs_offload: bool = False # True when this group's GPUs overlap with megatron model_path: str | None = None # checkpoint path for update_weights_from_disk router_ip: str | None = None @@ -139,14 +142,6 @@ def start_engines(self, port_cursors: dict[int, int] | None = None) -> tuple[lis env_vars = {name: "1" for name in NOSET_VISIBLE_DEVICES_ENV_VARS_LIST} | { key: os.environ.get(key, default_val) for key, default_val in { - "SGLANG_JIT_DEEPGEMM_PRECOMPILE": "true", - "SGLANG_JIT_DEEPGEMM_FAST_WARMUP": "true", - "SGL_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true", - "SGLANG_DISABLE_TP_MEMORY_INBALANCE_CHECK": "true", - "SGLANG_MEMORY_SAVER_CUDA_GRAPH": "true", - "SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_FALLBACK_VARIANT": "true", - "SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION": "false", - "SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_IDLE": "false", "SLIME_ENABLE_PROFILING": "true", }.items() } @@ -163,7 +158,7 @@ def start_engines(self, port_cursors: dict[int, int] | None = None) -> tuple[lis worker_type=self.worker_type, base_gpu_id=base_gpu_id, model_path=self.model_path, - sglang_overrides=self.sglang_overrides, + vllm_overrides=self.vllm_overrides, num_gpus_per_engine=self.num_gpus_per_engine, ) @@ -247,6 +242,7 @@ class RolloutServer: server_groups: list[ServerGroup] router_ip: str | None = None router_port: int | None = None + prometheus_port: int | None = None model_name: str = "default" update_weights: bool = True @@ -420,17 +416,19 @@ def __init__(self, args, pg): self._ci_fault_injection_pending = self.args.ci_test # Flag for CI fault injection def _get_metrics_router_addr(self) -> str | None: - """Return the router address for scraping SGLang engine metrics. - - The sglang_router gateway exposes ``/engine_metrics`` on its main port, - which aggregates Prometheus metrics from all backend sglang servers. - Returns ``http://{ip}:{port}`` for the first server, or ``None`` when - metrics are disabled or no servers are running. + """Return the full Prometheus scrape URL for the rollout router. + + vllm-router exposes Prometheus on a dedicated ``prometheus_port`` + (see ``router_args.prometheus_port`` in ``_start_router``), not via + a path on the main router port. The metrics endpoint is the default + ``/metrics`` served by the metrics-exporter-prometheus crate. + Returns ``http://{ip}:{prom_port}/metrics``, or ``None`` if metrics + are disabled or no servers are running. """ srv = self.server - if srv is None or srv.router_ip is None: + if srv is None or srv.router_ip is None or srv.prometheus_port is None: return None - return f"http://{srv.router_ip}:{srv.router_port}" + return f"http://{srv.router_ip}:{srv.prometheus_port}/metrics" def get_metrics_router_addr(self) -> str | None: """Public wrapper for remote calls from the driver process.""" @@ -460,7 +458,7 @@ def _try_ci_fault_injection(self): def dispose(self): for monitor in self._health_monitors: monitor.stop() - # Release inference workers (vLLM / SGLang). debug_rollout_only still hits this path at train.py end. + # Release vLLM inference workers. debug_rollout_only still hits this path at train.py end. shutdown_refs = [] for srv in self.servers.values(): for group in srv.server_groups: @@ -958,14 +956,14 @@ def addr(): def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool = False) -> tuple[str, int]: """Start the rollout HTTP gateway (vllm-router).""" - if not force_new and args.router_ip is not None: - return args.router_ip, args.router_port + if not force_new and args.vllm_router_ip is not None: + return args.vllm_router_ip, args.vllm_router_port router_ip = _wrap_ipv6(get_host_info()[1]) if force_new: router_port = find_available_port(random.randint(3000, 4000)) else: - router_port = args.router_port + router_port = args.vllm_router_port if router_port is None: router_port = find_available_port(random.randint(3000, 4000)) @@ -992,11 +990,11 @@ def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool if any(f.name == "disable_health_check" for f in dataclasses.fields(type(router_args))): router_args.disable_health_check = True - logger.info("Launch HTTP router (impl=vllm) with args: %s", router_args) + logger.info("Launch HTTP router with args: %s", router_args) process = multiprocessing.get_context("spawn").Process( target=run_router, - args=(("vllm", router_args),), + args=(router_args,), ) process.daemon = True process.start() @@ -1008,7 +1006,7 @@ def _start_router(args, *, has_pd_disaggregation: bool = False, force_new: bool "MiniLB is only valid with PD disaggregation. See slime.utils.http_utils run_router logs." ) logger.info(f"Router launched at {router_ip}:{router_port}, Prometheus port: {router_args.prometheus_port}") - return router_ip, router_port + return router_ip, router_port, router_args.prometheus_port def _compute_rollout_offset(args) -> int: @@ -1030,7 +1028,7 @@ def _compute_megatron_num_gpus(args) -> int: def start_rollout_servers(args, pg) -> dict[str, RolloutServer]: """Start rollout servers: one per model, each with its own router. - Each model defined in the sglang config gets its own router and set + Each model defined in the vLLM config gets its own router and set of server groups. Server groups within a model may have different ``num_gpus_per_engine`` (e.g. for PD disaggregation where prefill and decode use different TP sizes). @@ -1040,7 +1038,7 @@ def start_rollout_servers(args, pg) -> dict[str, RolloutServer]: Note: ``init_http_client`` should be called separately before this, as the HTTP client is shared across all servers. """ - config = _resolve_sglang_config(args) + config = _resolve_vllm_config(args) servers: dict[str, RolloutServer] = {} gpu_offset = 0 @@ -1054,13 +1052,13 @@ def start_rollout_servers(args, pg) -> dict[str, RolloutServer]: model_cfg.resolve(args) has_pd = model_cfg.has_pd_disaggregation - router_ip, router_port = _start_router(args, has_pd_disaggregation=has_pd, force_new=(model_idx > 0)) + router_ip, router_port, prom_port = _start_router(args, has_pd_disaggregation=has_pd, force_new=(model_idx > 0)) # Write back so downstream readers (vllm_rollout, vllm_engine) see the # router we just started (only relevant for first model in multi-model setups). if model_idx == 0: - args.router_ip = router_ip - args.router_port = router_port + args.vllm_router_ip = router_ip + args.vllm_router_port = router_port server_groups: list[ServerGroup] = [] port_cursors: dict[int, int] = {} @@ -1095,7 +1093,7 @@ def _make_group(group_cfg, router_ip, router_port, overrides_extra=None): worker_type=group_cfg.worker_type, rank_offset=engine_offset, gpu_offset=gpu_offset, - sglang_overrides=overrides, + vllm_overrides=overrides, needs_offload=needs_offload, model_path=overrides.get("model_path", args.hf_checkpoint), router_ip=router_ip, @@ -1157,29 +1155,30 @@ def _make_group(group_cfg, router_ip, router_port, overrides_extra=None): router_port=router_port, model_name=model_cfg.name, update_weights=model_cfg.update_weights, + prometheus_port=prom_port, ) # Expose per-model router info for custom rollout functions. - args.sglang_model_routers = {name: (srv.router_ip, srv.router_port) for name, srv in servers.items()} + args.vllm_model_routers = {name: (srv.router_ip, srv.router_port) for name, srv in servers.items()} return servers -def _resolve_sglang_config(args) -> SglangConfig: - """Build a SglangConfig from args, choosing the right source.""" - if getattr(args, "sglang_config", None) is not None: - config = SglangConfig.from_yaml(args.sglang_config) - # Validate total GPUs match. +def _resolve_vllm_config(args) -> VllmConfig: + """Build a VllmConfig from args, choosing the right source.""" + vllm_config_path = getattr(args, "vllm_config", None) + if vllm_config_path is not None: + config = VllmConfig.from_yaml(vllm_config_path) expected = args.rollout_num_gpus actual = config.total_num_gpus - assert actual == expected, f"sglang_config total GPUs ({actual}) != rollout_num_gpus ({expected})" + assert actual == expected, f"vllm_config total GPUs ({actual}) != rollout_num_gpus ({expected})" return config if args.prefill_num_servers is not None: - return SglangConfig.from_prefill_num_servers(args) + return VllmConfig.from_prefill_num_servers(args) # Default: single regular group. - return SglangConfig( + return VllmConfig( models=[ ModelConfig( name="default", diff --git a/slime/rollout/on_policy_distillation.py b/slime/rollout/on_policy_distillation.py index 919097434..949895513 100644 --- a/slime/rollout/on_policy_distillation.py +++ b/slime/rollout/on_policy_distillation.py @@ -1,67 +1,145 @@ +"""On-Policy Distillation (OPD): teacher logprobs via vLLM ``/inference/v1/generate``. + +The reward function sends the full (prompt + student response) token sequence to +an external vLLM teacher server's ``/inference/v1/generate`` endpoint with +``prompt_logprobs`` enabled; the post-process step reads the per-token teacher +logprobs out of the response and stores them on each sample for the OPD KL +penalty. + +Endpoint contract (vime vLLM disaggregated ``/inference/v1/generate``): + +- Request body:: + + { + "token_ids": [...], # full prompt+response token ids + "model": , # OPTIONAL on this endpoint; only sent when + # --opd-teacher-model is set (single-model + # teacher servers use their loaded model) + "sampling_params": { + "max_tokens": 1, # endpoint requires >=1; the generated + # token is ignored + "temperature": 0, + "prompt_logprobs": 1, # score every prompt token + "skip_special_tokens": False, + }, + } + +- Response body (``GenerateResponse``):: + + { + "choices": [{"token_ids": [...], "logprobs": {...}, ...}], + "prompt_logprobs": [ # TOP-LEVEL, aligned with the input token_ids + None, # position 0: no prior context + {: {"logprob": -3.21, "rank": 1, "decoded_token": "..."}, ...}, + ... + ], + } + +``prompt_logprobs[i]`` is a dict ``{token_id -> Logprob}`` for the token at +position ``i``. JSON serializes integer dict keys as strings, so we look up by +both ``int`` and ``str``. See ``vllm/entrypoints/serve/disagg/protocol.py`` +(``GenerateRequest`` / ``GenerateResponse``). +""" + +from __future__ import annotations + +from typing import Any + import aiohttp import torch -from slime.utils.processing_utils import encode_image_for_rollout_engine from slime.utils.types import Sample async def reward_func(args, sample, **kwargs): - payload = { - # "text": sample.prompt + sample.response, - "input_ids": sample.tokens, + if sample.multimodal_inputs and sample.multimodal_inputs.get("images"): + # ``/inference/v1/generate`` is token-only; a multimodal teacher requires the + # render -> generate flow (``/v1/chat/completions/render`` then attach + # ``features``), as in ``slime.rollout.vllm_rollout.generate``. Not yet wired + # for OPD — fail loudly rather than silently scoring a text-only sequence. + raise NotImplementedError( + "OPD multimodal teacher scoring over /inference/v1/generate is not implemented; " + "wire the /v1/chat/completions/render -> /inference/v1/generate flow first." + ) + + payload: dict[str, Any] = { + "token_ids": sample.tokens, "sampling_params": { + "max_tokens": 1, "temperature": 0, - "max_new_tokens": 0, + "prompt_logprobs": 1, "skip_special_tokens": False, }, - "return_logprob": True, - "logprob_start_len": 0, } + # ``model`` is optional on /inference/v1/generate. Only send it when the teacher + # model is explicitly named; never fall back to the *student* checkpoint + # (args.hf_checkpoint), which would mis-name a teacher!=student server. + teacher_model = getattr(args, "opd_teacher_model", None) + if teacher_model: + payload["model"] = teacher_model - if sample.multimodal_inputs and sample.multimodal_inputs.get("images"): - image_data = sample.multimodal_inputs["images"] - payload["image_data"] = [encode_image_for_rollout_engine(image) for image in image_data] - - session_kwargs = {} - async with aiohttp.ClientSession(**session_kwargs) as session: + async with aiohttp.ClientSession() as session: async with session.post(args.rm_url, json=payload) as resp: resp.raise_for_status() return await resp.json() -def post_process_rewards(args, samples: list[Sample], **kwargs): - """Process rewards from teacher model and extract teacher log probabilities. +def _logprob_for_token(pos_entry: dict | None, token_id: int) -> float: + """Pull the teacher's logprob for ``token_id`` out of one position's logprob dict. - This function: - 1. Extracts teacher log-probs from the reward response (which contains sglang's logprob output) - 2. Trims them to match the response length - 3. Stores them in sample.teacher_log_probs for OPD KL penalty computation - 4. Returns scalar rewards (0.0 for pure distillation) compatible with GRPO/PPO + Raises on a missing entry: vLLM always includes the actual prompt token in + ``prompt_logprobs`` (even when only top-1 is requested), so a missing token + means a malformed/misaligned response that must not be papered over with 0.0. + JSON serializes integer dict keys as strings, so we accept both. Each value + is either a dict (``{"logprob": float, ...}``) or a flattened float. + """ + if pos_entry is None: + raise ValueError("teacher prompt_logprobs has a None entry at a scored position") + entry = pos_entry.get(token_id) + if entry is None: + entry = pos_entry.get(str(token_id)) + if entry is None: + raise ValueError(f"teacher prompt_logprobs missing logprob for token_id={token_id}") + if isinstance(entry, dict): + return float(entry["logprob"]) + if isinstance(entry, (int, float)): + return float(entry) + return float(entry.logprob) + + +def post_process_rewards(args, samples: list[Sample], **kwargs): + """Extract teacher log-probs from the ``/inference/v1/generate`` responses. - Note: The reward_func calls the teacher server which returns token-level log-probs. - For pure on-policy distillation without task rewards, we return 0.0 for each sample. - The actual learning signal comes from the OPD KL penalty applied in compute_advantages_and_returns. + 1. Read top-level ``prompt_logprobs`` (aligned with the submitted token_ids). + 2. Pick out each actual token's logprob, skipping position 0 (always None). + 3. Trim to the response length and store on ``sample.teacher_log_probs``. + 4. Return scalar rewards (0.0 for pure distillation); the learning signal is + the OPD KL penalty applied in ``compute_advantages_and_returns``. """ raw_rewards = [sample.get_reward_value(args) for sample in samples] response_lengths = [sample.response_length for sample in samples] - # Extract teacher log-probs from the sglang response - teacher_log_probs = [ - torch.tensor([item[0] for item in reward["meta_info"]["input_token_logprobs"][1:]], dtype=torch.float32) - for reward in raw_rewards - ] - teacher_log_probs = [ - t_log_prob[-response_length:] - for t_log_prob, response_length in zip(teacher_log_probs, response_lengths, strict=False) - ] - - for sample, t_log_probs in zip(samples, teacher_log_probs, strict=False): + teacher_log_probs: list[torch.Tensor] = [] + for reward, sample in zip(raw_rewards, samples, strict=True): + plp = reward.get("prompt_logprobs") + assert plp is not None, "teacher response missing top-level prompt_logprobs" + assert len(plp) == len( + sample.tokens + ), f"prompt_logprobs length {len(plp)} != token_ids length {len(sample.tokens)}" + # plp[i] scores sample.tokens[i]; position 0 has no prior context. + per_pos = [_logprob_for_token(plp[i], sample.tokens[i]) for i in range(1, len(sample.tokens))] + teacher_log_probs.append(torch.tensor(per_pos, dtype=torch.float32)) + + trimmed: list[torch.Tensor] = [] + for t_log_prob, response_length in zip(teacher_log_probs, response_lengths, strict=True): + assert ( + len(t_log_prob) >= response_length + ), f"teacher logprobs ({len(t_log_prob)}) shorter than response_length ({response_length})" + trimmed.append(t_log_prob[-response_length:]) + + for sample, t_log_probs in zip(samples, trimmed, strict=True): sample.teacher_log_probs = t_log_probs - # Return scalar rewards for GRPO/PPO advantage estimator - # For pure on-policy distillation, we use 0.0 as the task reward. - # The learning signal comes entirely from the OPD KL penalty. - # If you have task rewards, you can add them here. + # Pure on-policy distillation: task reward is 0; KL penalty carries the signal. scalar_rewards = [0.0] * len(samples) - return scalar_rewards, scalar_rewards diff --git a/slime/rollout/sglang_rollout.py b/slime/rollout/sglang_rollout.py deleted file mode 100644 index 44c03858f..000000000 --- a/slime/rollout/sglang_rollout.py +++ /dev/null @@ -1,624 +0,0 @@ -import asyncio -import copy -import inspect -import logging -import uuid -from argparse import Namespace -from collections.abc import Callable -from contextlib import contextmanager -from typing import Any - -import numpy as np -import pybase64 -import sglang_router -from packaging.version import parse -from tqdm import tqdm - -from slime.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainOutput -from slime.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter -from slime.utils.async_utils import run -from slime.utils.data import Dataset -from slime.utils.eval_config import EvalDatasetConfig -from slime.utils.http_utils import get, post -from slime.utils.misc import SingletonMeta, load_function -from slime.utils.processing_utils import ( - build_processor_kwargs, - encode_image_for_rollout_engine, - load_processor, - load_tokenizer, -) -from slime.utils.trace_utils import build_sglang_meta_trace_attrs, trace_function, trace_span -from slime.utils.types import Sample - -from .rm_hub import async_rm, batched_async_rm - -__all__ = ["generate_rollout", "get_model_url"] - -logger = logging.getLogger(__name__) - -_PROCESSOR_PROMPT_KEYS = {"input_ids", "attention_mask"} - - -def _prepare_prompt_ids(sample: Sample, tokenizer, processor: Any) -> list[int]: - raw_multimodal_inputs = sample.multimodal_inputs or {} - has_multimodal_inputs = any(value is not None for value in raw_multimodal_inputs.values()) - reuse_existing_input_ids = bool(sample.tokens) and ( - sample.multimodal_train_inputs is not None or not has_multimodal_inputs - ) - - if processor and has_multimodal_inputs and not reuse_existing_input_ids: - processor_output = processor(text=sample.prompt, **build_processor_kwargs(raw_multimodal_inputs)) - prompt_ids = processor_output["input_ids"][0] - if sample.multimodal_train_inputs is None: - sample.multimodal_train_inputs = { - k: v for k, v in processor_output.items() if k not in _PROCESSOR_PROMPT_KEYS - } or None - return prompt_ids - - if reuse_existing_input_ids: - return sample.tokens - - return tokenizer.encode(sample.prompt, add_special_tokens=False) - - -def get_model_url(args: Namespace, model_name: str, endpoint: str = "/generate") -> str: - """Return the router URL for a named model. - - Use this in custom rollout functions to route requests to a specific - model when multiple models are deployed via ``--sglang-config``:: - - url = get_model_url(args, "ref", "/generate") - resp = await post(url, json=payload) - - Falls back to the default router if *model_name* is not found or - ``sglang_model_routers`` is not set. - """ - routers = getattr(args, "sglang_model_routers", None) - if routers and model_name in routers: - ip, port = routers[model_name] - return f"http://{ip}:{port}{endpoint}" - return f"http://{args.sglang_router_ip}:{args.sglang_router_port}{endpoint}" - - -class GenerateState(metaclass=SingletonMeta): - """ - The global state for the generation process. - """ - - def __init__(self, args: Namespace) -> None: - # persistent state for the generation process - self.args = args - self.tokenizer = load_tokenizer(args.hf_checkpoint, trust_remote_code=True) - self.processor = load_processor(args.hf_checkpoint, trust_remote_code=True) - - self.semaphore = asyncio.Semaphore( - args.sglang_server_concurrency * args.rollout_num_gpus // args.rollout_num_gpus_per_engine - ) - self.sampling_params: dict[str, Any] = dict( - temperature=args.rollout_temperature, - top_p=args.rollout_top_p, - top_k=args.rollout_top_k, - max_new_tokens=args.rollout_max_response_len, - stop=args.rollout_stop, - stop_token_ids=args.rollout_stop_token_ids, - skip_special_tokens=args.rollout_skip_special_tokens, - no_stop_trim=True, - spaces_between_special_tokens=False, - ) - - if getattr(args, "sglang_enable_deterministic_inference", False): - sampling_seed_base = args.rollout_seed - self.group_sampling_seeds = [sampling_seed_base + i for i in range(args.n_samples_per_prompt)] - - # dp rank balancing - self.dp_counts = [0] * (args.sglang_dp_size or 1) - self.dp_rank = 0 - - self.reset() - - @contextmanager - def dp_rank_context(self): - candidates = [i for i, count in enumerate(self.dp_counts) if count == min(self.dp_counts)] - dp_rank = int(np.random.choice(candidates)) - self.dp_counts[dp_rank] += 1 - self.dp_rank = dp_rank - try: - yield dp_rank - finally: - self.dp_counts[dp_rank] -= 1 - assert self.dp_counts[dp_rank] >= 0 - - def reset(self) -> None: - self.remaining_batch_size = 0 - self.pendings = set() - self.aborted = False - - def submit_generate_tasks(self, samples: list[list[Sample]]) -> None: - for group in samples: - self.pendings.add( - asyncio.create_task( - # submit a group of samples as a single task. - generate_and_rm_group( - self.args, - group, - sampling_params=self.sampling_params.copy(), - evaluation=False, - ) - ) - ) - self.remaining_batch_size += len(samples) - - -async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, Any]) -> Sample: - """Generate using traditional SGLang router with token-based workflow""" - if args.ci_test: - assert isinstance(sample.prompt, str) - - state = GenerateState(args) - url = f"http://{args.sglang_router_ip}:{args.sglang_router_port}/generate" - - assert ( - sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED - ), f"Sample status is {sample.status}" - - prompt_ids = _prepare_prompt_ids(sample, state.tokenizer, state.processor) - - assert ( - sampling_params["max_new_tokens"] >= 0 - ), f"max_new_tokens: {sampling_params['max_new_tokens']} should not be less than 0" - if sampling_params["max_new_tokens"] == 0: - sample.status = Sample.Status.TRUNCATED - return sample - - # Prepare payload for sglang server - payload = { - "sampling_params": sampling_params, - "return_logprob": True, - } - - if args.use_rollout_routing_replay: - payload["return_routed_experts"] = True - - images = sample.multimodal_inputs.get("images") if sample.multimodal_inputs else None - if images: - payload["image_data"] = [encode_image_for_rollout_engine(image) for image in images] - # For single-turn multimodal requests, send text so SGLang expands the - # image placeholders with its own processor rules. - payload["text"] = sample.prompt - else: - payload["input_ids"] = prompt_ids - - if not sample.tokens: - sample.tokens = prompt_ids - - # Use session_id for consistent hashing routing (SGLang Model Gateway) - headers = None - if sample.session_id: - if getattr(args, "router_policy", None) == "consistent_hashing": - headers = {"X-SMG-Routing-Key": sample.session_id} - - with trace_span(sample, "sglang_generate", attrs={"max_new_tokens": sampling_params["max_new_tokens"]}) as span: - output = await post(url, payload, headers=headers) - span.update(build_sglang_meta_trace_attrs(output["meta_info"])) - - if "output_token_logprobs" in output["meta_info"]: - new_response_tokens = [item[1] for item in output["meta_info"]["output_token_logprobs"]] - new_response_log_probs = [item[0] for item in output["meta_info"]["output_token_logprobs"]] - else: - new_response_tokens, new_response_log_probs = [], [] - - # Update sample with tokens directly - avoiding re-tokenization - sample.tokens = sample.tokens + new_response_tokens - sample.response_length += len(new_response_tokens) - sample.response += output["text"] - - # When partial rollout and masking off policy is enabled, update the loss mask - if sample.loss_mask is not None: - assert args.partial_rollout and args.mask_offpolicy_in_partial_rollout - sample.loss_mask += [1] * len(new_response_tokens) - - if sample.rollout_log_probs is None: - sample.rollout_log_probs = [] - sample.rollout_log_probs += new_response_log_probs - - if "routed_experts" in output["meta_info"]: - sample.rollout_routed_experts = np.frombuffer( - pybase64.b64decode(output["meta_info"]["routed_experts"].encode("ascii")), - dtype=np.int32, - ).reshape( - len(sample.tokens) - 1, - args.num_layers, - args.moe_router_topk, - ) - - sample.update_from_meta_info(args, output["meta_info"]) - - return sample - - -@trace_function("generate_and_rm", target="sample") -async def generate_and_rm( - args: Namespace, - sample: Sample | list[Sample], - sampling_params: dict[str, Any], - evaluation: bool = False, -) -> Sample | list[Sample]: - # mask previous off-policy generation for partial rollout - if args.partial_rollout and args.mask_offpolicy_in_partial_rollout and sample.response_length > 0: - sample.loss_mask = [0] * sample.response_length - - # For samples with existing response, check if they're complete - if sample.status == Sample.Status.COMPLETED or sample.status == Sample.Status.TRUNCATED: - assert sample.response is not None - if not args.group_rm: - assert sample.reward is not None - return sample - - state = GenerateState(args) - - # generate - async with state.semaphore: - if state.aborted: - sample.status = Sample.Status.ABORTED - return sample - - with state.dp_rank_context() as _: - # Check sample.generate_function_path for per-sample custom_generate_function_path (e.g., from eval dataset config) - custom_func_path = getattr(sample, "generate_function_path", None) or args.custom_generate_function_path - - if custom_func_path is not None: - custom_generate_func = load_function(custom_func_path) - # if signature has evaluation, pass evaluation - if "evaluation" in inspect.signature(custom_generate_func).parameters: - sample = await custom_generate_func(args, sample, sampling_params, evaluation=evaluation) - else: - sample = await custom_generate_func(args, sample, sampling_params) - else: - sample = await generate(args, sample, sampling_params) - - # for the rm that need the whole group, we will not do the rm here - if args.group_rm: - return sample - - if isinstance(sample, list): - samples = sample - if any(sample.status == Sample.Status.ABORTED for sample in samples): - return samples - - samples_need_reward = [sample for sample in samples if sample.reward is None] - with trace_span(samples_need_reward, "reward_model"): - rewards = await batched_async_rm(args, samples_need_reward) - for sample, reward in zip(samples_need_reward, rewards, strict=False): - sample.reward = reward - return samples - else: - if sample.status == Sample.Status.ABORTED: - return sample - # Some custom generate paths may have already filled the reward. - if sample.reward is None: - with trace_span(sample, "reward_model"): - sample.reward = await async_rm(args, sample) - - return sample - - -@trace_function( - "generate_and_rm_group", - target="group", - attrs_getter=lambda args, group, sampling_params, evaluation=False: {"group_size": len(group)}, -) -async def generate_and_rm_group( - args: Namespace, group: list[Sample], sampling_params: dict[str, Any], evaluation: bool = False -) -> list[Sample]: - state = GenerateState(args) - - if state.aborted: - return group - - # Generate a unique session_id for each sample in the group - for sample in group: - if sample.session_id is None: - sample.session_id = str(uuid.uuid4()) - - tasks = [] - for idx, sample in enumerate(group): - current_sampling_params = sampling_params.copy() - if getattr(args, "sglang_enable_deterministic_inference", False): - seed = state.group_sampling_seeds[idx] - current_sampling_params["sampling_seed"] = seed - tasks.append( - asyncio.create_task(generate_and_rm(args, sample, current_sampling_params, evaluation=evaluation)) - ) - - group = await asyncio.gather(*tasks) - - # for the rm that need the whole group, we will do the rm here - if not state.aborted and args.group_rm: - with trace_span(group, "group_reward_model"): - rewards = await batched_async_rm(args, group) - for sample, reward in zip(group, rewards, strict=False): - sample.reward = reward - - return group - - -async def abort(args: Namespace, rollout_id: int) -> list[list[Sample]]: - aborted_samples = [] - - state = GenerateState(args) - assert not state.aborted - state.aborted = True - - if parse(sglang_router.__version__) <= parse("0.2.1"): - response = await get(f"http://{args.sglang_router_ip}:{args.sglang_router_port}/list_workers") - urls = response["urls"] - else: - response = await get(f"http://{args.sglang_router_ip}:{args.sglang_router_port}/workers") - urls = [worker["url"] for worker in response["workers"]] - - logger.info(f"Abort request for {urls}") - abort_tasks = [post(f"{url}/abort_request", {"abort_all": True}) for url in urls] - abort_results = await asyncio.gather(*abort_tasks, return_exceptions=True) - for url, result in zip(urls, abort_results, strict=False): - if isinstance(result, Exception): - logger.warning(f"Failed to abort worker at {url}: {result}") - - # make sure all the pending tasks are finished - count = 0 - while state.pendings: - done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED) - - if not args.partial_rollout: - continue - - # for partial rollout, collect the partial samples into the data buffer - for task in done: - group = task.result() - for sample in group: - if sample.response and "start_rollout_id" not in sample.metadata: - sample.metadata["start_rollout_id"] = rollout_id - aborted_samples.append(group) - count += len(group) - - if args.partial_rollout: - logger.info(f"Collected {count} partial samples into the data buffer") - - return aborted_samples - - -async def generate_rollout_async( - args: Namespace, rollout_id: int, data_source: Callable[[int], list[list[Sample]]] -) -> tuple[RolloutFnTrainOutput, list[list[Sample]]]: - """An example to implement the generate_rollout function for an rule based rm rollout generation. - - Args: - args: the whole args - rollout_id: int, the id of the rollout, used for deterministic data generation - data_source: the data source to fetch - - Returns: - tuple[RolloutFnTrainOutput, list[list[Sample]]]: - - data: a list of groups of samples generated by the rollout, length equals `rollout_batch_size` - - aborted_samples: any partial groups collected during abort when partial_rollout is enabled - """ - assert args.rollout_global_dataset - - state = GenerateState(args) - - # instantiate data filters - dynamic_filter = ( - load_function(args.dynamic_sampling_filter_path) if args.dynamic_sampling_filter_path is not None else None - ) - - metric_gatherer = MetricGatherer() - - # target_data_size is the total number of valid samples to get - target_data_size = args.rollout_batch_size - - data = [] - all_data = [] - do_print = True - pbar = tqdm(total=target_data_size * args.n_samples_per_prompt, desc="Rollout generation") - while len(data) < target_data_size: - while state.remaining_batch_size < target_data_size: - # get samples from the buffer and submit the generation requests. - samples = data_source(args.over_sampling_batch_size) - state.submit_generate_tasks(samples) - - # wait for the generation to finish - done, state.pendings = await asyncio.wait(state.pendings, return_when=asyncio.FIRST_COMPLETED) - for task in done: - group: list[Sample] = task.result() - - if do_print: - sample = group[0][0] if isinstance(group[0], list) else group[0] - logger.info( - f"First rollout sample: {[str(sample.prompt) + sample.response]}, label: {str(sample.label)[:100]}, reward: {sample.reward}", - ) - do_print = False - - assert len(group) == args.n_samples_per_prompt - all_data.append(group) - dynamic_filter_output = call_dynamic_filter(dynamic_filter, args, group) - if not dynamic_filter_output.keep: - metric_gatherer.on_dynamic_filter_drop(reason=dynamic_filter_output.reason) - state.remaining_batch_size -= 1 - continue - - # add the samples to the data - # NOTE: here we have not stored all the unused samples back to the data buffer. - if len(data) < target_data_size: - data.append(group) - pbar.update(args.n_samples_per_prompt) - - pbar.close() - sample = data[-1][0][0] if isinstance(data[-1][0], list) else data[-1][0] - logger.info( - f"Finish rollout: {[str(sample.prompt) + sample.response]}, label: {str(sample.label)[:100]}, reward: {sample.reward}", - ) - - # there are still some unfinished requests, abort them - aborted_samples = await abort(args, rollout_id) - - assert len(data) == args.rollout_batch_size, f"Got {len(data)} samples, expected {args.rollout_batch_size}" - data = sorted(data, key=lambda group: group[0][0].index if isinstance(group[0], list) else group[0].index) - all_samples = sorted( - all_data, key=lambda group: group[0][0].index if isinstance(group[0], list) else group[0].index - ) - - # reset the global state to prevent effects on the next rollout or eval. - state.reset() - if args.rollout_sample_filter_path is not None: - filter_func = load_function(args.rollout_sample_filter_path) - filter_func(args, data) - - # There can be circumstances where users want to process all samples including filtered ones. - if args.rollout_all_samples_process_path is not None: - process_func = load_function(args.rollout_all_samples_process_path) - process_func(args, all_samples, data_source) - - return RolloutFnTrainOutput(samples=data, metrics=metric_gatherer.collect()), aborted_samples - - -EVAL_PROMPT_DATASET = {} - - -async def eval_rollout(args: Namespace, rollout_id: int) -> tuple[dict[str, dict[str, list[Any]]], list[list[Sample]]]: - assert not args.group_rm, "Group RM is not supported for eval rollout" - - coros = [] - for dataset_cfg in getattr(args, "eval_datasets", []) or []: - coros.append(eval_rollout_single_dataset(args, rollout_id, dataset_cfg)) - results_list = await asyncio.gather(*coros) - results = {} - for r in results_list: - results.update(r) - return RolloutFnEvalOutput(data=results), [] - - -async def eval_rollout_single_dataset( - args: Namespace, rollout_id: int, dataset_cfg: EvalDatasetConfig -) -> dict[str, dict[str, list[Any]]]: - """An example to implement the eval_rollout function for an rule based rm rollout generation. - - Args: - args: the whole args - rollout_id: int, the id of the rollout, used for deterministic data generation - dataset_cfg: configuration of the dataset - """ - assert not args.group_rm, "Group RM is not supported for eval rollout" - - global EVAL_PROMPT_DATASET - - cache_key = dataset_cfg.cache_key + (args.hf_checkpoint, args.apply_chat_template) - if cache_key not in EVAL_PROMPT_DATASET: - tokenizer = load_tokenizer(args.hf_checkpoint, trust_remote_code=True) - processor = load_processor(args.hf_checkpoint, trust_remote_code=True) - EVAL_PROMPT_DATASET[cache_key] = Dataset( - path=dataset_cfg.path, - tokenizer=tokenizer, - processor=processor, - max_length=args.eval_max_prompt_len, - prompt_key=dataset_cfg.input_key, - label_key=dataset_cfg.label_key, - multimodal_keys=args.multimodal_keys, - metadata_key=dataset_cfg.metadata_key, - tool_key=dataset_cfg.tool_key, - apply_chat_template=args.apply_chat_template, - apply_chat_template_kwargs=args.apply_chat_template_kwargs, - ) - dataset = EVAL_PROMPT_DATASET[cache_key] - - base_sampling_params = dict( - temperature=dataset_cfg.temperature, - top_p=dataset_cfg.top_p, - top_k=dataset_cfg.top_k, - max_new_tokens=dataset_cfg.max_response_len, - stop=args.rollout_stop, - stop_token_ids=args.rollout_stop_token_ids, - skip_special_tokens=args.rollout_skip_special_tokens, - no_stop_trim=True, - spaces_between_special_tokens=False, - ) - - tasks = [] - # do multiple samples for eval prompts - sample_index = 0 - for _i, prompt_sample in enumerate(dataset.samples): - for j in range(dataset_cfg.n_samples_per_eval_prompt): - # use the same prompt for multiple samples - sample = copy.deepcopy(prompt_sample) - sample.index = sample_index - sample_index += 1 - sample.metadata = dataset_cfg.inject_metadata(getattr(sample, "metadata", None)) - sample.generate_function_path = getattr(dataset_cfg, "custom_generate_function_path", None) - sampling_params = base_sampling_params - if getattr(args, "sglang_enable_deterministic_inference", False): - sampling_params = base_sampling_params.copy() - sampling_params["sampling_seed"] = args.rollout_seed + j - tasks.append( - asyncio.create_task( - generate_and_rm( - args, - sample, - sampling_params=sampling_params, - evaluation=True, - ) - ) - ) - - data = [] - do_print = True - pbar = tqdm(total=len(tasks), desc=f"Eval {dataset_cfg.name}", disable=not do_print) - for coro in asyncio.as_completed(tasks): - sample = await coro - if do_print: - logged_sample = sample[0] if isinstance(sample, list) else sample - logger.info( - "eval_rollout_single_dataset example data: " - f"{[str(logged_sample.prompt) + logged_sample.response]} " - f"reward={logged_sample.reward}" - ) - do_print = False - if isinstance(sample, list): - data.extend(sample) - else: - data.append(sample) - pbar.update(1) - pbar.close() - - data.sort(key=lambda sample: sample.index) - - reward_key = args.eval_reward_key or args.reward_key - return { - dataset_cfg.name: { - "rewards": [sample.reward if not reward_key else sample.reward[reward_key] for sample in data], - "truncated": [sample.status == Sample.Status.TRUNCATED for sample in data], - "samples": data, - } - } - - -def generate_rollout( - args: Namespace, rollout_id: int, data_source: Any, evaluation: bool = False -) -> RolloutFnTrainOutput | RolloutFnEvalOutput: - """An example to implement the generate_rollout function for an rule based rm rollout generation. - - Args: - args: the whole args - rollout_id: int, the id of the rollout, used for deterministic data generation - data_source: the data source to get and store samples - evaluation: bool, whether the rollout is for evaluation or not - - Returns: - RolloutFnTrainOutput | RolloutFnEvalOutput: the output of the rollout - """ - assert args.rollout_global_dataset - if evaluation: - output, _ = run(eval_rollout(args, rollout_id)) - return output - - output, aborted_samples = run(generate_rollout_async(args, rollout_id, data_source.get_samples)) - if aborted_samples: - data_source.add_samples(aborted_samples) - return output diff --git a/slime/rollout/vllm_rollout.py b/slime/rollout/vllm_rollout.py index 4a79b8425..da9484e5f 100644 --- a/slime/rollout/vllm_rollout.py +++ b/slime/rollout/vllm_rollout.py @@ -12,7 +12,7 @@ from typing import Any import numpy as np -import vllm_router # noqa: F401 — same side-effect as ``import sglang_router`` in sglang rollout +import vllm_router # noqa: F401 — ensures vllm-router is importable on startup from tqdm import tqdm from slime.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainOutput @@ -28,7 +28,7 @@ load_processor, load_tokenizer, ) -from slime.utils.trace_utils import trace_function, trace_span +from slime.utils.trace_utils import build_vllm_meta_trace_attrs, trace_function, trace_span from slime.utils.types import Sample from .rm_hub import async_rm, batched_async_rm @@ -86,7 +86,7 @@ def _prepare_prompt_ids(sample: Sample, tokenizer, processor: Any) -> list[int]: def _base_dataset_prompt_ids(sample: Sample, tokenizer, processor: Any) -> list[int]: """Token ids for the dataset prompt only (never reuse ``sample.tokens``). - Used for partial-continuation budgeting to match ``dev_vllm`` ``sglang_rollout``: + Used for partial-continuation budgeting: ``max_new_tokens -= len(sample.tokens) - len(base_prompt_ids)`` when ``sample.response`` is non-empty. """ raw_multimodal_inputs = sample.multimodal_inputs or {} @@ -102,24 +102,24 @@ def get_model_url(args: Namespace, model_name: str, endpoint: str = "/inference/ """Return the router URL for a named model. Use this in custom rollout functions to route requests to a specific - model when multiple models are deployed via ``--sglang-config``:: + model when multiple models are deployed via ``--vllm-config``:: url = get_model_url(args, "ref", "/inference/v1/generate") resp = await post(url, json=payload) Falls back to the default router if *model_name* is not found or - ``sglang_model_routers`` is not set. + ``vllm_model_routers`` is not set. """ - routers = getattr(args, "sglang_model_routers", None) + routers = getattr(args, "vllm_model_routers", None) if routers and model_name in routers: ip, port = routers[model_name] return f"http://{ip}:{port}{endpoint}" - return f"http://{args.router_ip}:{args.router_port}{endpoint}" + return f"http://{args.vllm_router_ip}:{args.vllm_router_port}{endpoint}" async def _router_worker_urls(args: Namespace) -> list[str]: - """Resolve worker base URLs from the vLLM router (same HTTP shape as SGLang router).""" - base = f"http://{args.router_ip}:{args.router_port}" + """Resolve worker base URLs from the vLLM router.""" + base = f"http://{args.vllm_router_ip}:{args.vllm_router_port}" try: response = await get(f"{base}/workers") return [worker["url"] for worker in response["workers"]] @@ -243,7 +243,7 @@ def _inference_generate_tokens_and_logprobs(choice: dict[str, Any]) -> tuple[lis def _align_engine_tokens_and_logprobs( new_response_tokens: list[int], new_response_log_probs: list[float] ) -> tuple[list[int], list[float]]: - """Pad or truncate logprobs so ``len(.) == len(new_response_tokens)`` (SGLang always has matched OTL pairs).""" + """Pad or truncate logprobs so ``len(.) == len(new_response_tokens)``.""" n = len(new_response_tokens) if n == 0: return [], [] @@ -281,7 +281,7 @@ def __init__(self, args: Namespace) -> None: spaces_between_special_tokens=False, ) - if getattr(args, "sglang_enable_deterministic_inference", False): + if getattr(args, "vllm_enable_deterministic_inference", False): sampling_seed_base = args.rollout_seed self.group_sampling_seeds = [sampling_seed_base + i for i in range(args.n_samples_per_prompt)] @@ -390,7 +390,7 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A assert isinstance(sample.prompt, str) state = GenerateState(args) - base = f"http://{args.router_ip}:{args.router_port}" + base = f"http://{args.vllm_router_ip}:{args.vllm_router_port}" assert ( sample.status == Sample.Status.PENDING or sample.status == Sample.Status.ABORTED @@ -417,11 +417,15 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A if not sample.tokens: sample.tokens = prompt_ids - # Use session_id for consistent hashing routing (SGLang Model Gateway) + # Use session_id for consistent_hash routing. vllm-router's + # ConsistentHashPolicy.extract_hash_key_from_headers recognizes + # x-session-id / x-user-id / x-tenant-id (and session_params.session_id + # in the request body) — see vllm_router_rs::policies::consistent_hash. + # NB: vllm-router's policy enum is "consistent_hash" (no -ing). headers = None if sample.session_id: - if getattr(args, "router_policy", None) == "consistent_hashing": - headers = {"X-SMG-Routing-Key": sample.session_id} + if getattr(args, "router_policy", None) == "consistent_hash": + headers = {"x-session-id": sample.session_id} if images: # Disaggregated MM flow: render (preprocess) then tokens-only generate — see vLLM docs @@ -456,8 +460,15 @@ async def generate(args: Namespace, sample: Sample, sampling_params: dict[str, A "token_ids": token_ids, "sampling_params": inference_sampling_params, } - with trace_span(sample, "vllm_inference_generate", attrs={"max_new_tokens": params["max_new_tokens"]}): + with trace_span( + sample, "vllm_inference_generate", attrs={"max_new_tokens": params["max_new_tokens"]} + ) as span: output = await post(url, payload, headers=headers) + # Enrich the span with response meta (finish reason + token usage), mirroring + # SGLang's build_sglang_meta_trace_attrs. trace_span yields the raw target when + # tracing is disabled, so guard on `.update`. + if hasattr(span, "update"): + span.update(build_vllm_meta_trace_attrs(output)) choice = output["choices"][0] skip_sp = params.get("skip_special_tokens") @@ -590,7 +601,7 @@ async def generate_and_rm_group( tasks = [] for idx, sample in enumerate(group): current_sampling_params = sampling_params.copy() - if getattr(args, "sglang_enable_deterministic_inference", False): + if getattr(args, "vllm_enable_deterministic_inference", False): seed = state.group_sampling_seeds[idx] current_sampling_params["seed"] = seed tasks.append( @@ -821,7 +832,7 @@ async def eval_rollout_single_dataset( sample.metadata = dataset_cfg.inject_metadata(getattr(sample, "metadata", None)) sample.generate_function_path = getattr(dataset_cfg, "custom_generate_function_path", None) sampling_params = base_sampling_params - if getattr(args, "sglang_enable_deterministic_inference", False): + if getattr(args, "vllm_enable_deterministic_inference", False): sampling_params = base_sampling_params.copy() sampling_params["seed"] = args.rollout_seed + j tasks.append( diff --git a/slime/utils/arguments.py b/slime/utils/arguments.py index ec525e6c3..a7976afa9 100644 --- a/slime/utils/arguments.py +++ b/slime/utils/arguments.py @@ -55,7 +55,7 @@ def add_cluster_arguments(parser): "--rollout-num-gpus-per-engine", type=int, default=1, - help="Number of GPUs per inference engine, just like the tp_size in sglang.", + help="Number of GPUs per inference engine, just like the tp_size in vllm.", ) parser.add_argument( "--num-gpus-per-node", @@ -134,7 +134,7 @@ def add_train_arguments(parser): "--megatron-to-hf-mode", choices=["raw", "bridge"], default="raw", - help="The method to convert megatron weights to hugging face weights for SGLang.", + help="The method to convert megatron weights to hugging face weights for vLLM.", ) parser.add_argument( "--custom-model-provider-path", @@ -211,8 +211,8 @@ def add_rollout_arguments(parser): default=None, help=( "The huggingface checkpoint of the trained model. " - "This is used to initialize sglang and also provide the tokenizer. " - "Note that, we will always update the parameters in sglang with that of megatron before training, " + "This is used to initialize vLLM and also provide the tokenizer. " + "Note that, we will always update the parameters in vLLM with that of megatron before training, " "so you only need to provide a huggingface checkpoint that has the same architecture as the model you want to train. " "It doesn't necessary need to contain the most up-to-date parameters." ), @@ -278,7 +278,7 @@ def add_rollout_arguments(parser): default=None, help=( "The maximum length of the response for the inference engine during rollout. " - "It is basically `max_tokens` in sglang." + "It is basically `max_tokens` in vllm." ), ) parser.add_argument( @@ -448,7 +448,7 @@ def add_rollout_arguments(parser): "--rollout-external", action="store_true", default=False, - help="Use external SGLang instances instead of launching them inside the framework.", + help="Use external vLLM instances instead of launching them inside the framework.", ) parser.add_argument( "--rollout-external-engine-addrs", @@ -985,11 +985,11 @@ def add_on_policy_distillation_arguments(parser): parser.add_argument( "--opd-type", type=str, - choices=["sglang", "megatron"], + choices=["vllm", "megatron"], default=None, help=( "Type of on-policy distillation. " - "'sglang': Teacher log-probs are obtained from external SGLang server during rollout. " + "'vllm': Teacher log-probs are obtained from external vLLM server during rollout. " "'megatron': Teacher model is loaded via --opd-teacher-load and forwarded during training." ), ) @@ -1011,12 +1011,23 @@ def add_on_policy_distillation_arguments(parser): parser.add_argument( "--opd-teacher-ckpt-step", type=int, default=None, help="The checkpoint step for OPD teacher model." ) + parser.add_argument( + "--opd-teacher-model", + type=str, + default=None, + help=( + "Served model name of the OPD teacher for --opd-type=vllm. Sent as the `model` field to " + "the teacher's /inference/v1/generate endpoint. Optional: if unset, the field is omitted " + "(single-model teacher servers use their loaded model). Never defaults to the student " + "hf_checkpoint, which would mis-name a teacher!=student server." + ), + ) return parser def add_router_arguments(parser): # vllm-router's full CLI surface (~30 knobs: policy, cache_threshold, # retries, health-check, …) under `--router-*` prefix (collision-safe). - # exclude_host_port=True because vime owns `--router-ip / --router-port` + # exclude_host_port=True because vime owns `--vllm-router-ip / --vllm-router-port` # (defined in slime/backends/vllm_utils/arguments.py:add_vllm_router_arguments). RouterArgs.add_cli_args(parser, use_router_prefix=True, exclude_host_port=True) return parser @@ -1601,7 +1612,7 @@ def slime_validate_args(args): # Validate on-policy distillation (OPD) arguments if args.use_opd: if args.opd_type is None: - raise ValueError("--opd-type must be specified when --use-opd is enabled. Choose 'sglang' or 'megatron'.") + raise ValueError("--opd-type must be specified when --use-opd is enabled. Choose 'vllm' or 'megatron'.") if args.opd_type == "megatron": if args.opd_teacher_load is None: @@ -1619,11 +1630,11 @@ def slime_validate_args(args): "please make sure it is a valid megatron checkpoint directory." ) - elif args.opd_type == "sglang": + elif args.opd_type == "vllm": if args.opd_teacher_load is not None: raise ValueError( - "--opd-teacher-load should not be set when --opd-type=sglang. " - "In sglang mode, teacher log-probs are obtained from external server during rollout." + "--opd-teacher-load should not be set when --opd-type=vllm. " + "In vllm mode, teacher log-probs are obtained from external server during rollout." ) else: # If OPD is not enabled, opd_teacher_load should not be set @@ -1702,7 +1713,7 @@ def slime_validate_args(args): if args.load_debug_rollout_data is not None: logger.info( f"load_debug_rollout_data {args.load_debug_rollout_data} is set, " - "will not instantiate sglang servers and will only run the training process." + "will not instantiate vLLM servers and will only run the training process." ) args.debug_train_only = True diff --git a/slime/utils/external_utils/command_utils.py b/slime/utils/external_utils/command_utils.py index cdf2cbe0b..859843571 100644 --- a/slime/utils/external_utils/command_utils.py +++ b/slime/utils/external_utils/command_utils.py @@ -108,17 +108,21 @@ def execute_train( master_addr = os.environ.get("MASTER_ADDR", "127.0.0.1") exec_command( - "pkill -9 sglang; " + # vLLM renames its VRAM-holding subprocesses via set_process_title() + # (VLLM::EngineCore, VLLM::Worker_TP*, vllm::router), so their cmdline no + # longer contains "vllm serve". Matching only the launcher would leave the + # engine/worker children holding GPU memory and leak it into the next run. + # Match both the launcher and the renamed children; the [v]/[M] bracket + # trick keeps this pattern from matching pkill's own cmdline. This targets + # exactly the vLLM tree, so the old indiscriminate `pkill -9 python` + # (dangerous on colocate/shared nodes) is no longer needed. + "pkill -9 -f '[v]llm serve|VLL[M]::'; " "sleep 3; " f"{'' if external_ray else 'ray stop --force; '}" f"{'' if external_ray else 'pkill -9 ray; '}" - # cannot be run in CI, o/w kill the parent script - # TODO: do we really need this kill? (or can we instead kill slime) - # "pkill -9 python; " "pkill -9 slime; " "sleep 3; " f"{'' if external_ray else 'pkill -9 ray; '}" - # "pkill -9 python; " "pkill -9 slime; " "pkill -9 redis; " "true; " @@ -222,7 +226,6 @@ def create_run_id() -> str: _warned_bool_env_var_keys = set() -# copied from SGLang def get_bool_env_var(name: str, default: str = "false") -> bool: value = os.getenv(name, default) value = value.lower() diff --git a/slime/utils/http_utils.py b/slime/utils/http_utils.py index 4f910c40e..e7fe05348 100644 --- a/slime/utils/http_utils.py +++ b/slime/utils/http_utils.py @@ -114,32 +114,17 @@ def _wrap_ipv6(host): return host -def run_router(payload): - """Start the HTTP router gateway in a child process. - - ``payload`` is either ``(impl, router_args)`` with ``impl`` in ``{"sglang","vllm"}``, - or a legacy single ``router_args`` object (treated as ``sglang``). - """ +def run_router(router_args): + """Start the vllm-router HTTP gateway in a child process.""" try: - if isinstance(payload, tuple) and len(payload) == 2: - impl, router_args = payload - else: - impl, router_args = "sglang", payload - - if impl == "vllm": - from vllm_router.launch_router import launch_router - else: - from sglang_router.launch_router import launch_router + from vllm_router.launch_router import launch_router router = launch_router(router_args) if router is None: return 1 return 0 except Exception: - logger.exception( - "run_router failed (impl=%s). For vllm-router, ensure the package is installed and RouterArgs are valid.", - payload[0] if isinstance(payload, tuple) and len(payload) == 2 else "sglang", - ) + logger.exception("run_router failed. Ensure vllm-router is installed and RouterArgs are valid.") return 1 @@ -225,7 +210,7 @@ def init_http_client(args): _http_client = httpx.AsyncClient( limits=httpx.Limits(max_connections=_client_concurrency), timeout=httpx.Timeout(None), - trust_env=False, # internal SGLang comm only — never route through system proxy + trust_env=False, # internal vLLM comm only — never route through system proxy ) # Optionally initialize distributed POST via Ray without changing interfaces @@ -260,7 +245,7 @@ def __init__(self, concurrency: int): self._client = httpx.AsyncClient( limits=httpx.Limits(max_connections=max(1, concurrency)), timeout=httpx.Timeout(None), - trust_env=False, # internal SGLang comm only — never route through system proxy + trust_env=False, # internal vLLM comm only — never route through system proxy ) async def do_post(self, url, payload, max_retries=60, headers=None): diff --git a/slime/utils/logging_utils.py b/slime/utils/logging_utils.py index 11348a407..408be24c9 100644 --- a/slime/utils/logging_utils.py +++ b/slime/utils/logging_utils.py @@ -8,7 +8,6 @@ _LOGGER_CONFIGURED = False -# ref: SGLang def configure_logger(prefix: str = ""): global _LOGGER_CONFIGURED if _LOGGER_CONFIGURED: diff --git a/slime/utils/trace_utils.py b/slime/utils/trace_utils.py index b51bc5b93..2dcaf8fb5 100644 --- a/slime/utils/trace_utils.py +++ b/slime/utils/trace_utils.py @@ -14,22 +14,6 @@ from slime.utils.types import Sample TRACE_VERSION = 1 -SGLANG_TRACE_META_KEYS = ( - "prompt_tokens", - "completion_tokens", - "cached_tokens", - "pd_prefill_bootstrap_queue_duration", - "pd_prefill_forward_duration", - "pd_prefill_transfer_queue_duration", - "pd_prefill_retry_count", - "pd_decode_prealloc_duration", - "pd_decode_transfer_duration", - "pd_decode_forward_duration", - "pd_bootstrap_duration", - "pd_alloc_waiting_duration", - "pd_transfer_speed_gb_s", - "pd_transfer_total_mb", -) logger = logging.getLogger(__name__) _TRACE_STACK: contextvars.ContextVar[tuple[tuple[str, str], ...]] = contextvars.ContextVar( @@ -121,9 +105,22 @@ def _new_span_id() -> str: return uuid.uuid4().hex -def build_sglang_meta_trace_attrs(meta: dict[str, Any]) -> dict[str, Any]: - attrs = {key: meta[key] for key in SGLANG_TRACE_META_KEYS if key in meta and meta[key] is not None} - attrs["finish_reason"] = meta["finish_reason"]["type"] +def build_vllm_meta_trace_attrs(output: dict[str, Any]) -> dict[str, Any]: + """Trace-span attributes from a vLLM ``/inference/v1/generate`` response. + + vLLM exposes far less per-request meta than SGLang's ``meta_info`` (no + PD-disaggregation timing): only the finish reason and token usage are + available on the response. Richer request timing lives in vLLM's own OTLP + traces (``gen_ai.latency.*``, enabled via ``--otlp-traces-endpoint``). + """ + attrs: dict[str, Any] = {} + choices = output.get("choices") or [] + if choices and choices[0].get("finish_reason") is not None: + attrs["finish_reason"] = choices[0]["finish_reason"] + usage = output.get("usage") or {} + for key in ("prompt_tokens", "completion_tokens", "cached_tokens"): + if usage.get(key) is not None: + attrs[key] = usage[key] return attrs diff --git a/slime/utils/types.py b/slime/utils/types.py index 8a6e1be7a..9ac3916c2 100644 --- a/slime/utils/types.py +++ b/slime/utils/types.py @@ -45,7 +45,7 @@ class Status(Enum): # metadata used during training, e.g., what loss to use for this sample. train_metadata: dict | None = None - # Session ID for consistent hashing routing (used when router policy is consistent_hashing) + # Session ID for consistent_hash routing (used when --router-policy consistent_hash) session_id: str | None = None non_generation_time: float = 0.0 # time spent in non-generation steps diff --git a/slime/utils/wandb_utils.py b/slime/utils/wandb_utils.py index 2ec859de5..d6fcf18fd 100644 --- a/slime/utils/wandb_utils.py +++ b/slime/utils/wandb_utils.py @@ -85,7 +85,7 @@ def reinit_wandb_primary_with_open_metrics(args, router_addr): The primary wandb init happens before rollout servers start (to obtain ``wandb_run_id`` for secondary processes). This function is called *after* servers are up so the router address is available for scraping - SGLang Prometheus metrics via the primary process's stats monitor. + vLLM Prometheus metrics via the primary process's stats monitor. """ if not args.use_wandb or _is_offline_mode(args): return @@ -97,15 +97,7 @@ def reinit_wandb_primary_with_open_metrics(args, router_addr): if wandb_run_id is None: return - import sglang_router - - if "slime" not in sglang_router.__version__: - logger.warning( - "Only customized sglang_router from https://github.com/zhuzilin/sgl-router supports uploading metrics." - ) - return - - logger.info(f"Re-initializing primary W&B with SGLang metrics at {router_addr}.") + logger.info(f"Re-initializing primary W&B with vLLM metrics at {router_addr}.") wandb.finish() @@ -119,10 +111,12 @@ def reinit_wandb_primary_with_open_metrics(args, router_addr): mode="shared", x_primary=True, x_stats_open_metrics_endpoints={ - "sgl_engine": f"{router_addr}/engine_metrics", + # router_addr already includes the /metrics path on the vllm-router + # prometheus port (see RolloutManager._get_metrics_router_addr). + "vllm_engine": router_addr, }, x_stats_open_metrics_filters={ - "sgl_engine.*": {}, + "vllm_engine.*": {}, }, ), } diff --git a/tests/unit/backends/vllm_utils/conftest.py b/tests/unit/backends/vllm_utils/conftest.py index 567c49777..298fcbc53 100644 --- a/tests/unit/backends/vllm_utils/conftest.py +++ b/tests/unit/backends/vllm_utils/conftest.py @@ -12,8 +12,8 @@ def vllm_args() -> SimpleNamespace: return SimpleNamespace( rollout_external=True, hf_checkpoint="/tmp/model", - router_ip=None, - router_port=None, + vllm_router_ip=None, + vllm_router_port=None, vllm_weight_transfer_timeout_sec=900.0, num_gpus_per_node=8, rollout_num_gpus_per_engine=4, diff --git a/tests/unit/backends/vllm_utils/test_arguments.py b/tests/unit/backends/vllm_utils/test_arguments.py index 841437f84..cde1b0e6a 100644 --- a/tests/unit/backends/vllm_utils/test_arguments.py +++ b/tests/unit/backends/vllm_utils/test_arguments.py @@ -145,7 +145,7 @@ def _ns(**overrides): vllm_data_parallel_size=1, vllm_pipeline_parallel_size=1, rollout_num_gpus_per_engine=4, - router_ip=None, + vllm_router_ip=None, ) base.update(overrides) return SimpleNamespace(**base) @@ -178,39 +178,39 @@ def test_validate_args_pp_indivisible_asserts(args_mod): @pytest.mark.unit def test_validate_args_router_ipv6_wrapped(args_mod): - ns = _ns(router_ip="::1") + ns = _ns(vllm_router_ip="::1") args_mod.validate_args(ns) - assert ns.router_ip == "[::1]" + assert ns.vllm_router_ip == "[::1]" @pytest.mark.unit def test_validate_args_router_ipv6_already_wrapped_unchanged(args_mod): - ns = _ns(router_ip="[::1]") + ns = _ns(vllm_router_ip="[::1]") args_mod.validate_args(ns) - assert ns.router_ip == "[::1]" + assert ns.vllm_router_ip == "[::1]" @pytest.mark.unit def test_validate_args_router_ipv4_unchanged(args_mod): - ns = _ns(router_ip="127.0.0.1") + ns = _ns(vllm_router_ip="127.0.0.1") args_mod.validate_args(ns) - assert ns.router_ip == "127.0.0.1" + assert ns.vllm_router_ip == "127.0.0.1" @pytest.mark.unit def test_validate_args_router_none_noop(args_mod): - ns = _ns(router_ip=None) + ns = _ns(vllm_router_ip=None) args_mod.validate_args(ns) - assert ns.router_ip is None + assert ns.vllm_router_ip is None @pytest.mark.unit -def test_add_vllm_router_arguments_registers_router_prefix(args_mod): +def test_add_vllm_router_arguments_registers_vllm_prefix(args_mod): parser = argparse.ArgumentParser(add_help=False) args_mod.add_vllm_router_arguments(parser) flags = {s for a in parser._actions for s in a.option_strings} - assert "--router-ip" in flags - assert "--router-port" in flags + assert "--vllm-router-ip" in flags + assert "--vllm-router-port" in flags assert "--router-request-timeout-secs" in flags @@ -219,21 +219,21 @@ def test_add_vllm_router_arguments_dests(args_mod): parser = argparse.ArgumentParser(add_help=False) args_mod.add_vllm_router_arguments(parser) dests = {a.dest for a in parser._actions if a.option_strings} - assert "router_ip" in dests - assert "router_port" in dests + assert "vllm_router_ip" in dests + assert "vllm_router_port" in dests assert "router_request_timeout_secs" in dests @pytest.mark.unit -def test_add_vllm_router_arguments_old_names_removed(args_mod): +def test_add_vllm_router_arguments_no_unprefixed_names(args_mod): parser = argparse.ArgumentParser(add_help=False) args_mod.add_vllm_router_arguments(parser) flags = {s for a in parser._actions for s in a.option_strings} dests = {a.dest for a in parser._actions if a.option_strings} - assert "--vllm-router-ip" not in flags - assert "--vllm-router-port" not in flags - assert "vllm_router_ip" not in dests - assert "vllm_router_port" not in dests + assert "--router-ip" not in flags + assert "--router-port" not in flags + assert "router_ip" not in dests + assert "router_port" not in dests @pytest.mark.unit @@ -241,21 +241,22 @@ def test_add_vllm_router_arguments_parses_real_values(args_mod): parser = argparse.ArgumentParser(add_help=False) args_mod.add_vllm_router_arguments(parser) parsed, _ = parser.parse_known_args( - ["--router-ip", "10.0.0.1", "--router-port", "8000", "--router-request-timeout-secs", "30"] + ["--vllm-router-ip", "10.0.0.1", "--vllm-router-port", "8000", "--router-request-timeout-secs", "30"] ) - assert parsed.router_ip == "10.0.0.1" - assert parsed.router_port == 8000 + assert parsed.vllm_router_ip == "10.0.0.1" + assert parsed.vllm_router_port == 8000 assert parsed.router_request_timeout_secs == 30 @pytest.mark.unit -def test_orchestration_dests_use_new_names(args_mod): - assert "router_ip" in args_mod._VIME_ORCHESTRATION_DESTS - assert "router_port" in args_mod._VIME_ORCHESTRATION_DESTS +def test_orchestration_dests_use_vllm_prefix(args_mod): + assert "vllm_router_ip" in args_mod._VIME_ORCHESTRATION_DESTS + assert "vllm_router_port" in args_mod._VIME_ORCHESTRATION_DESTS assert "router_request_timeout_secs" in args_mod._VIME_ORCHESTRATION_DESTS assert "vllm_weight_transfer_timeout_sec" in args_mod._VIME_ORCHESTRATION_DESTS - assert "vllm_router_ip" not in args_mod._VIME_ORCHESTRATION_DESTS - assert "vllm_router_port" not in args_mod._VIME_ORCHESTRATION_DESTS + assert "router_ip" not in args_mod._VIME_ORCHESTRATION_DESTS + assert "router_port" not in args_mod._VIME_ORCHESTRATION_DESTS + assert "vllm_router_request_timeout_secs" not in args_mod._VIME_ORCHESTRATION_DESTS @pytest.mark.unit @@ -272,8 +273,8 @@ def test_add_vllm_arguments_parses_weight_transfer_timeout(args_mod, monkeypatch def _realistic_add_vllm_arguments(parser): parser.add_argument("--vllm-gpu-memory-utilization", dest="vllm_gpu_memory_utilization", type=float, default=0.92) parser.add_argument("--vllm-enforce-eager", dest="vllm_enforce_eager", action="store_true", default=False) - parser.add_argument("--router-ip", dest="router_ip", type=str, default=None) - parser.add_argument("--router-port", dest="router_port", type=int, default=None) + parser.add_argument("--vllm-router-ip", dest="vllm_router_ip", type=str, default=None) + parser.add_argument("--vllm-router-port", dest="vllm_router_port", type=int, default=None) parser.add_argument("--vllm-server-concurrency", dest="vllm_server_concurrency", type=int, default=512) parser.add_argument( "--vllm-weight-transfer-timeout-sec", @@ -300,8 +301,8 @@ def test_action_table_excludes_orchestration(args_mod, monkeypatch): table = args_mod.get_vllm_cli_action_table() assert "vllm_gpu_memory_utilization" in table assert "vllm_enforce_eager" in table - assert "router_ip" not in table - assert "router_port" not in table + assert "vllm_router_ip" not in table + assert "vllm_router_port" not in table assert "vllm_server_concurrency" not in table assert "vllm_weight_transfer_timeout_sec" not in table diff --git a/tools/analyze_profile.py b/tools/analyze_profile.py index de511744f..eee2023db 100644 --- a/tools/analyze_profile.py +++ b/tools/analyze_profile.py @@ -1,8 +1,8 @@ #!/usr/bin/env python3 """ -SGLang Decode Profile Analyzer -============================== -Analyzes PyTorch profiler traces (.trace.json.gz) from SGLang decode workers. +vLLM Decode Profile Analyzer +============================ +Analyzes PyTorch profiler traces (.trace.json.gz) from vLLM decode workers. Usage: python tools/analyze_profile.py --profile-dir profiles/20260303_052303_my_run @@ -587,7 +587,7 @@ def print_analysis(r: TraceAnalysis): "portion that runs OUTSIDE the CUDA graph (DeepEP dispatch/combine + NCCL allgather).", "The 3-launch pattern per step = (1) pre-MoE graph, (2) post-MoE graph, (3) MoE-expert graph.", "Optimization: try increasing decode batch size to amortize graph launch overhead per token.", - "Check if `--sglang-disable-cuda-graph` helps isolate whether the overhead is in graph " + "Check if `--enforce-eager` helps isolate whether the overhead is in graph " "management vs. actual compute.", "Consider padding batch sizes to avoid frequent graph re-capture for different sizes.", ], @@ -633,7 +633,7 @@ def print_analysis(r: TraceAnalysis): ), "Increase batch size to improve GPU SM occupancy — many kernels are memory-bound at small batch.", "Speculative decoding could help if generation is latency-bound.", - "Verify `--sglang-mem-fraction-static` is set high enough for large KV cache.", + "Verify `--gpu-memory-utilization` is set high enough for large KV cache.", ], ) ) @@ -676,7 +676,7 @@ def print_cross_rank_summary(analyses: list[TraceAnalysis]): def main(): - parser = argparse.ArgumentParser(description="Analyze SGLang decode profile traces") + parser = argparse.ArgumentParser(description="Analyze vLLM decode profile traces") parser.add_argument("--profile-dir", type=str, required=True, help="Directory containing .trace.json.gz files") parser.add_argument("--rank", type=int, default=None, help="Specific rank to analyze (default: first file)") parser.add_argument("--all-ranks", action="store_true", help="Analyze all ranks and show comparison") diff --git a/tools/convert_hf_to_torch_dist.py b/tools/convert_hf_to_torch_dist.py index d14c22fc9..377486c53 100644 --- a/tools/convert_hf_to_torch_dist.py +++ b/tools/convert_hf_to_torch_dist.py @@ -25,7 +25,7 @@ def add_convertion_args(parser): "--megatron-to-hf-mode", choices=["raw", "bridge"], default="raw", - help="The method to convert megatron weights to hugging face weights for SGLang.", + help="The method to convert megatron weights to hugging face weights for vLLM.", ) try: parser.add_argument("--padded-vocab-size", type=int, default=None) diff --git a/tools/convert_torch_dist_to_hf_parallel.py b/tools/convert_torch_dist_to_hf_parallel.py index 763254d42..66a29e56f 100644 --- a/tools/convert_torch_dist_to_hf_parallel.py +++ b/tools/convert_torch_dist_to_hf_parallel.py @@ -272,9 +272,6 @@ def process_param(args, model_name, name, param, vocab_size=None): def save_tensors(args, model_name, state_dict, output_dir, chunk_size, vocab_size=None, max_workers=1, worker_id=None): - # for slime update_weight compatible - args.sglang_enable_ep_moe = False - print(f"start saving to {output_dir}") os.makedirs(output_dir, exist_ok=True) param_list = list(get_named_params(args, state_dict)) diff --git a/tools/profile_rollout.py b/tools/profile_rollout.py index 8869a1993..c3802d2eb 100644 --- a/tools/profile_rollout.py +++ b/tools/profile_rollout.py @@ -43,10 +43,10 @@ def stop_profile(worker_url): def main(): - parser = argparse.ArgumentParser(description="Automate SGLang profiling across all workers via router.") + parser = argparse.ArgumentParser(description="Automate vLLM profiling across all workers via router.") parser.add_argument("--router-url", type=str, required=True, help="Router URL (e.g., http://127.0.0.1:3000)") parser.add_argument("--action", type=str, choices=["start", "stop"], default="start", help="Action to perform") - parser.add_argument("--output-dir", type=str, default="/tmp/sglang_profile", help="Output directory for traces") + parser.add_argument("--output-dir", type=str, default="/tmp/vllm_profile", help="Output directory for traces") parser.add_argument("--num-steps", type=int, default=3, help="Number of steps to profile (default: 3)") parser.add_argument("--activities", type=str, nargs="+", default=["GPU"], help="Activities to profile (CPU, GPU)") parser.add_argument("--profile-by-stage", action="store_true", help="Profile by stage (prefill/decode)") diff --git a/tools/replay_openai_jsonl.py b/tools/replay_openai_jsonl.py index fd7b4a2f9..1bd31d9f4 100644 --- a/tools/replay_openai_jsonl.py +++ b/tools/replay_openai_jsonl.py @@ -45,7 +45,7 @@ def record(self, result: dict[str, Any]) -> None: def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( - description="Replay OpenAI-compatible chat completion payloads from a JSONL file against an SGLang router.", + description="Replay OpenAI-compatible chat completion payloads from a JSONL file against a vLLM router.", formatter_class=argparse.ArgumentDefaultsHelpFormatter, epilog=( "Examples:\n" diff --git a/train.py b/train.py index 2404a0bbd..628f927b3 100644 --- a/train.py +++ b/train.py @@ -12,11 +12,11 @@ def train(args): pgs = create_placement_groups(args) init_tracking(args) - # create the rollout manager, with sglang engines inside. + # create the rollout manager, with vLLM engines inside. # need to initialize rollout manager first to calculate num_rollout rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) - # Update primary W&B with SGLang metrics endpoint now that servers are up. + # Update primary W&B with vLLM metrics endpoint now that servers are up. router_addr = ray.get(rollout_manager.get_metrics_router_addr.remote()) update_tracking_open_metrics(args, router_addr) diff --git a/train_async.py b/train_async.py index 6960bd055..8e4bc6629 100644 --- a/train_async.py +++ b/train_async.py @@ -14,11 +14,11 @@ def train(args): pgs = create_placement_groups(args) init_tracking(args) - # create the rollout manager, with sglang engines inside. + # create the rollout manager, with vLLM engines inside. # need to initialize rollout manager first to calculate num_rollout rollout_manager, num_rollout_per_epoch = create_rollout_manager(args, pgs["rollout"]) - # Update primary W&B with SGLang metrics endpoint now that servers are up. + # Update primary W&B with vLLM metrics endpoint now that servers are up. router_addr = ray.get(rollout_manager.get_metrics_router_addr.remote()) update_tracking_open_metrics(args, router_addr)