[Deps] Upgrade CUDA PyTorch stack to 2.13 (#28836)
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
This commit is contained in:
co-authored by
Brayden Zhong
parent
4ad990ba7d
commit
434e646282
@@ -120,7 +120,7 @@ jobs:
|
|||||||
--platform linux/amd64 \
|
--platform linux/amd64 \
|
||||||
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=12.9.1 \
|
--build-arg CUDA_VERSION=12.9.2 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg GRACE_BLACKWELL=0 \
|
--build-arg GRACE_BLACKWELL=0 \
|
||||||
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
||||||
@@ -149,7 +149,7 @@ jobs:
|
|||||||
--platform linux/amd64 \
|
--platform linux/amd64 \
|
||||||
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=13.0.1 \
|
--build-arg CUDA_VERSION=13.0.3 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg GRACE_BLACKWELL=0 \
|
--build-arg GRACE_BLACKWELL=0 \
|
||||||
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
||||||
@@ -232,7 +232,7 @@ jobs:
|
|||||||
--platform linux/arm64 \
|
--platform linux/arm64 \
|
||||||
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=12.9.1 \
|
--build-arg CUDA_VERSION=12.9.2 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg GRACE_BLACKWELL=1 \
|
--build-arg GRACE_BLACKWELL=1 \
|
||||||
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
||||||
@@ -261,7 +261,7 @@ jobs:
|
|||||||
--platform linux/arm64 \
|
--platform linux/arm64 \
|
||||||
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
--output type=image,name=${{ inputs.image_repo }},push-by-digest=true,name-canonical=true,push=true \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=13.0.1 \
|
--build-arg CUDA_VERSION=13.0.3 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg GRACE_BLACKWELL=1 \
|
--build-arg GRACE_BLACKWELL=1 \
|
||||||
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
||||||
|
|||||||
@@ -68,6 +68,7 @@ on:
|
|||||||
# defines must be redeclared here for the called job to see the same context.
|
# defines must be redeclared here for the called job to see the same context.
|
||||||
env:
|
env:
|
||||||
SGLANG_IS_IN_CI: true
|
SGLANG_IS_IN_CI: true
|
||||||
|
NCCL_NVLS_ENABLE: "0"
|
||||||
# base-a runs the fwd-occupancy sanity kit; async-assert probes add GPU work
|
# base-a runs the fwd-occupancy sanity kit; async-assert probes add GPU work
|
||||||
# and skew the occupancy measurement, so keep them off in that stage.
|
# and skew the occupancy measurement, so keep them off in that stage.
|
||||||
SGLANG_ENABLE_ASYNC_ASSERT: ${{ startsWith(inputs.self_name, 'base-a') && 'false' || 'true' }}
|
SGLANG_ENABLE_ASYNC_ASSERT: ${{ startsWith(inputs.self_name, 'base-a') && 'false' || 'true' }}
|
||||||
|
|||||||
@@ -195,7 +195,7 @@ jobs:
|
|||||||
--output type=image,name=${IMAGE_REF},push=true \
|
--output type=image,name=${IMAGE_REF},push=true \
|
||||||
--target framework_final \
|
--target framework_final \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=13.0.1 \
|
--build-arg CUDA_VERSION=13.0.3 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg CMAKE_BUILD_PARALLEL_LEVEL=$(nproc) \
|
--build-arg CMAKE_BUILD_PARALLEL_LEVEL=$(nproc) \
|
||||||
--build-arg GRACE_BLACKWELL=1 \
|
--build-arg GRACE_BLACKWELL=1 \
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ on:
|
|||||||
# The github context (including github.event_name) IS inherited from the caller.
|
# The github context (including github.event_name) IS inherited from the caller.
|
||||||
env:
|
env:
|
||||||
SGLANG_IS_IN_CI: true
|
SGLANG_IS_IN_CI: true
|
||||||
|
NCCL_NVLS_ENABLE: "0"
|
||||||
SGLANG_CUDA_COREDUMP: "1"
|
SGLANG_CUDA_COREDUMP: "1"
|
||||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
||||||
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ on:
|
|||||||
# The github context (including github.event_name) IS inherited from the caller.
|
# The github context (including github.event_name) IS inherited from the caller.
|
||||||
env:
|
env:
|
||||||
SGLANG_IS_IN_CI: true
|
SGLANG_IS_IN_CI: true
|
||||||
|
NCCL_NVLS_ENABLE: "0"
|
||||||
SGLANG_CUDA_COREDUMP: "1"
|
SGLANG_CUDA_COREDUMP: "1"
|
||||||
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
||||||
SKIP_PR_TEST_HEALTH_CHECK: ${{ inputs.skip_pr_test_health_check == 'true' }}
|
SKIP_PR_TEST_HEALTH_CHECK: ${{ inputs.skip_pr_test_health_check == 'true' }}
|
||||||
|
|||||||
@@ -59,6 +59,8 @@ env:
|
|||||||
SGLANG_CUDA_COREDUMP: "1"
|
SGLANG_CUDA_COREDUMP: "1"
|
||||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
||||||
SKIP_PR_TEST_HEALTH_CHECK: ${{ (inputs.skip_pr_test_health_check == true || inputs.test_parallel_dispatch == true || inputs.run_all_tests == true) && 'true' || 'false' }}
|
SKIP_PR_TEST_HEALTH_CHECK: ${{ (inputs.skip_pr_test_health_check == true || inputs.test_parallel_dispatch == true || inputs.run_all_tests == true) && 'true' || 'false' }}
|
||||||
|
# TEMP: force DeepEP rebuilds during the initial Torch 2.13 rollout so every CI runner replaces its Torch 2.11 build; remove after the runner fleet has cycled.
|
||||||
|
FORCE_REBUILD_DEEPEP: '1'
|
||||||
# Schedule / main-branch dispatch / workflow_call from main use refs/heads/main; PR events use refs/pull/*/merge
|
# Schedule / main-branch dispatch / workflow_call from main use refs/heads/main; PR events use refs/pull/*/merge
|
||||||
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
PR_TEST_BYPASS_MAINTENANCE_ON_MAIN: ${{ github.ref == 'refs/heads/main' && 'true' || 'false' }}
|
||||||
USE_VENV: false
|
USE_VENV: false
|
||||||
|
|||||||
@@ -176,7 +176,7 @@ jobs:
|
|||||||
--load \
|
--load \
|
||||||
-t "$IMAGE" \
|
-t "$IMAGE" \
|
||||||
-f docker/Dockerfile \
|
-f docker/Dockerfile \
|
||||||
--build-arg CUDA_VERSION=13.0.1 \
|
--build-arg CUDA_VERSION=13.0.3 \
|
||||||
--build-arg BUILD_TYPE=all \
|
--build-arg BUILD_TYPE=all \
|
||||||
--build-arg GRACE_BLACKWELL=0 \
|
--build-arg GRACE_BLACKWELL=0 \
|
||||||
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
--build-arg INSTALL_FLASHINFER_JIT_CACHE=1 \
|
||||||
|
|||||||
@@ -27,9 +27,10 @@ concurrency:
|
|||||||
cancel-in-progress: true
|
cancel-in-progress: true
|
||||||
|
|
||||||
env:
|
env:
|
||||||
|
NCCL_NVLS_ENABLE: "0"
|
||||||
# Must match TORCH_VER in docker/sgl-deep-gemm.Dockerfile — the wheel's
|
# Must match TORCH_VER in docker/sgl-deep-gemm.Dockerfile — the wheel's
|
||||||
# pre-compiled _C.so links against this torch ABI.
|
# pre-compiled _C.so links against this torch ABI.
|
||||||
TORCH_VER: "2.11.0"
|
TORCH_VER: "2.13.0"
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
build-cu129-matrix:
|
build-cu129-matrix:
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ env:
|
|||||||
SGLANG_ENABLE_ASYNC_ASSERT: true
|
SGLANG_ENABLE_ASYNC_ASSERT: true
|
||||||
SGLANG_CUDA_COREDUMP: "1"
|
SGLANG_CUDA_COREDUMP: "1"
|
||||||
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
SGLANG_JIT_DEEPGEMM_FAST_WARMUP: true
|
||||||
# TEMP: rebuild deepep against the new torch for torch-211-merge PR only — revert before merging to main.
|
# TEMP: force DeepEP rebuilds during the initial Torch 2.13 rollout so every CI runner replaces its Torch 2.11 build; remove after the runner fleet has cycled.
|
||||||
FORCE_REBUILD_DEEPEP: '1'
|
FORCE_REBUILD_DEEPEP: '1'
|
||||||
|
|
||||||
permissions:
|
permissions:
|
||||||
|
|||||||
+21
-27
@@ -1,4 +1,4 @@
|
|||||||
ARG CUDA_VERSION=13.0.1
|
ARG CUDA_VERSION=13.0.3
|
||||||
FROM nvidia/cuda:${CUDA_VERSION}-cudnn-devel-ubuntu24.04 AS base
|
FROM nvidia/cuda:${CUDA_VERSION}-cudnn-devel-ubuntu24.04 AS base
|
||||||
|
|
||||||
ARG TARGETARCH
|
ARG TARGETARCH
|
||||||
@@ -10,9 +10,9 @@ ARG HOPPER_SBO=0
|
|||||||
ARG HOPPER_SBO_DEEPEP_COMMIT=9f2fc4b3182a51044ae7ecb6610f7c9c3258c4d6
|
ARG HOPPER_SBO_DEEPEP_COMMIT=9f2fc4b3182a51044ae7ecb6610f7c9c3258c4d6
|
||||||
ARG DEEPEP_COMMIT=9af0e0d0e74f3577af1979c9b9e1ac2cad0104ee
|
ARG DEEPEP_COMMIT=9af0e0d0e74f3577af1979c9b9e1ac2cad0104ee
|
||||||
ARG BUILD_AND_DOWNLOAD_PARALLEL=8
|
ARG BUILD_AND_DOWNLOAD_PARALLEL=8
|
||||||
ARG SGL_KERNEL_VERSION=0.4.5
|
ARG SGL_KERNEL_VERSION=0.4.6.post1
|
||||||
ARG SGL_VERSION
|
ARG SGL_VERSION
|
||||||
ARG SGL_DEEP_GEMM_VERSION=0.1.5.post1
|
ARG SGL_DEEP_GEMM_VERSION=0.1.5.post2
|
||||||
ARG USE_LATEST_SGLANG=0
|
ARG USE_LATEST_SGLANG=0
|
||||||
ARG GDRCOPY_VERSION=2.5.1
|
ARG GDRCOPY_VERSION=2.5.1
|
||||||
ARG PIP_DEFAULT_INDEX
|
ARG PIP_DEFAULT_INDEX
|
||||||
@@ -194,19 +194,18 @@ RUN curl --proto '=https' --tlsv1.2 --retry 3 --retry-delay 2 -sSf https://sh.ru
|
|||||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||||
python3 -m pip install --upgrade pip setuptools wheel html5lib six \
|
python3 -m pip install --upgrade pip setuptools wheel html5lib six \
|
||||||
&& case "$CUDA_VERSION" in \
|
&& case "$CUDA_VERSION" in \
|
||||||
12.6.1) CUINDEX=126 ;; \
|
12.6.3) CUINDEX=126 ;; \
|
||||||
12.8.1) CUINDEX=128 ;; \
|
12.9.2) CUINDEX=129 ;; \
|
||||||
12.9.1) CUINDEX=129 ;; \
|
13.0.3) CUINDEX=130 ;; \
|
||||||
13.0.1) CUINDEX=130 ;; \
|
|
||||||
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
||||||
esac \
|
esac \
|
||||||
&& if [ "$CUDA_VERSION" = "12.6.1" ]; then \
|
&& if [ "$CUDA_VERSION" = "12.6.3" ]; then \
|
||||||
python3 -m pip install https://${GITHUB_ARTIFACTORY}/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu124-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \
|
python3 -m pip install https://${GITHUB_ARTIFACTORY}/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu124-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \
|
||||||
; \
|
; \
|
||||||
elif [ "$CUDA_VERSION" = "12.8.1" ] || [ "$CUDA_VERSION" = "12.9.1" ]; then \
|
elif [ "$CUDA_VERSION" = "12.9.2" ]; then \
|
||||||
python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu129-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \
|
python3 -m pip install https://${GITHUB_ARTIFACTORY}/sgl-project/whl/releases/download/v${SGL_KERNEL_VERSION}/sglang_kernel-${SGL_KERNEL_VERSION}+cu129-cp310-abi3-manylinux2014_$(uname -m).whl --force-reinstall --no-deps \
|
||||||
; \
|
; \
|
||||||
elif [ "$CUDA_VERSION" = "13.0.1" ]; then \
|
elif [ "$CUDA_VERSION" = "13.0.3" ]; then \
|
||||||
# --no-deps prevents pip from pulling torch from default PyPI
|
# --no-deps prevents pip from pulling torch from default PyPI
|
||||||
python3 -m pip install sglang-kernel==${SGL_KERNEL_VERSION} --force-reinstall --no-deps \
|
python3 -m pip install sglang-kernel==${SGL_KERNEL_VERSION} --force-reinstall --no-deps \
|
||||||
; \
|
; \
|
||||||
@@ -229,10 +228,9 @@ COPY proto /tmp/sglang_deps/proto
|
|||||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||||
--mount=type=cache,target=/root/.cargo/registry \
|
--mount=type=cache,target=/root/.cargo/registry \
|
||||||
case "$CUDA_VERSION" in \
|
case "$CUDA_VERSION" in \
|
||||||
12.6.1) CUINDEX=126 ;; \
|
12.6.3) CUINDEX=126 ;; \
|
||||||
12.8.1) CUINDEX=128 ;; \
|
12.9.2) CUINDEX=129 ;; \
|
||||||
12.9.1) CUINDEX=129 ;; \
|
13.0.3) CUINDEX=130 ;; \
|
||||||
13.0.1) CUINDEX=130 ;; \
|
|
||||||
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
||||||
esac \
|
esac \
|
||||||
&& cd /tmp/sglang_deps/python \
|
&& cd /tmp/sglang_deps/python \
|
||||||
@@ -251,7 +249,7 @@ RUN --mount=type=cache,target=/root/.cache/pip \
|
|||||||
pip list --format=freeze | awk -F'==' '/-cu13(==|$)/ {print $1}' \
|
pip list --format=freeze | awk -F'==' '/-cu13(==|$)/ {print $1}' \
|
||||||
| xargs -r python3 -m pip uninstall -y && \
|
| xargs -r python3 -m pip uninstall -y && \
|
||||||
python3 -m pip install --index-url https://download.pytorch.org/whl/cu${CUINDEX} \
|
python3 -m pip install --index-url https://download.pytorch.org/whl/cu${CUINDEX} \
|
||||||
torch==2.11.0 torchvision==0.26.0 torchaudio==2.11.0 --force-reinstall; \
|
torch==2.13.0 torchvision==0.28.0 torchaudio==2.11.0 --force-reinstall; \
|
||||||
python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_DEEP_GEMM_VERSION}/sgl_deep_gemm-${SGL_DEEP_GEMM_VERSION}+cu129-py3-none-manylinux2014_$(uname -m).whl --force-reinstall; \
|
python3 -m pip install https://github.com/sgl-project/whl/releases/download/v${SGL_DEEP_GEMM_VERSION}/sgl_deep_gemm-${SGL_DEEP_GEMM_VERSION}+cu129-py3-none-manylinux2014_$(uname -m).whl --force-reinstall; \
|
||||||
fi \
|
fi \
|
||||||
&& cd /sgl-workspace \
|
&& cd /sgl-workspace \
|
||||||
@@ -315,13 +313,10 @@ RUN set -eux; \
|
|||||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||||
cd /build/DeepEP && \
|
cd /build/DeepEP && \
|
||||||
case "$CUDA_VERSION" in \
|
case "$CUDA_VERSION" in \
|
||||||
12.6.1) \
|
12.6.3) \
|
||||||
CHOSEN_TORCH_CUDA_ARCH_LIST='9.0' \
|
CHOSEN_TORCH_CUDA_ARCH_LIST='9.0' \
|
||||||
;; \
|
;; \
|
||||||
12.8.1) \
|
12.9.2|13.0.3) \
|
||||||
CHOSEN_TORCH_CUDA_ARCH_LIST='9.0;10.0' \
|
|
||||||
;; \
|
|
||||||
12.9.1|13.0.1) \
|
|
||||||
CHOSEN_TORCH_CUDA_ARCH_LIST='9.0;10.0;10.3' \
|
CHOSEN_TORCH_CUDA_ARCH_LIST='9.0;10.0;10.3' \
|
||||||
;; \
|
;; \
|
||||||
*) \
|
*) \
|
||||||
@@ -370,10 +365,9 @@ ARG FLASHINFER_VERSION
|
|||||||
# Stage jit-cache/cubin artifacts into /flashinfer_jit_output for clean COPY later
|
# Stage jit-cache/cubin artifacts into /flashinfer_jit_output for clean COPY later
|
||||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||||
case "$CUDA_VERSION" in \
|
case "$CUDA_VERSION" in \
|
||||||
12.6.1) CUINDEX=126 ;; \
|
12.6.3) CUINDEX=126 ;; \
|
||||||
12.8.1) CUINDEX=128 ;; \
|
12.9.2) CUINDEX=129 ;; \
|
||||||
12.9.1) CUINDEX=129 ;; \
|
13.0.3) CUINDEX=130 ;; \
|
||||||
13.0.1) CUINDEX=130 ;; \
|
|
||||||
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
*) echo "Unsupported CUDA version: $CUDA_VERSION" && exit 1 ;; \
|
||||||
esac \
|
esac \
|
||||||
&& mkdir -p /flashinfer_jit_output \
|
&& mkdir -p /flashinfer_jit_output \
|
||||||
@@ -616,10 +610,10 @@ RUN --mount=type=cache,target=/root/.cache/pip \
|
|||||||
# to avoid shipping wrong-CUDA libs on cu13 images.
|
# to avoid shipping wrong-CUDA libs on cu13 images.
|
||||||
RUN --mount=type=cache,target=/root/.cache/pip if [ "${CUDA_VERSION%%.*}" = "12" ]; then \
|
RUN --mount=type=cache,target=/root/.cache/pip if [ "${CUDA_VERSION%%.*}" = "12" ]; then \
|
||||||
python3 -m pip install nixl nixl-cu12 --no-deps ; \
|
python3 -m pip install nixl nixl-cu12 --no-deps ; \
|
||||||
python3 -m pip install cuda-python==12.9 ; \
|
python3 -m pip install "cuda-python>=12,<13" ; \
|
||||||
elif [ "${CUDA_VERSION%%.*}" = "13" ]; then \
|
elif [ "${CUDA_VERSION%%.*}" = "13" ]; then \
|
||||||
python3 -m pip install nixl nixl-cu13 --no-deps ; \
|
python3 -m pip install nixl nixl-cu13 --no-deps ; \
|
||||||
python3 -m pip install cuda-python==13.2.0 ; \
|
python3 -m pip install "cuda-python>=13,<14" ; \
|
||||||
fi
|
fi
|
||||||
|
|
||||||
# Add yank script
|
# Add yank script
|
||||||
|
|||||||
@@ -9,7 +9,7 @@
|
|||||||
# 2. DeepEP patch + rebuild:
|
# 2. DeepEP patch + rebuild:
|
||||||
# topk 11->16, SWITCH_HIDDEN += 3584, EP>8 SourceMeta alignment,
|
# topk 11->16, SWITCH_HIDDEN += 3584, EP>8 SourceMeta alignment,
|
||||||
# and cross-node timeout headroom; rebuilt for sm_90 and sm_100a only
|
# and cross-node timeout headroom; rebuilt for sm_90 and sm_100a only
|
||||||
# 3. DeepGEMM upgrade to 0.1.5.post1:
|
# 3. DeepGEMM upgrade to 0.1.5.post2:
|
||||||
# official MegaMoE runtime-JIT header with Kimi-K3 SiTU support
|
# official MegaMoE runtime-JIT header with Kimi-K3 SiTU support
|
||||||
# 4. FlashInfer CuTeDSL MLA DCP patch:
|
# 4. FlashInfer CuTeDSL MLA DCP patch:
|
||||||
# apply the seven runtime-file diffs; exclude tests absent from the wheel
|
# apply the seven runtime-file diffs; exclude tests absent from the wheel
|
||||||
@@ -25,7 +25,7 @@
|
|||||||
|
|
||||||
FROM lmsysorg/sglang:v0.5.16-cu129 AS base
|
FROM lmsysorg/sglang:v0.5.16-cu129 AS base
|
||||||
|
|
||||||
ARG SGL_DEEP_GEMM_VERSION="0.1.5.post1"
|
ARG SGL_DEEP_GEMM_VERSION="0.1.5.post2"
|
||||||
|
|
||||||
# Current Kimi-K3 source auto-discovers and builds its PyO3 extensions.
|
# Current Kimi-K3 source auto-discovers and builds its PyO3 extensions.
|
||||||
ARG RUST_VERSION="1.90.0"
|
ARG RUST_VERSION="1.90.0"
|
||||||
|
|||||||
@@ -10,7 +10,7 @@
|
|||||||
# topk 11->16, SWITCH_HIDDEN += 3584, EP>8 SourceMeta alignment,
|
# topk 11->16, SWITCH_HIDDEN += 3584, EP>8 SourceMeta alignment,
|
||||||
# cross-node timeout headroom, CUDA-13 cccl include; rebuilt for
|
# cross-node timeout headroom, CUDA-13 cccl include; rebuilt for
|
||||||
# sm_90, sm_100a, and sm_103a
|
# sm_90, sm_100a, and sm_103a
|
||||||
# 3. DeepGEMM upgrade to 0.1.5.post1:
|
# 3. DeepGEMM upgrade to 0.1.5.post2:
|
||||||
# official MegaMoE runtime-JIT header with Kimi-K3 SiTU support
|
# official MegaMoE runtime-JIT header with Kimi-K3 SiTU support
|
||||||
# 4. FlashInfer CuTeDSL MLA DCP patch:
|
# 4. FlashInfer CuTeDSL MLA DCP patch:
|
||||||
# apply the seven runtime-file diffs; exclude tests absent from the wheel
|
# apply the seven runtime-file diffs; exclude tests absent from the wheel
|
||||||
@@ -26,7 +26,7 @@
|
|||||||
|
|
||||||
FROM lmsysorg/sglang:v0.5.16 AS base
|
FROM lmsysorg/sglang:v0.5.16 AS base
|
||||||
|
|
||||||
ARG SGL_DEEP_GEMM_VERSION="0.1.5.post1"
|
ARG SGL_DEEP_GEMM_VERSION="0.1.5.post2"
|
||||||
|
|
||||||
# Current Kimi-K3 source auto-discovers and builds its PyO3 extensions.
|
# Current Kimi-K3 source auto-discovers and builds its PyO3 extensions.
|
||||||
ARG RUST_VERSION="1.90.0"
|
ARG RUST_VERSION="1.90.0"
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ ARG ARCH=x86_64
|
|||||||
ARG CUDA_VERSION=13.0
|
ARG CUDA_VERSION=13.0
|
||||||
ARG PYTHON_VERSION=3.12
|
ARG PYTHON_VERSION=3.12
|
||||||
ARG PYTHON_TAG=cp312-cp312
|
ARG PYTHON_TAG=cp312-cp312
|
||||||
ARG TORCH_VER=2.11.0
|
ARG TORCH_VER=2.13.0
|
||||||
ARG TVM_FFI_VER=0.1.11
|
ARG TVM_FFI_VER=0.1.11
|
||||||
ARG PIP_DEFAULT_INDEX=https://pypi.python.org/simple
|
ARG PIP_DEFAULT_INDEX=https://pypi.python.org/simple
|
||||||
ARG PYTORCH_MIRROR=download.pytorch.org
|
ARG PYTORCH_MIRROR=download.pytorch.org
|
||||||
|
|||||||
@@ -11,9 +11,9 @@ metatags:
|
|||||||
The command below uses `/tmp/huggingface` as an example.
|
The command below uses `/tmp/huggingface` as an example.
|
||||||
|
|
||||||
```
|
```
|
||||||
docker pull nvidia/cuda:12.9.1-devel-ubuntu22.04
|
docker pull nvidia/cuda:13.0.3-devel-ubuntu22.04
|
||||||
# Nvidia
|
# Nvidia
|
||||||
docker run --shm-size 128g -it -v /tmp/huggingface:/hf_home --gpus all nvidia/cuda:12.9.1-devel-ubuntu22.04 /bin/bash
|
docker run --shm-size 128g -it -v /tmp/huggingface:/hf_home --gpus all nvidia/cuda:13.0.3-devel-ubuntu22.04 /bin/bash
|
||||||
# AMD
|
# AMD
|
||||||
docker run --rm --device=/dev/kfd --device=/dev/dri --group-add video --shm-size 128g -it -v /tmp/huggingface:/hf_home lmsysorg/sglang:v0.5.8-rocm700-mi30x /bin/bash
|
docker run --rm --device=/dev/kfd --device=/dev/dri --group-add video --shm-size 128g -it -v /tmp/huggingface:/hf_home lmsysorg/sglang:v0.5.8-rocm700-mi30x /bin/bash
|
||||||
# AMD just the last 2 GPUs
|
# AMD just the last 2 GPUs
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ The major version of Cuda is 13 by default. To install sglang under Cuda 12 with
|
|||||||
pip install --upgrade pip
|
pip install --upgrade pip
|
||||||
pip install uv
|
pip install uv
|
||||||
uv pip install --prerelease=allow sglang
|
uv pip install --prerelease=allow sglang
|
||||||
uv pip install --force-reinstall torch==2.11.0 torchaudio==2.11.0 torchvision --index-url https://download.pytorch.org/whl/cu129
|
uv pip install --force-reinstall torch==2.13.0 torchaudio==2.11.0 torchvision --index-url https://download.pytorch.org/whl/cu129
|
||||||
uv pip install --force-reinstall sglang-kernel --index-url https://docs.sglang.ai/whl/cu129/
|
uv pip install --force-reinstall sglang-kernel --index-url https://docs.sglang.ai/whl/cu129/
|
||||||
uv pip install --force-reinstall sgl-deep-gemm --index-url https://docs.sglang.ai/whl/cu129/ --no-deps
|
uv pip install --force-reinstall sgl-deep-gemm --index-url https://docs.sglang.ai/whl/cu129/ --no-deps
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -67,19 +67,19 @@ dependencies = [
|
|||||||
"scipy",
|
"scipy",
|
||||||
"sentencepiece",
|
"sentencepiece",
|
||||||
"setproctitle",
|
"setproctitle",
|
||||||
"sgl-deep-gemm==0.1.5.post1",
|
"sgl-deep-gemm==0.1.5.post2",
|
||||||
"sglang-kernel==0.4.5",
|
"sglang-kernel==0.4.6.post1",
|
||||||
"smg-grpc-servicer>=0.5.0",
|
"smg-grpc-servicer>=0.5.0",
|
||||||
"soundfile==0.13.1",
|
"soundfile==0.13.1",
|
||||||
"tiktoken",
|
"tiktoken",
|
||||||
"tilelang==0.1.11",
|
"tilelang==0.1.11",
|
||||||
"timm==1.0.16",
|
"timm==1.0.16",
|
||||||
"tokenspeed_mla==0.1.8",
|
"tokenspeed_mla==0.1.8",
|
||||||
"torch==2.11.0",
|
"torch==2.13.0",
|
||||||
"torch_memory_saver>=0.0.9.post1",
|
"torch_memory_saver>=0.0.9.post1",
|
||||||
"torchao==0.17.0",
|
"torchao==0.17.0",
|
||||||
"torchaudio==2.11.0",
|
"torchaudio==2.11.0",
|
||||||
"torchcodec==0.11.1 ; sys_platform != 'linux' or (sys_platform == 'linux' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'armv7l')", # torchcodec 0.11.1 for torch 2.11.x (0.10 is ABI-incompatible: references the pre-2.11 c10::MessageLogger ctor signature). Not available on Linux ARM.
|
"torchcodec==0.15.0 ; sys_platform != 'linux' or (sys_platform == 'linux' and platform_machine != 'aarch64' and platform_machine != 'arm64' and platform_machine != 'armv7l')", # Not available on Linux ARM.
|
||||||
"torchvision",
|
"torchvision",
|
||||||
"tqdm",
|
"tqdm",
|
||||||
"transformers==5.12.1",
|
"transformers==5.12.1",
|
||||||
|
|||||||
@@ -766,6 +766,9 @@ def launch_disagg_role(server_args: ServerArgs):
|
|||||||
|
|
||||||
def dispatch_launch(server_args: ServerArgs):
|
def dispatch_launch(server_args: ServerArgs):
|
||||||
"""Route to the correct launch function based on --disagg-role."""
|
"""Route to the correct launch function based on --disagg-role."""
|
||||||
|
if "NCCL_NVLS_ENABLE" not in os.environ or server_args.enable_nccl_nvls:
|
||||||
|
os.environ["NCCL_NVLS_ENABLE"] = str(int(server_args.enable_nccl_nvls))
|
||||||
|
|
||||||
role = server_args.disagg_role
|
role = server_args.disagg_role
|
||||||
if role == RoleType.MONOLITHIC:
|
if role == RoleType.MONOLITHIC:
|
||||||
launch_server(server_args)
|
launch_server(server_args)
|
||||||
|
|||||||
@@ -321,6 +321,10 @@ class RMSNormNoWeight(CustomOp):
|
|||||||
return F.rms_norm(x, normalized_shape=(x.shape[-1],), eps=eps)
|
return F.rms_norm(x, normalized_shape=(x.shape[-1],), eps=eps)
|
||||||
|
|
||||||
def forward_cuda(self, x: torch.Tensor, eps: float) -> torch.Tensor:
|
def forward_cuda(self, x: torch.Tensor, eps: float) -> torch.Tensor:
|
||||||
|
# Torch 2.12+ runs rms_norm in fp32 under CUDA autocast. This operator
|
||||||
|
# historically preserved the activation dtype, and callers rely on
|
||||||
|
# that contract for both memory use and downstream kernel selection.
|
||||||
|
with torch.autocast(device_type="cuda", enabled=False):
|
||||||
return self.forward_native(x, eps=eps)
|
return self.forward_native(x, eps=eps)
|
||||||
|
|
||||||
def forward_npu(self, x: torch.Tensor, eps: float) -> torch.Tensor:
|
def forward_npu(self, x: torch.Tensor, eps: float) -> torch.Tensor:
|
||||||
|
|||||||
@@ -847,8 +847,17 @@ class LTX2Attention(nn.Module):
|
|||||||
else:
|
else:
|
||||||
if self.qk_norm:
|
if self.qk_norm:
|
||||||
assert self.q_norm is not None and self.k_norm is not None
|
assert self.q_norm is not None and self.k_norm is not None
|
||||||
q = self.q_norm(q)
|
# Torch 2.12+ places rms_norm on the autocast fp32 list. A
|
||||||
k = self.k_norm(k)
|
# cast after the norm preserves the attention contract but
|
||||||
|
# still runs the much slower fp32 kernel. Torch 2.11 ran
|
||||||
|
# this operation in the input dtype, so disable autocast
|
||||||
|
# around Q/K norm to preserve both its precision path and
|
||||||
|
# performance.
|
||||||
|
q_dtype = q.dtype
|
||||||
|
k_dtype = k.dtype
|
||||||
|
with torch.autocast(device_type=q.device.type, enabled=False):
|
||||||
|
q = self.q_norm(q).to(dtype=q_dtype)
|
||||||
|
k = self.k_norm(k).to(dtype=k_dtype)
|
||||||
|
|
||||||
if pe is not None and cos.dim() == 3:
|
if pe is not None and cos.dim() == 3:
|
||||||
q = apply_interleaved_rotary_emb(q, (cos, sin))
|
q = apply_interleaved_rotary_emb(q, (cos, sin))
|
||||||
|
|||||||
@@ -203,6 +203,7 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
|
|
||||||
# Distributed executor backend
|
# Distributed executor backend
|
||||||
nccl_port: Optional[int] = None
|
nccl_port: Optional[int] = None
|
||||||
|
enable_nccl_nvls: bool = False
|
||||||
|
|
||||||
# HuggingFace specific parameters
|
# HuggingFace specific parameters
|
||||||
trust_remote_code: bool = False
|
trust_remote_code: bool = False
|
||||||
@@ -1382,6 +1383,12 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
# Parallelism
|
# Parallelism
|
||||||
|
parser.add_argument(
|
||||||
|
"--enable-nccl-nvls",
|
||||||
|
action=StoreBoolean,
|
||||||
|
default=ServerArgs.enable_nccl_nvls,
|
||||||
|
help="Enable NCCL NVLS when available.",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--num-gpus",
|
"--num-gpus",
|
||||||
type=int,
|
type=int,
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ logger = init_logger(__name__)
|
|||||||
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
# NPU/ascend) is read from sgl-project/ci-data-diffusion, where the GT-gen workflows
|
||||||
# publish.
|
# publish.
|
||||||
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
SGL_TEST_FILES_CI_DATA_REPO = "sgl-project/ci-data-diffusion"
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "d05810e3ea3eff1d137dec723f6e66d9c11b470f"
|
SGL_TEST_FILES_CI_DATA_REVISION = "dc0e1bb34f2776313a259bcfab3e30daed85160e"
|
||||||
|
|
||||||
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
|
||||||
# when it's regenerated on its own cadence.
|
# when it's regenerated on its own cadence.
|
||||||
|
|||||||
@@ -2114,5 +2114,23 @@ class TestDisaggTransferBackendArgs(unittest.TestCase):
|
|||||||
self.assertEqual(args.disagg_transfer_backend, "mock")
|
self.assertEqual(args.disagg_transfer_backend, "mock")
|
||||||
|
|
||||||
|
|
||||||
|
class TestNcclNvlsArgs(unittest.TestCase):
|
||||||
|
def test_enable_nccl_nvls_cli_arg(self):
|
||||||
|
parser = FlexibleArgumentParser()
|
||||||
|
ServerArgs.add_cli_args(parser)
|
||||||
|
|
||||||
|
default_args, _ = parser.parse_known_args(["--model-path", "/fake"])
|
||||||
|
enabled_args, _ = parser.parse_known_args(
|
||||||
|
["--model-path", "/fake", "--enable-nccl-nvls"]
|
||||||
|
)
|
||||||
|
disabled_args, _ = parser.parse_known_args(
|
||||||
|
["--model-path", "/fake", "--enable-nccl-nvls", "false"]
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertFalse(default_args.enable_nccl_nvls)
|
||||||
|
self.assertTrue(enabled_args.enable_nccl_nvls)
|
||||||
|
self.assertFalse(disabled_args.enable_nccl_nvls)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -1651,7 +1651,7 @@ def _set_envs_and_config(server_args: ServerArgs):
|
|||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
assert_pkg_version(
|
assert_pkg_version(
|
||||||
"sglang-kernel",
|
"sglang-kernel",
|
||||||
"0.4.5",
|
"0.4.6.post1",
|
||||||
"Please reinstall the latest version with `pip install sglang-kernel --force-reinstall`",
|
"Please reinstall the latest version with `pip install sglang-kernel --force-reinstall`",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -7,17 +7,14 @@ from __future__ import annotations
|
|||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from triton_kernels.matmul_ogs import (
|
from triton_kernels.matmul import (
|
||||||
FlexCtx,
|
FlexCtx,
|
||||||
FnSpecs,
|
FnSpecs,
|
||||||
FusedActivation,
|
FusedActivation,
|
||||||
GatherIndx,
|
|
||||||
PrecisionConfig,
|
PrecisionConfig,
|
||||||
RoutingData,
|
matmul,
|
||||||
ScatterIndx,
|
|
||||||
matmul_ogs,
|
|
||||||
)
|
)
|
||||||
from triton_kernels.matmul_ogs_details.opt_flags import update_opt_flags_constraints
|
from triton_kernels.matmul_details.opt_flags import update_opt_flags_constraints
|
||||||
from triton_kernels.numerics import InFlexData
|
from triton_kernels.numerics import InFlexData
|
||||||
from triton_kernels.swiglu import swiglu_fn
|
from triton_kernels.swiglu import swiglu_fn
|
||||||
from triton_kernels.tensor import FP4
|
from triton_kernels.tensor import FP4
|
||||||
@@ -35,6 +32,8 @@ else:
|
|||||||
from sgl_kernel import gelu_and_mul, silu_and_mul
|
from sgl_kernel import gelu_and_mul, silu_and_mul
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from triton_kernels.tensor_details.ragged_tensor import RaggedTensorMetadata
|
||||||
|
|
||||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
|
|
||||||
@@ -86,15 +85,17 @@ def triton_kernel_moe_forward(
|
|||||||
|
|
||||||
assert TopKOutputChecker.format_is_triton_kernels(topk_output)
|
assert TopKOutputChecker.format_is_triton_kernels(topk_output)
|
||||||
|
|
||||||
routing_data, gather_idx, scatter_idx = topk_output
|
a_ragged_metadata, gather_idx, scatter_idx, gate_scal, n_expts_act = topk_output
|
||||||
|
|
||||||
return triton_kernel_fused_experts(
|
return triton_kernel_fused_experts(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
w1,
|
w1,
|
||||||
w2,
|
w2,
|
||||||
routing_data,
|
a_ragged_metadata,
|
||||||
gather_idx,
|
gather_idx,
|
||||||
scatter_idx,
|
scatter_idx,
|
||||||
|
gate_scal,
|
||||||
|
n_expts_act,
|
||||||
inplace=False, # triton kernel doesn't support inplace
|
inplace=False, # triton kernel doesn't support inplace
|
||||||
activation=moe_runner_config.activation,
|
activation=moe_runner_config.activation,
|
||||||
apply_router_weight_on_input=apply_router_weight_on_input,
|
apply_router_weight_on_input=apply_router_weight_on_input,
|
||||||
@@ -115,9 +116,11 @@ def triton_kernel_fused_experts(
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
w1: torch.Tensor,
|
w1: torch.Tensor,
|
||||||
w2: torch.Tensor,
|
w2: torch.Tensor,
|
||||||
routing_data: RoutingData,
|
a_ragged_metadata: RaggedTensorMetadata,
|
||||||
gather_indx: GatherIndx,
|
gather_indx: torch.Tensor,
|
||||||
scatter_indx: ScatterIndx,
|
scatter_indx: Optional[torch.Tensor],
|
||||||
|
gate_scal: torch.Tensor,
|
||||||
|
n_expts_act: int,
|
||||||
inplace: bool = False,
|
inplace: bool = False,
|
||||||
activation: str = "silu",
|
activation: str = "silu",
|
||||||
apply_router_weight_on_input: bool = False,
|
apply_router_weight_on_input: bool = False,
|
||||||
@@ -162,7 +165,6 @@ def triton_kernel_fused_experts(
|
|||||||
|
|
||||||
M, K = hidden_states.shape
|
M, K = hidden_states.shape
|
||||||
E, _, N = w1.shape
|
E, _, N = w1.shape
|
||||||
n_expts_act = routing_data.n_expts_act
|
|
||||||
dtype = hidden_states.dtype
|
dtype = hidden_states.dtype
|
||||||
|
|
||||||
if global_num_experts == -1:
|
if global_num_experts == -1:
|
||||||
@@ -170,16 +172,16 @@ def triton_kernel_fused_experts(
|
|||||||
|
|
||||||
# consistent with default implementation
|
# consistent with default implementation
|
||||||
intermediate_cache2 = torch.empty(
|
intermediate_cache2 = torch.empty(
|
||||||
(M * n_expts_act, N // 2), device="cuda", dtype=dtype
|
(M * n_expts_act, N // 2), device=hidden_states.device, dtype=dtype
|
||||||
)
|
)
|
||||||
|
|
||||||
intermediate_cache1 = matmul_ogs(
|
intermediate_cache1 = matmul(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
w1,
|
w1,
|
||||||
None,
|
None,
|
||||||
routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
gather_indx=gather_indx,
|
gather_indx=gather_indx,
|
||||||
gammas=routing_data.gate_scal if apply_router_weight_on_input else None,
|
gammas=gate_scal if apply_router_weight_on_input else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
if activation == "silu":
|
if activation == "silu":
|
||||||
@@ -189,13 +191,13 @@ def triton_kernel_fused_experts(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported FusedMoe activation: {activation}")
|
raise ValueError(f"Unsupported FusedMoe activation: {activation}")
|
||||||
|
|
||||||
intermediate_cache3 = matmul_ogs(
|
intermediate_cache3 = matmul(
|
||||||
intermediate_cache2,
|
intermediate_cache2,
|
||||||
w2,
|
w2,
|
||||||
None,
|
None,
|
||||||
routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
scatter_indx=scatter_indx,
|
scatter_indx=scatter_indx,
|
||||||
gammas=None if apply_router_weight_on_input else routing_data.gate_scal,
|
gammas=None if apply_router_weight_on_input else gate_scal,
|
||||||
)
|
)
|
||||||
|
|
||||||
return intermediate_cache3
|
return intermediate_cache3
|
||||||
@@ -226,7 +228,7 @@ def triton_kernel_moe_with_bias_forward(
|
|||||||
|
|
||||||
assert TopKOutputChecker.format_is_triton_kernels(topk_output)
|
assert TopKOutputChecker.format_is_triton_kernels(topk_output)
|
||||||
|
|
||||||
routing_data, gather_idx, scatter_idx = topk_output
|
a_ragged_metadata, gather_idx, scatter_idx, gate_scal, n_expts_act = topk_output
|
||||||
|
|
||||||
return triton_kernel_fused_experts_with_bias(
|
return triton_kernel_fused_experts_with_bias(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
@@ -236,9 +238,11 @@ def triton_kernel_moe_with_bias_forward(
|
|||||||
w2=w2,
|
w2=w2,
|
||||||
w2_pcg=w2_pcg,
|
w2_pcg=w2_pcg,
|
||||||
b2=b2,
|
b2=b2,
|
||||||
routing_data=routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
gather_indx=gather_idx,
|
gather_indx=gather_idx,
|
||||||
scatter_indx=scatter_idx,
|
scatter_indx=scatter_idx,
|
||||||
|
gate_scal=gate_scal,
|
||||||
|
n_expts_act=n_expts_act,
|
||||||
inplace=False, # triton kernel doesn't support inplace
|
inplace=False, # triton kernel doesn't support inplace
|
||||||
activation=moe_runner_config.activation,
|
activation=moe_runner_config.activation,
|
||||||
apply_router_weight_on_input=apply_router_weight_on_input,
|
apply_router_weight_on_input=apply_router_weight_on_input,
|
||||||
@@ -264,9 +268,11 @@ def triton_kernel_fused_experts_with_bias(
|
|||||||
w2: torch.Tensor,
|
w2: torch.Tensor,
|
||||||
w2_pcg,
|
w2_pcg,
|
||||||
b2: torch.Tensor,
|
b2: torch.Tensor,
|
||||||
routing_data: RoutingData,
|
a_ragged_metadata: RaggedTensorMetadata,
|
||||||
gather_indx: GatherIndx,
|
gather_indx: torch.Tensor,
|
||||||
scatter_indx: ScatterIndx,
|
scatter_indx: Optional[torch.Tensor],
|
||||||
|
gate_scal: torch.Tensor,
|
||||||
|
n_expts_act: int,
|
||||||
inplace: bool = False,
|
inplace: bool = False,
|
||||||
activation: str = "silu",
|
activation: str = "silu",
|
||||||
apply_router_weight_on_input: bool = False,
|
apply_router_weight_on_input: bool = False,
|
||||||
@@ -315,7 +321,6 @@ def triton_kernel_fused_experts_with_bias(
|
|||||||
|
|
||||||
M, K = hidden_states.shape
|
M, K = hidden_states.shape
|
||||||
E, _, N = w1.shape
|
E, _, N = w1.shape
|
||||||
n_expts_act = routing_data.n_expts_act
|
|
||||||
|
|
||||||
if global_num_experts == -1:
|
if global_num_experts == -1:
|
||||||
global_num_experts = E
|
global_num_experts = E
|
||||||
@@ -335,35 +340,24 @@ def triton_kernel_fused_experts_with_bias(
|
|||||||
(gemm1_alpha, gemm1_clamp_limit),
|
(gemm1_alpha, gemm1_clamp_limit),
|
||||||
)
|
)
|
||||||
|
|
||||||
intermediate_cache = torch.empty(
|
intermediate_cache = matmul(
|
||||||
(1, M * n_expts_act, N // 2),
|
|
||||||
device=hidden_states.device,
|
|
||||||
dtype=hidden_states.dtype,
|
|
||||||
)
|
|
||||||
output = torch.empty(
|
|
||||||
(1, M, K), device=hidden_states.device, dtype=hidden_states.dtype
|
|
||||||
)
|
|
||||||
|
|
||||||
matmul_ogs(
|
|
||||||
hidden_states,
|
hidden_states,
|
||||||
w1,
|
w1,
|
||||||
b1,
|
b1,
|
||||||
routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
gather_indx=gather_indx,
|
gather_indx=gather_indx,
|
||||||
precision_config=w1_pcg,
|
precision_config=w1_pcg,
|
||||||
gammas=routing_data.gate_scal if apply_router_weight_on_input else None,
|
gammas=gate_scal if apply_router_weight_on_input else None,
|
||||||
fused_activation=act,
|
fused_activation=act,
|
||||||
y=intermediate_cache,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
matmul_ogs(
|
output = matmul(
|
||||||
intermediate_cache.view(M * n_expts_act, N // 2),
|
intermediate_cache.view(M * n_expts_act, N // 2),
|
||||||
w2,
|
w2,
|
||||||
b2,
|
b2,
|
||||||
routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
scatter_indx=scatter_indx,
|
scatter_indx=scatter_indx,
|
||||||
precision_config=w2_pcg,
|
precision_config=w2_pcg,
|
||||||
gammas=None if apply_router_weight_on_input else routing_data.gate_scal,
|
gammas=None if apply_router_weight_on_input else gate_scal,
|
||||||
y=output,
|
|
||||||
)
|
)
|
||||||
return output.view(M, K)
|
return output.view(-1, K)
|
||||||
|
|||||||
@@ -19,12 +19,8 @@ from sglang.srt.layers.moe.moe_runner.base import (
|
|||||||
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from triton_kernels.matmul_ogs import (
|
from triton_kernels.matmul import PrecisionConfig
|
||||||
GatherIndx,
|
from triton_kernels.tensor_details.ragged_tensor import RaggedTensorMetadata
|
||||||
PrecisionConfig,
|
|
||||||
RoutingData,
|
|
||||||
ScatterIndx,
|
|
||||||
)
|
|
||||||
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.standard import (
|
from sglang.srt.layers.moe.token_dispatcher.standard import (
|
||||||
StandardCombineInput,
|
StandardCombineInput,
|
||||||
@@ -42,9 +38,11 @@ class TritonKernelsRunnerInput(RunnerInput):
|
|||||||
"""Input bundle passed to the triton-kernels runner core."""
|
"""Input bundle passed to the triton-kernels runner core."""
|
||||||
|
|
||||||
hidden_states: torch.Tensor
|
hidden_states: torch.Tensor
|
||||||
routing_data: RoutingData
|
a_ragged_metadata: RaggedTensorMetadata
|
||||||
gather_indx: GatherIndx
|
gather_indx: torch.Tensor
|
||||||
scatter_indx: ScatterIndx
|
scatter_indx: torch.Tensor
|
||||||
|
gate_scal: torch.Tensor
|
||||||
|
n_expts_act: int
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def runner_backend(self) -> MoeRunnerBackend:
|
def runner_backend(self) -> MoeRunnerBackend:
|
||||||
@@ -102,9 +100,11 @@ class TritonKernelsRunnerCore(MoeRunnerCore):
|
|||||||
hidden_states = runner_input.hidden_states
|
hidden_states = runner_input.hidden_states
|
||||||
|
|
||||||
common_kwargs = dict(
|
common_kwargs = dict(
|
||||||
routing_data=runner_input.routing_data,
|
a_ragged_metadata=runner_input.a_ragged_metadata,
|
||||||
gather_indx=runner_input.gather_indx,
|
gather_indx=runner_input.gather_indx,
|
||||||
scatter_indx=None if self.config.no_combine else runner_input.scatter_indx,
|
scatter_indx=None if self.config.no_combine else runner_input.scatter_indx,
|
||||||
|
gate_scal=runner_input.gate_scal,
|
||||||
|
n_expts_act=runner_input.n_expts_act,
|
||||||
inplace=False,
|
inplace=False,
|
||||||
activation=self.config.activation,
|
activation=self.config.activation,
|
||||||
apply_router_weight_on_input=self.config.apply_router_weight_on_input,
|
apply_router_weight_on_input=self.config.apply_router_weight_on_input,
|
||||||
@@ -137,12 +137,14 @@ class TritonKernelsRunnerCore(MoeRunnerCore):
|
|||||||
**common_kwargs,
|
**common_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.config.no_combine:
|
|
||||||
tokens = runner_input.hidden_states.shape[0]
|
tokens = runner_input.hidden_states.shape[0]
|
||||||
hidden = runner_input.hidden_states.shape[-1]
|
hidden = runner_input.hidden_states.shape[-1]
|
||||||
total_rows = output.shape[0]
|
top_k = runner_input.n_expts_act
|
||||||
top_k = total_rows // tokens
|
|
||||||
|
if self.config.no_combine:
|
||||||
output = output.view(tokens, top_k, hidden)
|
output = output.view(tokens, top_k, hidden)
|
||||||
|
else:
|
||||||
|
output = output.view(tokens, top_k, hidden).sum(dim=1)
|
||||||
|
|
||||||
return TritonKernelsRunnerOutput(hidden_states=output)
|
return TritonKernelsRunnerOutput(hidden_states=output)
|
||||||
|
|
||||||
@@ -172,13 +174,15 @@ def pre_permute_standard_to_triton_kernels(
|
|||||||
topk_output
|
topk_output
|
||||||
), "Triton-kernel runner expects TritonKernelTopKOutput"
|
), "Triton-kernel runner expects TritonKernelTopKOutput"
|
||||||
|
|
||||||
routing_data, gather_indx, scatter_indx = topk_output
|
a_ragged_metadata, gather_indx, scatter_indx, gate_scal, n_expts_act = topk_output
|
||||||
|
|
||||||
return TritonKernelsRunnerInput(
|
return TritonKernelsRunnerInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
routing_data=routing_data,
|
a_ragged_metadata=a_ragged_metadata,
|
||||||
gather_indx=gather_indx,
|
gather_indx=gather_indx,
|
||||||
scatter_indx=scatter_indx,
|
scatter_indx=scatter_indx,
|
||||||
|
gate_scal=gate_scal,
|
||||||
|
n_expts_act=n_expts_act,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -32,10 +32,12 @@ from typing import (
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from triton_kernels.tensor_details.ragged_tensor import RaggedTensorMetadata
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_exec, get_lora, get_parallel
|
from sglang.srt.runtime_context import get_exec, get_lora, get_parallel
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
|
|
||||||
from triton_kernels.tensor import make_ragged_tensor_metadata
|
from triton_kernels.tensor import make_ragged_tensor_metadata
|
||||||
from triton_kernels.topk import topk as triton_kernels_topk
|
from triton_kernels.topk import topk as triton_kernels_topk
|
||||||
|
|
||||||
@@ -49,7 +51,7 @@ try:
|
|||||||
):
|
):
|
||||||
if simulated_ep != 1:
|
if simulated_ep != 1:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"simulated_ep routing is not supported with triton_kernels 3.6.0"
|
"simulated_ep routing is not supported with triton_kernels 3.7.1"
|
||||||
)
|
)
|
||||||
|
|
||||||
if sm_first:
|
if sm_first:
|
||||||
@@ -64,20 +66,13 @@ try:
|
|||||||
)
|
)
|
||||||
dispatch_indx = sparse_logits.mask_metadata.row_sorted_indx
|
dispatch_indx = sparse_logits.mask_metadata.row_sorted_indx
|
||||||
combine_indx = sparse_logits.mask_metadata.col_sorted_indx
|
combine_indx = sparse_logits.mask_metadata.col_sorted_indx
|
||||||
|
gather_indx = torch.div(combine_indx, n_expts_act, rounding_mode="trunc")
|
||||||
|
scatter_indx = combine_indx
|
||||||
ragged_metadata = make_ragged_tensor_metadata(
|
ragged_metadata = make_ragged_tensor_metadata(
|
||||||
sparse_logits.mask_metadata.col_sum, dispatch_indx.shape[0]
|
sparse_logits.mask_metadata.col_sum, dispatch_indx.shape[0]
|
||||||
)
|
)
|
||||||
gate_scal = sparse_logits.vals.flatten()[combine_indx]
|
gate_scal = sparse_logits.vals.flatten()[combine_indx]
|
||||||
routing_data = RoutingData(
|
return ragged_metadata, gather_indx, scatter_indx, gate_scal, n_expts_act
|
||||||
gate_scal,
|
|
||||||
ragged_metadata.slice_sizes,
|
|
||||||
logits.shape[-1],
|
|
||||||
n_expts_act,
|
|
||||||
ragged_metadata,
|
|
||||||
)
|
|
||||||
gather_indx = GatherIndx(combine_indx, dispatch_indx)
|
|
||||||
scatter_indx = ScatterIndx(dispatch_indx, combine_indx)
|
|
||||||
return routing_data, gather_indx, scatter_indx
|
|
||||||
|
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
@@ -319,9 +314,11 @@ class StandardTopKOutputPacked(NamedTuple):
|
|||||||
class TritonKernelTopKOutput(NamedTuple):
|
class TritonKernelTopKOutput(NamedTuple):
|
||||||
"""Triton kernel top-k output format."""
|
"""Triton kernel top-k output format."""
|
||||||
|
|
||||||
routing_data: RoutingData
|
a_ragged_metadata: RaggedTensorMetadata
|
||||||
gather_indx: GatherIndx
|
gather_indx: torch.Tensor
|
||||||
scatter_indx: ScatterIndx
|
scatter_indx: torch.Tensor
|
||||||
|
gate_scal: torch.Tensor
|
||||||
|
n_expts_act: int
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def format(self) -> TopKOutputFormat:
|
def format(self) -> TopKOutputFormat:
|
||||||
@@ -537,12 +534,24 @@ class TopK(BaseFusedOp):
|
|||||||
|
|
||||||
if output_format == TopKOutputFormat.TRITON_KERNEL:
|
if output_format == TopKOutputFormat.TRITON_KERNEL:
|
||||||
# renormalize=True is equivalent to sm_first=False
|
# renormalize=True is equivalent to sm_first=False
|
||||||
routing_data, gather_idx, scatter_idx = routing(
|
(
|
||||||
|
a_ragged_metadata,
|
||||||
|
gather_idx,
|
||||||
|
scatter_idx,
|
||||||
|
gate_scal,
|
||||||
|
n_expts_act,
|
||||||
|
) = routing(
|
||||||
router_logits,
|
router_logits,
|
||||||
self.topk_config.top_k,
|
self.topk_config.top_k,
|
||||||
sm_first=not self.topk_config.renormalize,
|
sm_first=not self.topk_config.renormalize,
|
||||||
)
|
)
|
||||||
return TritonKernelTopKOutput(routing_data, gather_idx, scatter_idx)
|
return TritonKernelTopKOutput(
|
||||||
|
a_ragged_metadata,
|
||||||
|
gather_idx,
|
||||||
|
scatter_idx,
|
||||||
|
gate_scal,
|
||||||
|
n_expts_act,
|
||||||
|
)
|
||||||
elif output_format == TopKOutputFormat.BYPASSED:
|
elif output_format == TopKOutputFormat.BYPASSED:
|
||||||
return BypassedTopKOutput(
|
return BypassedTopKOutput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
|
|||||||
@@ -263,7 +263,7 @@ class Fp8Config(QuantizationConfig):
|
|||||||
if weight_block_size is not None:
|
if weight_block_size is not None:
|
||||||
if not is_checkpoint_fp8_serialized:
|
if not is_checkpoint_fp8_serialized:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The block-wise quantization only supports fp8-serialized checkpoint for now."
|
"The block-wise quantization only supports fp8-serialized checkpoint for now."
|
||||||
)
|
)
|
||||||
if len(weight_block_size) != 2:
|
if len(weight_block_size) != 2:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -1788,11 +1788,10 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
from triton_kernels.tensor import convert_layout, wrap_torch_tensor
|
from triton_kernels.tensor import convert_layout, wrap_torch_tensor
|
||||||
from triton_kernels.tensor_details import layout
|
from triton_kernels.tensor_details import layout
|
||||||
|
|
||||||
scale_layout, scale_layout_opts = (
|
scale_layout = layout.make_default_matmul_mxfp4_w_scale_layout(
|
||||||
layout.make_default_matmul_mxfp4_w_scale_layout(
|
mx_axis=-2, num_warps=num_warps
|
||||||
mx_axis=1, num_warps=num_warps
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
scale_layout_opts = {}
|
||||||
scale = scale.transpose(-2, -1)
|
scale = scale.transpose(-2, -1)
|
||||||
scale = convert_layout(
|
scale = convert_layout(
|
||||||
wrap_torch_tensor(scale), scale_layout, **scale_layout_opts
|
wrap_torch_tensor(scale), scale_layout, **scale_layout_opts
|
||||||
|
|||||||
@@ -169,17 +169,17 @@ if _is_hip:
|
|||||||
|
|
||||||
def _swizzle_mxfp4(quant_tensor, scale, num_warps):
|
def _swizzle_mxfp4(quant_tensor, scale, num_warps):
|
||||||
"""weight swizzle for mxfp4 moe, used for OAI mxfp4 kernel"""
|
"""weight swizzle for mxfp4 moe, used for OAI mxfp4 kernel"""
|
||||||
import triton_kernels.matmul_ogs_details.opt_flags as opt_flags
|
import triton_kernels.matmul_details.opt_flags as opt_flags
|
||||||
from triton_kernels.numerics import InFlexData
|
from triton_kernels.numerics import InFlexData
|
||||||
from triton_kernels.tensor import FP4, convert_layout, wrap_torch_tensor
|
from triton_kernels.tensor import FP4, convert_layout, wrap_torch_tensor
|
||||||
from triton_kernels.tensor_details import layout
|
from triton_kernels.tensor_details import layout
|
||||||
|
|
||||||
value_layout, value_layout_opts = layout.make_default_matmul_mxfp4_w_layout(
|
value_layout = layout.make_default_matmul_mxfp4_w_layout(mx_axis=-2)
|
||||||
mx_axis=1
|
value_layout_opts = {}
|
||||||
)
|
scale_layout = layout.make_default_matmul_mxfp4_w_scale_layout(
|
||||||
scale_layout, scale_layout_opts = layout.make_default_matmul_mxfp4_w_scale_layout(
|
mx_axis=-2, num_warps=num_warps
|
||||||
mx_axis=1, num_warps=num_warps
|
|
||||||
)
|
)
|
||||||
|
scale_layout_opts = {}
|
||||||
if is_sm100_supported():
|
if is_sm100_supported():
|
||||||
constraints = {
|
constraints = {
|
||||||
"is_persistent": True,
|
"is_persistent": True,
|
||||||
@@ -931,7 +931,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
|
|
||||||
if self.use_triton_kernels:
|
if self.use_triton_kernels:
|
||||||
|
|
||||||
from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig
|
from triton_kernels.matmul import FlexCtx, PrecisionConfig
|
||||||
|
|
||||||
w13_weight_bias = layer.w13_weight_bias.to(torch.float32)
|
w13_weight_bias = layer.w13_weight_bias.to(torch.float32)
|
||||||
w2_weight_bias = layer.w2_weight_bias.to(torch.float32)
|
w2_weight_bias = layer.w2_weight_bias.to(torch.float32)
|
||||||
@@ -949,10 +949,10 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.w13_precision_config = PrecisionConfig(
|
self.w13_precision_config = PrecisionConfig(
|
||||||
weight_scale=w13_scale, flex_ctx=FlexCtx(rhs_data=w13_flex)
|
b_mx_scale=w13_scale, flex_ctx=FlexCtx(rhs_data=w13_flex)
|
||||||
)
|
)
|
||||||
self.w2_precision_config = PrecisionConfig(
|
self.w2_precision_config = PrecisionConfig(
|
||||||
weight_scale=w2_scale, flex_ctx=FlexCtx(rhs_data=w2_flex)
|
b_mx_scale=w2_scale, flex_ctx=FlexCtx(rhs_data=w2_flex)
|
||||||
)
|
)
|
||||||
|
|
||||||
self.w13_weight_triton_tensor = w13_weight
|
self.w13_weight_triton_tensor = w13_weight
|
||||||
|
|||||||
@@ -191,7 +191,6 @@ from sglang.srt.models.deepseek_common.utils import (
|
|||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_device,
|
get_device,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_flags,
|
|
||||||
get_forward,
|
get_forward,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -910,12 +909,6 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
and self.num_fused_shared_experts == 0
|
and self.num_fused_shared_experts == 0
|
||||||
and hidden_states.shape[0] > 0
|
and hidden_states.shape[0] > 0
|
||||||
and get_is_capture_mode()
|
and get_is_capture_mode()
|
||||||
and not (
|
|
||||||
get_flags().capture.enable_torch_compile
|
|
||||||
and hidden_states.shape[0]
|
|
||||||
<= get_exec().graph.torch_compile_max_bs
|
|
||||||
* (get_spec().speculative_num_draft_tokens or 1)
|
|
||||||
)
|
|
||||||
):
|
):
|
||||||
return self.forward_normal_dual_stream(
|
return self.forward_normal_dual_stream(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
|
|||||||
@@ -279,16 +279,10 @@ class NemotronHMoE(nn.Module):
|
|||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
||||||
overlap = _is_cuda and not torch.compiler.is_compiling()
|
if _is_cuda and (
|
||||||
if (
|
not get_moe_a2a_backend().is_flashinfer() or get_is_capture_mode()
|
||||||
overlap
|
|
||||||
and get_moe_a2a_backend().is_flashinfer()
|
|
||||||
and not get_is_capture_mode()
|
|
||||||
):
|
):
|
||||||
overlap = False
|
|
||||||
if overlap:
|
|
||||||
return self._forward_core_shared_routed_overlap(hidden_states)
|
return self._forward_core_shared_routed_overlap(hidden_states)
|
||||||
else:
|
|
||||||
return self._forward_core_normal(hidden_states)
|
return self._forward_core_normal(hidden_states)
|
||||||
|
|
||||||
def _forward_core_normal(
|
def _forward_core_normal(
|
||||||
|
|||||||
@@ -576,7 +576,11 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
shared_output = None
|
shared_output = None
|
||||||
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
||||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||||
elif self.alt_stream is not None and get_is_capture_mode():
|
elif (
|
||||||
|
self.alt_stream is not None
|
||||||
|
and get_is_capture_mode()
|
||||||
|
and not torch.compiler.is_compiling()
|
||||||
|
):
|
||||||
final_hidden_states, shared_output = self.forward_normal_dual_stream(
|
final_hidden_states, shared_output = self.forward_normal_dual_stream(
|
||||||
hidden_states, use_fused_gate=use_fused_gate
|
hidden_states, use_fused_gate=use_fused_gate
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4259,15 +4259,7 @@ class ConcurrentCounter:
|
|||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
def is_triton_kernels_available() -> bool:
|
def is_triton_kernels_available() -> bool:
|
||||||
if importlib.util.find_spec("triton_kernels") is None:
|
return importlib.util.find_spec("triton_kernels") is not None
|
||||||
return False
|
|
||||||
try:
|
|
||||||
ragged_metadata_spec = importlib.util.find_spec(
|
|
||||||
"triton_kernels.tensor_details.ragged_tensor"
|
|
||||||
)
|
|
||||||
except ModuleNotFoundError:
|
|
||||||
return False
|
|
||||||
return ragged_metadata_spec is not None
|
|
||||||
|
|
||||||
|
|
||||||
def json_list_type(value):
|
def json_list_type(value):
|
||||||
|
|||||||
@@ -352,6 +352,22 @@ uninstall_stale_flashinfer() {
|
|||||||
mark_step_done "${FUNCNAME[0]}"
|
mark_step_done "${FUNCNAME[0]}"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
install_pytorch_stack() {
|
||||||
|
PYTORCH_SPECS=()
|
||||||
|
for package in torch torchaudio torchvision torchao torchcodec; do
|
||||||
|
spec=$(grep -Po -m1 "\"${package}([<>=!~ ;][^\"]*)?\"" python/pyproject.toml | tr -d '"' || true)
|
||||||
|
if [ -n "$spec" ]; then
|
||||||
|
PYTORCH_SPECS+=("$spec")
|
||||||
|
fi
|
||||||
|
done
|
||||||
|
|
||||||
|
$PIP_CMD install \
|
||||||
|
"${PYTORCH_SPECS[@]}" \
|
||||||
|
--index-url "https://download.pytorch.org/whl/${CU_VERSION}"
|
||||||
|
|
||||||
|
mark_step_done "${FUNCNAME[0]}"
|
||||||
|
}
|
||||||
|
|
||||||
require_prebuilt_rust_exts() {
|
require_prebuilt_rust_exts() {
|
||||||
# Stages whose download succeeded set this to none. Runs before
|
# Stages whose download succeeded set this to none. Runs before
|
||||||
# setup_pip_toolchain uninstalls sglang, so clearing it here still reaches
|
# setup_pip_toolchain uninstalls sglang, so clearing it here still reaches
|
||||||
@@ -470,43 +486,6 @@ install_sglang_kernel() {
|
|||||||
fi
|
fi
|
||||||
fi
|
fi
|
||||||
|
|
||||||
# Reinstall torch with matching CUDA version if needed
|
|
||||||
# TODO: Remove after torch 2.11 where cu13 is enabled by default
|
|
||||||
REINSTALL_TORCH=false
|
|
||||||
if TORCH_CUDA_VER=$(python3 -c "import torch; v=torch.version.cuda; parts=v.split('.'); print(f'cu{parts[0]}{parts[1]}')" 2>&1); then
|
|
||||||
echo "Detected torch CUDA version: ${TORCH_CUDA_VER}"
|
|
||||||
else
|
|
||||||
TORCH_IMPORT_ERROR="${TORCH_CUDA_VER}"
|
|
||||||
TORCH_CUDA_VER=""
|
|
||||||
echo "WARNING: importing torch failed while probing CUDA version; force-reinstalling torch packages."
|
|
||||||
printf '%s\n' "${TORCH_IMPORT_ERROR}"
|
|
||||||
REINSTALL_TORCH=true
|
|
||||||
fi
|
|
||||||
TORCHAUDIO_CUDA_VER=$(pip show torchaudio 2>/dev/null | grep "^Version:" | awk '{print $2}' | sed -n 's/.*+\(cu[0-9][0-9]*\)$/\1/p' || true)
|
|
||||||
TORCHVISION_CUDA_VER=$(pip show torchvision 2>/dev/null | grep "^Version:" | awk '{print $2}' | sed -n 's/.*+\(cu[0-9][0-9]*\)$/\1/p' || true)
|
|
||||||
if [ "${TORCH_CUDA_VER}" != "${CU_VERSION}" ]; then
|
|
||||||
REINSTALL_TORCH=true
|
|
||||||
else
|
|
||||||
for cuda_ver in "${TORCHAUDIO_CUDA_VER}" "${TORCHVISION_CUDA_VER}"; do
|
|
||||||
if [ -n "${cuda_ver}" ] && [ "${cuda_ver}" != "${CU_VERSION}" ]; then
|
|
||||||
REINSTALL_TORCH=true
|
|
||||||
break
|
|
||||||
fi
|
|
||||||
done
|
|
||||||
fi
|
|
||||||
if [ "${REINSTALL_TORCH}" = true ]; then
|
|
||||||
TORCH_VER=$(pip show torch 2>/dev/null | grep "^Version:" | awk '{print $2}' | sed 's/+.*//')
|
|
||||||
TORCHAUDIO_VER=$(pip show torchaudio 2>/dev/null | grep "^Version:" | awk '{print $2}' | sed 's/+.*//')
|
|
||||||
TORCHVISION_VER=$(pip show torchvision 2>/dev/null | grep "^Version:" | awk '{print $2}' | sed 's/+.*//')
|
|
||||||
if [ -z "${TORCH_VER}" ] || [ -z "${TORCHAUDIO_VER}" ] || [ -z "${TORCHVISION_VER}" ]; then
|
|
||||||
echo "ERROR: could not determine installed torch package versions before reinstall."
|
|
||||||
pip show torch torchaudio torchvision || true
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
echo "Reinstalling torch==${TORCH_VER} torchaudio==${TORCHAUDIO_VER} torchvision==${TORCHVISION_VER} from ${CU_VERSION} index to match torch..."
|
|
||||||
$PIP_CMD install "torch==${TORCH_VER}" "torchaudio==${TORCHAUDIO_VER}" "torchvision==${TORCHVISION_VER}" --index-url "https://download.pytorch.org/whl/${CU_VERSION}" --force-reinstall --no-deps $PIP_INSTALL_SUFFIX
|
|
||||||
fi
|
|
||||||
|
|
||||||
if [ "${CUSTOM_BUILD_SGL_KERNEL:-}" != "true" ]; then
|
if [ "${CUSTOM_BUILD_SGL_KERNEL:-}" != "true" ]; then
|
||||||
# The PyPI default wheel tracks one CUDA version (currently cu130); other
|
# The PyPI default wheel tracks one CUDA version (currently cu130); other
|
||||||
# runners (e.g. h20 / cu129) need the +${CU_VERSION}-tagged wheel from the
|
# runners (e.g. h20 / cu129) need the +${CU_VERSION}-tagged wheel from the
|
||||||
@@ -779,6 +758,7 @@ main() {
|
|||||||
setup_pip_toolchain
|
setup_pip_toolchain
|
||||||
remove_stale_cuda12_nvidia_wheels
|
remove_stale_cuda12_nvidia_wheels
|
||||||
uninstall_stale_flashinfer
|
uninstall_stale_flashinfer
|
||||||
|
install_pytorch_stack
|
||||||
install_sglang
|
install_sglang
|
||||||
# Diffusion B200 CI imports torch inside install_sglang_kernel after removing
|
# Diffusion B200 CI imports torch inside install_sglang_kernel after removing
|
||||||
# stale CUDA 12 NVIDIA wheels, so opt into one early LD_LIBRARY_PATH refresh.
|
# stale CUDA 12 NVIDIA wheels, so opt into one early LD_LIBRARY_PATH refresh.
|
||||||
|
|||||||
@@ -18,7 +18,12 @@ from sglang.test.test_utils import (
|
|||||||
try_cached_model,
|
try_cached_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=189, stage="base-c", runner_config="deepep-4-gpu-h100")
|
register_cuda_ci(
|
||||||
|
est_time=189,
|
||||||
|
stage="base-c",
|
||||||
|
runner_config="deepep-4-gpu-h100",
|
||||||
|
disabled="Temporarily disabled until the next Mooncake release includes the PyTorch 2.13 collective forwarding fix.",
|
||||||
|
)
|
||||||
|
|
||||||
ib_devices = get_rdma_devices_args()
|
ib_devices = get_rdma_devices_args()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user