[AMD] support gfx1250 on ROCM 10 (#36871)

Co-authored-by: HAI <hixiao@gmail.com>
Co-authored-by: Kao <akao@amd.com>
Co-authored-by: wunhuang <wunhuang@amd.com>
Co-authored-by: Thomas Wang <1am9trash@gmail.com>
Co-authored-by: Xinyi Song <86638975+RolaoDenthu@users.noreply.github.com>
Co-authored-by: Lin, Soga <soga.lin@amd.com>
Co-authored-by: kk <43161300+kkHuang-amd@users.noreply.github.com>
Co-authored-by: Bingxu Chen <bingxche@amd.com>
Co-authored-by: sogalin_codegen <39478626+sogalin@users.noreply.github.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
YC Yen-Ching Tseng
2026-08-31 01:19:11 -07:00
committed by GitHub
co-authored by HAI Kao wunhuang Thomas Wang Xinyi Song Lin, Soga kk Bingxu Chen sogalin_codegen Thomas Wang
parent 712a720c8a
commit 3865efc9f7
35 changed files with 2200 additions and 298 deletions
@@ -1,109 +0,0 @@
name: Release Docker Images Nightly ROCm7.15 (AMD)
on:
workflow_dispatch:
inputs:
job_select:
description: 'Select which release job to run'
required: false
type: choice
default: 'all'
options:
- 'all'
- publish
# schedule:
# - cron: '0 12 * * *'
concurrency:
# A PR number if a pull request and otherwise the commit hash. This cancels
# queued and in-progress runs for the same PR (presubmit) or commit
# (postsubmit). The workflow name is prepended to avoid conflicts between
# different workflows.
group: ${{ github.workflow }}-${{ github.event.number || github.sha }}
cancel-in-progress: false
jobs:
publish:
if: github.repository == 'sgl-project/sglang' && (github.event_name != 'workflow_dispatch' || inputs.job_select == 'all' || inputs.job_select == 'publish')
runs-on: linux-mi300-1gpu-sglang
environment: 'prod'
strategy:
fail-fast: false
matrix:
gpu_arch: ['gfx1250-rocm7_15']
build_type: ['all']
steps:
- name: Checkout repository
uses: actions/checkout@v4
with:
fetch-depth: 0 # Required for git describe to find tags
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: "Set Date"
run: |
echo "DATE=$(date +%Y%m%d)" >> $GITHUB_ENV
- name: Get version from latest tag
id: version
run: |
# Use the shared helper so stable/post releases sort above rc tags.
VERSION=$(python3 scripts/release/get_version_tag.py --tag-only | sed 's/^v//')
if [ -z "$VERSION" ]; then
echo "::error::Could not determine version from git tags"
exit 1
fi
# Get short commit hash of current HEAD
COMMIT_HASH=$(git rev-parse --short HEAD)
# Compose pretend version for setuptools_scm: e.g., 0.5.8.post1.dev20260211+g1a2b3c4
PRETEND_VERSION="${VERSION}.dev${{ env.DATE }}+g${COMMIT_HASH}"
echo "version=${VERSION}" >> $GITHUB_OUTPUT
echo "pretend_version=${PRETEND_VERSION}" >> $GITHUB_OUTPUT
echo "Detected version: ${VERSION}"
echo "Pretend version for pip: ${PRETEND_VERSION}"
# - name: Login to Docker Hub (AMD)
# uses: docker/login-action@v2
# with:
# username: ${{ secrets.DOCKERHUB_AMD_USERNAME }}
# password: ${{ secrets.DOCKERHUB_AMD_TOKEN }}
- name: Build and Push to rocm/sgl-dev
run: |
version=${{ steps.version.outputs.version }}
pretend_version=${{ steps.version.outputs.pretend_version }}
echo "Version: ${version}"
echo "Pretend version: ${pretend_version}"
if [ "${{ matrix.gpu_arch }}" = "gfx1250-rocm7_15" ]; then
rocm_tag="rocm7.15-mi45x"
else
echo "Unsupported gfx arch"
exit 1
fi
tag=v${version}-${rocm_tag}
echo "IMAGE_TAG=${tag}-${{ env.DATE }}" >> $GITHUB_ENV
# remove --build-arg NIC_BACKEND=ainic for auto detection nic support in mori
# UBUNTU_MIRROR forces apt over HTTPS to dodge port-80 reachability flakes
# to Canonical's archive.ubuntu.com mirror IPs from the build runner.
docker build . -f docker/rocm.Dockerfile --build-arg SGL_BRANCH=${{ github.sha }} --build-arg BUILD_TYPE=${{ matrix.build_type }} --build-arg GPU_ARCH=${{ matrix.gpu_arch }} --build-arg GPU_ARCH_LIST_ARG=gfx1250 --build-arg ENABLE_MORI=1 --build-arg SGLANG_VERSION=${pretend_version} --build-arg UBUNTU_MIRROR=https://archive.ubuntu.com -t rocm/sgl-dev:${tag}-${{ env.DATE }} --no-cache
# docker push rocm/sgl-dev:${tag}-${{ env.DATE }}
- name: Login to Docker Hub (lmsys)
uses: docker/login-action@v2
with:
username: ${{ secrets.DOCKERHUB_USERNAME }}
password: ${{ secrets.DOCKERHUB_TOKEN }}
- name: Push to lmsysorg/sglang-rocm
run: |
docker tag rocm/sgl-dev:${{ env.IMAGE_TAG }} lmsysorg/sglang-rocm:${{ env.IMAGE_TAG }}
docker push lmsysorg/sglang-rocm:${{ env.IMAGE_TAG }}
+157 -42
View File
@@ -1,12 +1,13 @@
# Usage (to build SGLang ROCm docker image):
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942 -t v0.5.10.post1-rocm700-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942-rocm720 -t v0.5.10.post1-rocm720-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942-rocm724 -t v0.5.10.post1-rocm724-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950 -t v0.5.10.post1-rocm700-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950-rocm720 -t v0.5.10.post1-rocm720-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950-rocm724 -t v0.5.10.post1-rocm724-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942-rocm1000 -t v0.5.10.post1-rocm1000-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950-rocm1000 -t v0.5.10.post1-rocm1000-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942 -t v0.5.17-rocm700-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942-rocm720 -t v0.5.17-rocm720-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942-rocm724 -t v0.5.17-rocm724-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950 -t v0.5.17-rocm700-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950-rocm720 -t v0.5.17-rocm720-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950-rocm724 -t v0.5.17-rocm724-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942-rocm1000 -t v0.5.17-rocm1000-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950-rocm1000 -t v0.5.17-rocm1000-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx1250-rocm1000 -t v0.5.17-rocm1000-mi45x -f rocm.Dockerfile .
#
# Flavor notes:
# GPU_ARCH=*-rocm724 is built on a Python 3.12 base and upgrades the stack to
@@ -15,22 +16,25 @@
# GPU_ARCH=*-rocm1000 is Python 3.12 + torch 2.11, and takes the
# whole ROCm stack from AMD's stable wheel channel rather than an apt
# ROCm base image; see the rocm1000-base stage for what that changes.
# GPU_ARCH=gfx1250-rocm1000 carries the gfx1250 bring-up onto that same GA
# wheel channel. The gfx1250 workarounds key off GPU_ARCH_LIST=gfx1250 rather
# than the flavor name.
# Usage (to build SGLang ROCm + Mori docker image):
# remove --build-arg NIC_BACKEND=ainic since new MoRI JIT will do NIC auto detection on target
# Keep the build-arg for user to select the desired nic support, current choice: [ainic, bxnt]
# if no set this arg, it will support nic auto detection. On a target with more than 1 type of
# RDMA NICs installed (rare), overwrite w. runtime env MORI_DEVICE_NIC = "bnxt"|"ionic"|"mlx5"
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942 --build-arg ENABLE_MORI=1 -t v0.5.10.post1-rocm700-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx942-rocm720 --build-arg ENABLE_MORI=1 -t v0.5.10.post1-rocm720-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950 --build-arg ENABLE_MORI=1 -t v0.5.10.post1-rocm700-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950-rocm720 --build-arg ENABLE_MORI=1 -t v0.5.10.post1-rocm720-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942 --build-arg ENABLE_MORI=1 -t v0.5.17-rocm700-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx942-rocm720 --build-arg ENABLE_MORI=1 -t v0.5.17-rocm720-mi30x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950 --build-arg ENABLE_MORI=1 -t v0.5.17-rocm700-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950-rocm720 --build-arg ENABLE_MORI=1 -t v0.5.17-rocm720-mi35x -f rocm.Dockerfile .
# Usage (to build SGLang ROCm + NIXL docker image, for prefill/decode disaggregation):
# Builds UCX (--with-rocm) and upstream ai-dynamo/nixl from source by default.
# Set ENABLE_NIXL=0 to skip NIXL.
# At runtime use --disaggregation-transfer-backend nixl (env is wired via /etc/bash.bashrc).
# docker build --build-arg SGL_BRANCH=v0.5.10.post1 --build-arg GPU_ARCH=gfx950-rocm720 -t v0.5.10.post1-rocm720-mi35x -f rocm.Dockerfile .
# docker build --build-arg SGL_BRANCH=v0.5.17 --build-arg GPU_ARCH=gfx950-rocm720 -t v0.5.17-rocm720-mi35x -f rocm.Dockerfile .
# Default base images
ARG BASE_IMAGE_942="rocm/sgl-dev:rocm7-vllm-20250904"
@@ -43,6 +47,7 @@ ARG BASE_IMAGE_950_ROCM724="rocm/pytorch:rocm7.2.4_ubuntu24.04_py3.12_pytorch_re
# than a published image; point these at one to build on a prebuilt base.
ARG BASE_IMAGE_942_ROCM1000="rocm1000-base"
ARG BASE_IMAGE_950_ROCM1000="rocm1000-base"
ARG BASE_IMAGE_1250_ROCM1000="rocm1000-base"
ARG BASE_IMAGE_ROCM1000="ubuntu:24.04"
# This is necessary for scope purpose
@@ -67,6 +72,7 @@ ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="c16d44b93a528b2a4bfd6d8d3409116d465872a9"
ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840"
# ===============================
# Base image 942 with rocm724 and args (Python 3.12 + torch 2.11)
@@ -111,6 +117,7 @@ ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1"
ENV AITER_COMMIT_DEFAULT="c16d44b93a528b2a4bfd6d8d3409116d465872a9"
ENV TRITON_COMMIT_DEFAULT="42270451990532c67e69d753fbd026f28fcc4840"
# ===============================
# Base image 950 with rocm724 and args (Python 3.12 + torch 2.11)
@@ -297,6 +304,29 @@ ENV AITER_COMMIT_DEFAULT="c16d44b93a528b2a4bfd6d8d3409116d465872a9"
ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt"
RUN mkdir -p /etc/sglang/constraints && : > /etc/sglang/constraints/torch-rocm.txt
# ===============================
# Base image 1250 with ROCm 10.0.0 GA and args (Python 3.12 + torch 2.11)
# The gfx1250 build paths are keyed on GPU_ARCH_LIST=gfx1250 rather than the
# flavor name, so they apply here unchanged.
FROM $BASE_IMAGE_1250_ROCM1000 AS gfx1250-rocm1000
ENV BUILD_VLLM="0"
# Unlike the gfx942/gfx950 images, this one replaces the SDK's Triton: the
# revision below is what the MI45x bring-up ran on, and it carries a fix the
# SDK build does not have yet.
ENV BUILD_TRITON="1"
ENV BUILD_LLVM="0"
ENV BUILD_AITER_ALL="1"
ENV BUILD_MOONCAKE="1"
# gfx1250 tracks a different AITER than the gfx942/gfx950 baseline: this commit
# plus the four reverts applied at clone time are what the gfx1250 kernels were
# brought up against.
ENV AITER_COMMIT_DEFAULT="a6d2b564fd671724a3720b8edf70e8d674e4d694"
# The upstream Triton the gfx1250 bring-up was validated against, carried over
# from the ROCm 7.14 flavor this image replaced. Built from source below.
ENV TRITON_COMMIT_DEFAULT="76940ad348795521b3dc9f6c79acd7309ff924e3"
ENV PIP_CONSTRAINT="/etc/sglang/constraints/torch-rocm.txt"
RUN mkdir -p /etc/sglang/constraints && : > /etc/sglang/constraints/torch-rocm.txt
# Local source stage: with BRANCH_TYPE=local the build context is copied here and
# used instead of git clone (mirrors docker/Dockerfile's local_src stage).
FROM scratch AS local_src
@@ -308,13 +338,16 @@ FROM ${GPU_ARCH}
# This is necessary for scope purpose, again
ARG GPU_ARCH=gfx950
RUN echo GPU_ARCH="${GPU_ARCH}" \
&& echo "GPU_ARCH_LIST=${GPU_ARCH%%-*}" >> /etc/environment \
&& echo "export GPU_ARCH_LIST=${GPU_ARCH%%-*}" >> /etc/bash.bashrc
# ARG is build-time only. Stamp the stage name (gfx950-rocm724, gfx942, ...)
# so CI can read which AITER_COMMIT_DEFAULT block to use instead of guessing
# from torch or HIP — 720 may also ship torch 2.11 later, and both 7.2 flavors
# report HIP 7.2*.
ENV GPU_ARCH=${GPU_ARCH}
ENV GPU_ARCH_LIST=${GPU_ARCH%-*}
ENV PYTORCH_ROCM_ARCH=gfx942;gfx950
ENV PYTORCH_ROCM_ARCH="gfx942;gfx950;gfx1250"
ARG SGL_REPO="https://github.com/sgl-project/sglang.git"
ARG SGL_DEFAULT="main"
@@ -324,6 +357,9 @@ ARG BRANCH_TYPE=remote
# Version override for setuptools_scm (used in nightly builds)
ARG SETUPTOOLS_SCM_PRETEND_VERSION=""
ARG TRITON_REPO="https://github.com/triton-lang/triton.git"
ENV TRITON_COMMIT="${TRITON_COMMIT:-${TRITON_COMMIT_DEFAULT}}"
# ROCm 7.2 Triton (BUILD_TRITON=1 stages only). Both wheels are the same
# upstream revision, triton-lang/triton@89002410. AITER only requires
# triton>=3.6.0 and treats the base image as the owner of the version, so the
@@ -519,9 +555,8 @@ RUN if [ "$BUILD_LLVM" = "1" ]; then \
# -----------------------
# AITER
# Unset setuptools_scm override so AITER gets its own version (AITER_COMMIT), not SGLang's
# (SETUPTOOLS_SCM_PRETEND_VERSION is set later for SGLang nightly builds and would otherwise
# leak into AITER's version when AITER uses setuptools_scm)
# Clear any inherited setuptools_scm override so AITER gets its own version
# from AITER_COMMIT rather than SGLang's nightly version.
ENV SETUPTOOLS_SCM_PRETEND_VERSION=
# Compile AITER against the base image's Triton; the Triton step at the end of
@@ -539,12 +574,21 @@ RUN git clone ${AITER_REPO} \
&& git checkout -f ${AITER_COMMIT} \
&& git cherry-pick --no-commit 8578af153f4fa1e007fede7e3c1e1b373f07af4c \
&& git submodule update --init --recursive \
&& pip install -r requirements.txt
&& pip install -r requirements.txt \
&& if [ "${GPU_ARCH_LIST}" = "gfx1250" ]; then \
git revert --no-edit --no-commit 8d581370b; \
git revert --no-edit --no-commit 79c11b677; \
git revert --no-edit --no-commit 1ecb760a5; \
git revert --no-edit --no-commit e708f6c15; \
fi
RUN cd aiter \
&& echo "[AITER] GPU_ARCH=${GPU_ARCH}" \
&& echo "[AITER] AITER_USE_SYSTEM_TRITON=${AITER_USE_SYSTEM_TRITON}" \
&& if [ "$BUILD_AITER_ALL" = "1" ] && [ "$BUILD_LLVM" = "1" ]; then \
&& if [ "${GPU_ARCH_LIST}" = "gfx1250" ]; then \
PATH=$PATH:$ROCM_HOME/llvm/bin ENABLE_CK=0 GPU_ARCHS="${GPU_ARCH_LIST}" python setup.py build_ext --inplace \
&& PATH=$PATH:$ROCM_HOME/llvm/bin ENABLE_CK=0 GPU_ARCHS="${GPU_ARCH_LIST}" pip install --no-build-isolation -e .; \
elif [ "$BUILD_AITER_ALL" = "1" ] && [ "$BUILD_LLVM" = "1" ]; then \
sh -c "HIP_CLANG_PATH=/sgl-workspace/llvm-project/build/bin/ PREBUILD_KERNELS=1 GPU_ARCHS=$GPU_ARCH_LIST python setup.py build_ext --inplace" \
&& sh -c "HIP_CLANG_PATH=/sgl-workspace/llvm-project/build/bin/ GPU_ARCHS=$GPU_ARCH_LIST pip install --config-settings editable_mode=compat -e ."; \
elif [ "$BUILD_AITER_ALL" = "1" ]; then \
@@ -645,21 +689,40 @@ RUN if [ "$BRANCH_TYPE" = "local" ]; then \
&& cd python/sglang/kernels/aot \
&& rm -f pyproject.toml \
&& mv pyproject_rocm.toml pyproject.toml \
&& AMDGPU_TARGET=$GPU_ARCH_LIST python setup_rocm.py install \
&& cd ../../../.. \
&& rm -rf python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml \
# srt_hip pins compressed-tensors==0.15.0, which requires torch<2.11. Use
# the torch-2.11 ROCm 7.2.4 extra for both upgraded stacks; all other
# flavors keep the extras they used before.
&& AMDGPU_TARGET=$GPU_ARCH_LIST python setup_rocm.py install
RUN pip list --format=freeze | grep -E '^(torch|triton)' > /tmp/constraints.txt
# srt_hip pins compressed-tensors==0.15.0, which requires torch<2.11 and so
# cannot be satisfied on the ROCm 7.2.4/1000 torch 2.11 stack. The *_rocm724
# extras carry a 0.16.0 pin instead; all other flavors keep the extras they used before.
RUN cd sglang \
&& cp python/pyproject_other.toml python/pyproject.toml \
&& case "${GPU_ARCH}" in \
*-rocm724) srt_extras="srt_hip_rocm724,diffusion_hip"; all_extras="all_hip_rocm724" ;; \
*-rocm1000) srt_extras="srt_hip_rocm724,diffusion_hip"; all_extras="all_hip_rocm724" ;; \
*) srt_extras="srt_hip,diffusion_hip"; all_extras="all_hip" ;; \
*-rocm1000) srt_extras="srt_hip_rocm724,diffusion_hip"; \
all_extras="all_hip_rocm724" ; \
CONS="-c /tmp/constraints.txt" ; \
;; \
*-rocm724) srt_extras="srt_hip_rocm724,diffusion_hip"; \
all_extras="all_hip_rocm724" ; \
CONS="-c /tmp/constraints.txt" ; \
;; \
*-rocm720) srt_extras="srt_hip,diffusion_hip"; \
all_extras="all_hip" ; \
CONS="-c /tmp/constraints.txt" ; \
;; \
*) srt_extras="srt_hip,diffusion_hip"; \
all_extras="all_hip" ; \
CONS="-c /tmp/constraints.txt" ; \
echo 'diffusers==0.37.0' >> /tmp/constraints.txt ; \
echo 'transformers==5.12.1' >> /tmp/constraints.txt ; \
echo 'tokenizers==0.22.2' >> /tmp/constraints.txt ; \
echo 'huggingface_hub==1.27.0' >> /tmp/constraints.txt ; \
;; \
esac \
&& if [ "$BUILD_TYPE" = "srt" ]; then \
export SETUPTOOLS_SCM_PRETEND_VERSION="${SETUPTOOLS_SCM_PRETEND_VERSION}" && python -m pip --no-cache-dir install -e "python[${srt_extras}]"; \
export SETUPTOOLS_SCM_PRETEND_VERSION="${SETUPTOOLS_SCM_PRETEND_VERSION}" && python -m pip --no-cache-dir install $CONS -e "python[${srt_extras}]"; \
else \
export SETUPTOOLS_SCM_PRETEND_VERSION="${SETUPTOOLS_SCM_PRETEND_VERSION}" && python -m pip --no-cache-dir install -e "python[${all_extras}]"; \
export SETUPTOOLS_SCM_PRETEND_VERSION="${SETUPTOOLS_SCM_PRETEND_VERSION}" && python -m pip --no-cache-dir install $CONS -e "python[${all_extras}]"; \
fi
RUN python -m pip cache purge
@@ -764,7 +827,12 @@ RUN /bin/bash -lc 'set -euo pipefail; \
git fetch --depth=1 origin "${TILELANG_COMMIT}" || true && \
git checkout -f "${TILELANG_COMMIT}" && \
git submodule update --init --recursive && \
export CMAKE_ARGS="-DUSE_CUDA=OFF -DUSE_ROCM=ON -DROCM_PATH=/opt/rocm -DLLVM_CONFIG=${LLVM_CONFIG} -DSKBUILD_SABI_VERSION= ${CMAKE_ARGS:-}" && \
if [ "${GPU_ARCH_LIST}" = "gfx1250" ]; then \
export ROCM_PATH=${ROCM_HOME}; \
else \
export ROCM_PATH=/opt/rocm; \
fi; \
export CMAKE_ARGS="-DUSE_CUDA=OFF -DUSE_ROCM=ON -DROCM_PATH=${ROCM_PATH} -DLLVM_CONFIG=${LLVM_CONFIG} -DSKBUILD_SABI_VERSION= ${CMAKE_ARGS:-}" && \
"$VENV_PIP" install -e . -v --no-build-isolation --no-deps; \
if [ -f pyproject.toml ]; then sed -i "/^[[:space:]]*\"torch/d" pyproject.toml || true; fi; \
"$VENV_PIP" cache purge || true; \
@@ -776,7 +844,7 @@ RUN /bin/bash -lc 'set -euo pipefail; \
git clone --branch "${FHT_BRANCH}" "${FHT_REPO}" fast-hadamard-transform; \
cd fast-hadamard-transform; \
git checkout -f "${FHT_COMMIT}"; \
python setup.py install'
PYTORCH_ROCM_ARCH=${GPU_ARCH_LIST} python setup.py install'
# -----------------------
# Python tools
@@ -804,6 +872,19 @@ RUN /bin/bash -lc 'set -euo pipefail; \
initramfs-tools \
&& rm -rf /var/lib/apt/lists/*; \
\
# gfx1250 brought its own MORI build up against a wider set of RDMA, MPI and
# gRPC packages. Scoped to that arch so the CDNA images keep the package set
# they were released with.
if [ "${GPU_ARCH_LIST}" = "gfx1250" ]; then \
apt-get update && apt-get install -y --no-install-recommends \
libdrm-dev \
librdmacm-dev rdmacm-utils infiniband-diags ibverbs-utils perftest ethtool \
libibverbs-dev rdma-core \
openssh-server openmpi-bin openmpi-common libopenmpi-dev \
libgrpc++-dev protobuf-compiler-grpc \
&& rm -rf /var/lib/apt/lists/*; \
fi; \
\
# NIC backend deps — mori auto-detects NIC at runtime (MORI_DEVICE_NIC env var override).
# Only vendor packages are installed here for dlopen (e.g. libionic.so); no compile-time flags needed.
case "${NIC_BACKEND}" in \
@@ -875,7 +956,18 @@ RUN /bin/bash -lc 'set -euo pipefail; \
# so it shadows nothing system-wide. Scope this explicitly to ROCm 10 so the
# ROCm 7.2 and 7.2.4 MORI build paths remain byte-for-byte equivalent here.
ROCM_SYSDEPS="${ROCM_HOME:-/opt/rocm}/lib/rocm_sysdeps"; \
if [ "${GPU_ARCH##*-}" = "rocm1000" ] && [ -d "${ROCM_SYSDEPS}" ]; then \
if [ "${GPU_ARCH_LIST}" = "gfx1250" ] && [ -d "${ROCM_SYSDEPS}" ]; then \
# gfx1250 was brought up needing the SDK's own cmake trees on the prefix
# path as well: that is what lets hsakmt-config.cmake resolve
# find_dependency(NUMA) without patching MORI's CMakeLists.
export PATH="${ROCM_HOME}/bin:${PATH}"; \
export CMAKE_PREFIX_PATH="${ROCM_SYSDEPS}/lib/cmake:${ROCM_SYSDEPS}:${ROCM_HOME}/lib/cmake:${ROCM_HOME}${CMAKE_PREFIX_PATH:+:${CMAKE_PREFIX_PATH}}"; \
export CPATH="${ROCM_SYSDEPS}/include${CPATH:+:${CPATH}}"; \
export LIBRARY_PATH="${ROCM_SYSDEPS}/lib${LIBRARY_PATH:+:${LIBRARY_PATH}}"; \
echo "${ROCM_SYSDEPS}/lib" > /etc/ld.so.conf.d/rocm-sysdeps.conf; \
ldconfig; \
echo "[MORI] rocm_sysdeps prefix: ${ROCM_SYSDEPS}"; \
elif [ "${GPU_ARCH##*-}" = "rocm1000" ] && [ -d "${ROCM_SYSDEPS}" ]; then \
export CMAKE_PREFIX_PATH="${ROCM_SYSDEPS}${CMAKE_PREFIX_PATH:+:${CMAKE_PREFIX_PATH}}"; \
export CPATH="${ROCM_SYSDEPS}/include${CPATH:+:${CPATH}}"; \
export LIBRARY_PATH="${ROCM_SYSDEPS}/lib${LIBRARY_PATH:+:${LIBRARY_PATH}}"; \
@@ -903,13 +995,20 @@ RUN /bin/bash -lc 'set -euo pipefail; \
build-essential autoconf automake libtool pkg-config git \
libibverbs-dev librdmacm-dev rdma-core && rm -rf /var/lib/apt/lists/*; \
# Mooncake's dependencies.sh apt-installs Ubuntu's libabsl-dev (20220623 on
# the noble base used by rocm724 and rocm1000). NIXL's meson then finds
# absl_base but no absl_log and refuses to fall back to its bundled Abseil --
# "that would result in a mix of Abseil versions at runtime" -- so nixl fails
# at metadata generation. Drop just the -dev package (headers and pkg-config
# files); the runtime library that already-built components link against
# stays in place.
# the noble base used by rocm724 and rocm1000). NIXL's meson then
# finds absl_base but no absl_log and refuses to fall back to its bundled
# Abseil -- "that would result in a mix of Abseil versions at runtime" -- so
# nixl fails at metadata generation. Drop just the -dev package (headers and
# pkg-config files); the runtime library that already-built components link
# against stays in place.
#
# gfx1250 was brought up dropping the runtime package as well, so keep that
# for gfx1250 only rather than changing what the CDNA images were released with.
if [ "${GPU_ARCH_LIST}" = "gfx1250" ]; then \
apt-get remove -y libabsl-dev libabsl20220623 || true; \
else \
case "${GPU_ARCH}" in *-rocm724|*-rocm1000) apt-get remove -y libabsl-dev ;; esac; \
fi; \
pip install --no-cache-dir meson ninja pybind11 meson-python patchelf pyyaml; \
git clone --depth=1 -b "${UCX_BRANCH}" "${UCX_REPO}" /sgl-workspace/ucx; \
cd /sgl-workspace/ucx && ./autogen.sh && mkdir build && cd build && \
@@ -1013,9 +1112,25 @@ RUN python3 -c "from pathlib import Path; import transformers.dynamic_module_uti
# install pull CUDA torch instead. The hip check below is the tripwire.
# torch 2.11 names this `triton-rocm`; uninstall it so the pin is the only Triton.
RUN if [ "$BUILD_TRITON" = "1" ]; then \
pip uninstall -y triton-rocm || true \
&& PIP_NO_CACHE_DIR=1 pip install --extra-index-url ${TRITON_INDEX_URL} \
"triton==${TRITON_VERSION}" "triton-kernels==${TRITON_KERNELS_VERSION}" \
case "${GPU_ARCH}" in \
gfx1250-rocm1000) \
echo "[Triton] gfx1250: building ${TRITON_COMMIT} from source"; \
pip uninstall -y triton triton-rocm || true \
&& apt-get update && apt-get install -y --no-install-recommends cmake && rm -rf /var/lib/apt/lists/* \
&& git clone ${TRITON_REPO} triton-custom \
&& cd triton-custom \
&& git checkout ${TRITON_COMMIT} \
&& pip install -r python/requirements.txt \
&& pip install -e . \
&& if [ -d python/triton_kernels ]; then pip install -e python/triton_kernels --no-deps; fi; \
;; \
*rocm72*) \
echo "[Triton] ROCm 7.2: installing pinned wheels from ${TRITON_INDEX_URL}"; \
pip uninstall -y triton-rocm || true && \
PIP_NO_CACHE_DIR=1 pip install --extra-index-url ${TRITON_INDEX_URL} \
"triton==${TRITON_VERSION}" "triton-kernels==${TRITON_KERNELS_VERSION}"; \
;; \
esac \
&& python3 -c "import torch; from importlib.metadata import version; v = version('triton'); k = version('triton-kernels'); assert torch.version.hip is not None, torch.__version__; print(f'[Triton] ROCm Torch {torch.__version__}, Triton {v}, triton-kernels {k}')"; \
fi
+1 -5
View File
@@ -112,13 +112,9 @@ tracing = [
# in its compressed-tensors and torch pins. A dependency added here has to be
# added there too; nothing enforces that.
srt_hip = [
# Pin to 0.15.0: 0.16.0 needs torch>=2.10 (incompatible with ROCm torch
# 2.9.1). An open-ended `<0.16.0` made pip backtrack into an unbuildable
# ancient setuptools sdist; an exact pin keeps the resolver converging.
"compressed-tensors==0.15.0",
"petit_kernel==0.0.2",
"sglang[runtime_common]",
"torch",
"petit_kernel==0.0.2",
"wave-lang==3.8.2",
]
@@ -91,6 +91,8 @@ union BufferResource {
};
// llvm.amdgcn.raw.buffer.* instructions do not exist on RDNA4 (gfx12).
// Stubs satisfy the compiler; these functions must not be called on gfx1250.
// Mirrors vLLM PR #46516 csrc/quickreduce/base.h.
// QuickReduce remains runtime-disabled on gfx1250; these stubs only allow the
// shared ROCm extension to compile for that target.
#if !defined(__gfx1250__)
@@ -102,7 +104,10 @@ buffer_store_dwordx4(int32x4_t data, int32x4_t srsrc, int32_t voffset, int32_t s
"llvm.amdgcn.raw.buffer.store.v4i32");
#else
__quickreduce_device_inline__ static int32x4_t
buffer_load_dwordx4(int32x4_t srsrc, int32_t voffset, int32_t soffset, int32_t aux) {}
buffer_load_dwordx4(int32x4_t srsrc, int32_t voffset, int32_t soffset, int32_t aux) {
__builtin_trap();
return int32x4_t{};
}
__quickreduce_device_inline__ static void
buffer_store_dwordx4(int32x4_t data, int32x4_t srsrc, int32_t voffset, int32_t soffset, int32_t aux) {}
+1 -1
View File
@@ -82,7 +82,7 @@ if amdgpu_target not in ["gfx942", "gfx950", "gfx1250"]:
fp8_macro = (
"-DHIP_FP8_TYPE_FNUZ" if amdgpu_target == "gfx942" else "-DHIP_FP8_TYPE_E4M3"
)
) # gfx950 and gfx1250 use E4M3
# Dynamic shared-memory budget for the TopK kernels.
# - gfx942 (MI300/MI325): LDS is typically 64KB per workgroup -> keep dynamic smem <= ~48KB
@@ -10,6 +10,15 @@
#include <cfloat>
#include <cstdint>
// gfx1250 needs a 64-bit __shfl_*_sync mask (static_assert sizeof == 8); it's
// masked to wave32 internally, so 64-bit is fine everywhere. __gfx1250__ is
// device-pass only, so widen in the host pass too or the launch stub won't build.
#if defined(__gfx1250__) || (defined(__HIP_PLATFORM_AMD__) && !defined(__HIP_DEVICE_COMPILE__))
#define SGL_WARP_SYNC_MASK 0xFFFFFFFFFFFFFFFFULL
#else
#define SGL_WARP_SYNC_MASK 0xFFFFFFFF
#endif
namespace sglang {
constexpr uint32_t kWarpSize = 32;
@@ -104,8 +113,8 @@ __global__ void moe_fused_gate_kernel_small_token(const MoEFusedGateParams __gri
#pragma unroll
for (int offset = 16; offset > 0; offset /= 2) {
float other_val = __shfl_down_sync(0xFFFFFFFF, warp_max_val, offset);
int other_expert = __shfl_down_sync(0xFFFFFFFF, warp_max_expert, offset);
float other_val = __shfl_down_sync(SGL_WARP_SYNC_MASK, warp_max_val, offset);
int other_expert = __shfl_down_sync(SGL_WARP_SYNC_MASK, warp_max_expert, offset);
if (other_val > warp_max_val) {
warp_max_val = other_val;
warp_max_expert = other_expert;
@@ -125,8 +134,8 @@ __global__ void moe_fused_gate_kernel_small_token(const MoEFusedGateParams __gri
#pragma unroll
for (int offset = 16; offset > 0; offset /= 2) {
float other_val = __shfl_down_sync(0xFFFFFFFF, final_max, offset);
int other_expert = __shfl_down_sync(0xFFFFFFFF, final_expert, offset);
float other_val = __shfl_down_sync(SGL_WARP_SYNC_MASK, final_max, offset);
int other_expert = __shfl_down_sync(SGL_WARP_SYNC_MASK, final_expert, offset);
if (other_val > final_max) {
final_max = other_val;
final_expert = other_expert;
@@ -219,8 +228,8 @@ __global__ void moe_fused_gate_kernel(const MoEFusedGateParams __grid_constant__
}
for (int offset = kWarpSize / 2; offset > 0; offset /= 2) {
float other_val = __shfl_down_sync(0xFFFFFFFF, max_val, offset);
int other_expert = __shfl_down_sync(0xFFFFFFFF, max_expert, offset);
float other_val = __shfl_down_sync(SGL_WARP_SYNC_MASK, max_val, offset);
int other_expert = __shfl_down_sync(SGL_WARP_SYNC_MASK, max_expert, offset);
if (other_val > max_val || (other_val == max_val && other_expert < max_expert)) {
max_val = other_val;
@@ -30,9 +30,15 @@ import triton.language as tl
from sglang.kernels.ops.attention.score_mod import unpack_aux_tensors
from sglang.srt.environ import envs
from sglang.srt.utils import get_device_core_count, is_gfx95_supported, is_hip
from sglang.srt.utils import (
get_device_core_count,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
)
_is_hip = is_hip()
_is_gfx1250 = _is_hip and is_gfx1250_supported()
logger = logging.getLogger(__name__)
@@ -539,6 +545,7 @@ def _fwd_grouped_kernel_stage1(
Lv: tl.constexpr,
HAS_MLA: tl.constexpr = False,
USE_PDL: tl.constexpr = False,
IS_GFX1250: tl.constexpr = False,
PAGE_SIZE: tl.constexpr = 1,
SCORE_MOD: tl.constexpr = None,
Aux0=None,
@@ -613,6 +620,16 @@ def _fwd_grouped_kernel_stage1(
if split_kv_end > split_kv_start:
q = tl.load(Q + offs_q, mask=(mask_h[:, None]) & (mask_d[None, :]), other=0.0)
# gfx1250: triton tl.dot(fp8, fp8) returns garbage (~1e34+) for contraction
# dim K>=128 (verified K=64 ok, K>=128 broken; bf16 fine at all K). The MLA
# nope QK dot has K=512, so an fp8 KV cache MUST NOT be consumed as an fp8 dot
# here: keep q in bf16 and upcast the fp8 K to bf16 for the dot. No-op for a
# bf16 cache. (Do NOT "optimize" this back to q.to(fp8) on gfx1250.)
# On all other platforms keep the original downcast of q to the KV dtype.
# TODO: remove this branch once the gfx1250 fp8 tl.dot issue is resolved.
if IS_GFX1250:
q_k = q
else:
q_k = q.to(K_Buffer.dtype.element_ty)
if BLOCK_DPE > 0:
qpe = tl.load(
@@ -641,6 +658,9 @@ def _fwd_grouped_kernel_stage1(
mask=(offs_n[None, :] < split_kv_end) & (mask_d[:, None]),
other=0.0,
)
if IS_GFX1250:
qk = tl.dot(q_k, k.to(q_k.dtype))
else:
qk = tl.dot(q_k, k)
if BLOCK_DPE > 0:
if PAGE_SIZE == 1:
@@ -703,6 +723,14 @@ def _fwd_grouped_kernel_stage1(
re_scale = tl.exp(e_max - n_e_max)
p = tl.exp(qk - n_e_max[:, None])
acc *= re_scale[:, None]
# Keep the softmax weights p in fp32 for the P·V dot (do NOT downcast p to
# bf16) on gfx1250. The bf16 downcast of p was the accuracy loss vs a torch
# fp32 SDPA reference (recovers gfx1250 R1 GSM8K ~0.82 -> ~0.92 with
# attention idealized). On other platforms restore the p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
acc += tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
acc += tl.dot(p.to(v.dtype), v)
e_sum = e_sum * re_scale + tl.sum(p, 1)
@@ -859,6 +887,7 @@ def _decode_grouped_att_m_fwd(
Lv=Lv,
HAS_MLA=has_mla,
USE_PDL=use_pdl,
IS_GFX1250=_is_gfx1250,
PAGE_SIZE=page_size,
SCORE_MOD=score_mod,
Aux0=aux0,
@@ -59,6 +59,9 @@ import triton.language as tl
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
from sglang.srt.utils import is_hip
from sglang.srt.utils.common import is_gfx1250_supported
_is_gfx1250_supported = is_gfx1250_supported()
LOG2E = 1.4426950408889634 # log2(e); folded into qk_scale so softmax can use exp2.
@@ -903,6 +906,22 @@ def sparse_attn_v4_paged_decode(
When ``kv_scales`` is provided, ``unified_kv`` must be fp8 (e4m3fnuz) and
will be dequantized in-kernel using 1xGROUP_SIZE (default 64) block scales.
"""
if _is_gfx1250_supported:
# aiter ships only on ROCm, and this module is imported by a CPU-registered
# test, so the import has to sit behind the same gate as the call.
from aiter.ops.triton.attention.pa_decode_sparse import pa_decode_sparse
return pa_decode_sparse(
q,
unified_kv,
kv_indices,
kv_indptr,
attn_sink,
softmax_scale,
has_invalid=False,
kv_scales=kv_scales,
)
else:
return _sparse_attn_v4_paged_decode_triton(
q,
unified_kv,
@@ -49,14 +49,14 @@ import torch
import triton
import triton.language as tl
from sglang.srt.utils.common import is_gfx95_supported
from sglang.srt.utils.common import is_gfx95_supported, is_gfx1250_supported
# OPUS gfx950 paged-prefill kernel is preferred when importable; otherwise fall
# back to the Triton implementation below.
try:
from aiter.ops.pa_sparse_prefill_opus import pa_sparse_prefill_opus
_HAS_OPUS = is_gfx95_supported()
_HAS_OPUS = is_gfx95_supported() and not is_gfx1250_supported()
except ImportError:
pa_sparse_prefill_opus = None
_HAS_OPUS = False
@@ -26,7 +26,12 @@ from sglang.kernels.ops.attention.prefill_attention import (
)
from sglang.kernels.ops.attention.score_mod import unpack_aux_tensors
from sglang.srt.environ import envs
from sglang.srt.utils import is_cuda, is_gfx95_supported, is_hip
from sglang.srt.utils import (
is_cuda,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
)
_is_cuda = is_cuda()
if _is_cuda:
@@ -34,6 +39,7 @@ if _is_cuda:
_is_hip = is_hip()
_is_gfx95 = _is_hip and is_gfx95_supported()
_is_gfx1250 = _is_hip and is_gfx1250_supported()
try:
_triton_version_parts = tuple(
@@ -358,7 +364,8 @@ def _fwd_kernel(
SKIP_EXTEND: tl.constexpr,
STORE_TRANSPOSE: tl.constexpr,
HAS_SINK: tl.constexpr,
USE_COMPACT_TILE_GRID: tl.constexpr,
IS_GFX1250: tl.constexpr = False,
USE_COMPACT_TILE_GRID: tl.constexpr = False,
PAGE_SIZE: tl.constexpr = 1,
SCORE_MOD: tl.constexpr = None,
Aux0=None,
@@ -519,6 +526,16 @@ def _fwd_kernel(
mask=(mask_n[None, :]) & (mask_d[:, None]),
other=0.0,
)
# gfx1250: triton tl.dot(fp8, fp8) returns garbage (~1e34+) for contraction
# dim K>=128 (K=64 ok). This prefix read fires when a radix-cache prefix is
# reused (prefill reads the cached fp8 KV), and the MLA nope dot has K=512,
# so we must upcast the fp8 K to q's dtype and dot in bf16 rather than
# downcasting q to fp8. No-op for a bf16 cache. (Do NOT revert to q.to(fp8).)
# On all other platforms keep the original q.to(k.dtype) downcast.
# TODO: remove this branch once the gfx1250 fp8 tl.dot issue is resolved.
if IS_GFX1250:
qk = tl.dot(q, k.to(q.dtype))
else:
qk = tl.dot(q.to(k.dtype), k)
if BLOCK_DPE > 0:
if PAGE_SIZE == 1:
@@ -539,6 +556,9 @@ def _fwd_kernel(
mask=mask_n[None, :],
other=0.0,
)
if IS_GFX1250:
qk += tl.dot(qpe, kpe.to(qpe.dtype))
else:
qk += tl.dot(qpe.to(kpe.dtype), kpe)
qk *= sm_scale * k_scale
@@ -592,8 +612,14 @@ def _fwd_kernel(
mask=mask_n[:, None] & mask_dv[None, :],
other=0.0,
)
p = p.to(v.dtype)
acc = acc * re_scale[:, None] + tl.dot(p, v) * v_scale
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot * v_scale
e_max = n_e_max
@@ -722,8 +748,14 @@ def _fwd_kernel(
v = tl.load(
V_Extend + offs_v, mask=mask_n[:, None] & mask_dv[None, :], other=0.0
)
p = p.to(v.dtype)
acc = acc * re_scale[:, None] + tl.dot(p, v)
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot
e_max = n_e_max
@@ -917,6 +949,7 @@ def extend_attention_fwd(
SKIP_PREFIX=skip_prefix,
SKIP_EXTEND=skip_extend,
HAS_SINK=HAS_SINK,
IS_GFX1250=_is_gfx1250,
STORE_TRANSPOSE=_is_hip,
USE_COMPACT_TILE_GRID=use_compact_tile_grid,
PAGE_SIZE=page_size,
@@ -1011,6 +1044,7 @@ def _fwd_kernel_unified(
IS_CAUSAL: tl.constexpr,
USE_CUSTOM_MASK: tl.constexpr,
HAS_SINK: tl.constexpr,
IS_GFX1250: tl.constexpr = False,
PAGE_SIZE: tl.constexpr = 1,
SCORE_MOD: tl.constexpr = None,
Aux0=None,
@@ -1179,6 +1213,16 @@ def _fwd_kernel_unified(
other=0.0,
)
# gfx1250: triton tl.dot(fp8, fp8) returns garbage (~1e34+) for contraction
# dim K>=128 (K=64 ok). This prefix read fires when a radix-cache prefix is
# reused (prefill reads the cached fp8 KV), and the MLA nope dot has K=512,
# so we must upcast the fp8 K to q's dtype and dot in bf16 rather than
# downcasting q to fp8. No-op for a bf16 cache. (Do NOT revert to q.to(fp8).)
# On all other platforms keep the original q.to(k.dtype) downcast.
# TODO: remove this branch once the gfx1250 fp8 tl.dot issue is resolved.
if IS_GFX1250:
qk = tl.dot(q, k.to(q.dtype))
else:
qk = tl.dot(q.to(k.dtype), k)
if BLOCK_DPE > 0:
if PAGE_SIZE == 1:
@@ -1199,6 +1243,9 @@ def _fwd_kernel_unified(
mask=mask_n[None, :],
other=0.0,
)
if IS_GFX1250:
qk += tl.dot(qpe, kpe.to(qpe.dtype))
else:
qk += tl.dot(qpe.to(kpe.dtype), kpe)
qk *= sm_scale_withk
@@ -1253,8 +1300,14 @@ def _fwd_kernel_unified(
mask=mask_n[:, None] & mask_dv[None, :],
other=0.0,
)
p = p.to(v.dtype)
acc = acc * re_scale[:, None] + tl.dot(p, v)
# keep softmax weights p in fp32 for the P·V dot (do not downcast to bf16)
# on gfx1250; on other platforms restore the original p.to(v.dtype) cast.
# TODO: remove this branch once the gfx1250 bf16 P·V issue is resolved.
if IS_GFX1250:
dot = tl.dot(p, v.to(tl.float32), out_dtype=tl.float32)
else:
dot = tl.dot(p.to(v.dtype), v)
acc = acc * re_scale[:, None] + dot
e_max = n_e_max
@@ -1406,6 +1459,7 @@ def extend_attention_fwd_unified(
IS_CAUSAL=is_causal,
USE_CUSTOM_MASK=USE_CUSTOM_MASK,
HAS_SINK=HAS_SINK,
IS_GFX1250=_is_gfx1250,
PAGE_SIZE=page_size,
SCORE_MOD=score_mod,
Aux0=aux0,
+171
View File
@@ -6,6 +6,8 @@ import threading
from typing import Tuple
import torch
import triton
import triton.language as tl
from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
@@ -16,6 +18,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_round_robin_split
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.utils.common import strict_contiguous
from sglang.srt.utils.common import is_gfx1250_supported
logger = logging.getLogger(__name__)
@@ -189,6 +192,168 @@ def hc_split_sinkhorn_kernel(hc: int, sinkhorn_iters: int, eps: float):
return hc_split_sinkhorn_kernel_
def _hc_split_sinkhorn_torch(
mixes: torch.Tensor,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
hc_mult: int = 4,
sinkhorn_iters: int = 20,
eps: float = 1e-6,
):
"""Pure-torch equivalent of hc_split_sinkhorn_kernel.
TileLang's CK-backed buffer addressing does not compile on gfx1250, so the
sinkhorn kernel is reimplemented here. Layout mirrors the kernel exactly:
the flattened ``mixes`` row holds ``pre`` (hc), ``post`` (hc) and the
``comb`` matrix (hc * hc) consecutively.
"""
b, s, _ = mixes.size()
hc = hc_mult
flat = mixes.reshape(-1, (2 + hc) * hc).float()
scale = hc_scale.float()
base = hc_base.float()
pre = torch.sigmoid(flat[:, :hc] * scale[0] + base[:hc]) + eps
post = 2 * torch.sigmoid(flat[:, hc : 2 * hc] * scale[1] + base[hc : 2 * hc])
comb = flat[:, 2 * hc :] * scale[2] + base[2 * hc :]
comb = comb.reshape(-1, hc, hc)
# Initial row softmax (numerically stabilized) then column normalize.
row_max = comb.amax(dim=2, keepdim=True)
comb = torch.exp(comb - row_max)
comb = comb / comb.sum(dim=2, keepdim=True) + eps
comb = comb / (comb.sum(dim=1, keepdim=True) + eps)
for _ in range(sinkhorn_iters - 1):
comb = comb / (comb.sum(dim=2, keepdim=True) + eps)
comb = comb / (comb.sum(dim=1, keepdim=True) + eps)
pre = pre.reshape(b, s, hc).to(mixes.dtype)
post = post.reshape(b, s, hc).to(mixes.dtype)
comb = comb.reshape(b, s, hc, hc).to(mixes.dtype)
return pre, post, comb
@triton.jit
def _hc_split_sinkhorn_triton_kernel(
mixes_ptr,
hc_scale_ptr,
hc_base_ptr,
pre_ptr,
post_ptr,
comb_ptr,
n,
HC: tl.constexpr,
MIX_HC: tl.constexpr,
SINKHORN_ITERS: tl.constexpr,
EPS: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Triton port of hc_split_sinkhorn_kernel (one program per token row).
Layout mirrors the TileLang/torch reference exactly: the flattened ``mixes``
row holds ``pre`` (HC), ``post`` (HC) and the ``comb`` matrix (HC*HC)
consecutively. gfx1250 can't compile TileLang's CK-backed addressing, so this
replaces it while keeping the numerics identical.
"""
row = tl.program_id(0)
if row >= n:
return
scale0 = tl.load(hc_scale_ptr + 0)
scale1 = tl.load(hc_scale_ptr + 1)
scale2 = tl.load(hc_scale_ptr + 2)
j = tl.arange(0, BLOCK)
jmask = j < HC
# pre = sigmoid(mixes[:HC] * scale0 + base[:HC]) + eps
base_pre = tl.load(hc_base_ptr + j, mask=jmask, other=0.0)
mix_pre = tl.load(mixes_ptr + row * MIX_HC + j, mask=jmask, other=0.0)
pre = tl.sigmoid(mix_pre * scale0 + base_pre) + EPS
tl.store(pre_ptr + row * HC + j, pre, mask=jmask)
# post = 2 * sigmoid(mixes[HC:2*HC] * scale1 + base[HC:2*HC])
base_post = tl.load(hc_base_ptr + HC + j, mask=jmask, other=0.0)
mix_post = tl.load(mixes_ptr + row * MIX_HC + HC + j, mask=jmask, other=0.0)
post = 2.0 * tl.sigmoid(mix_post * scale1 + base_post)
tl.store(post_ptr + row * HC + j, post, mask=jmask)
# comb[j, k] = mixes[2*HC + j*HC + k] * scale2 + base[2*HC + j*HC + k]
jj = j[:, None]
kk = j[None, :]
mmask = (jj < HC) & (kk < HC)
coff = 2 * HC + jj * HC + kk
base_c = tl.load(hc_base_ptr + coff, mask=mmask, other=0.0)
mix_c = tl.load(mixes_ptr + row * MIX_HC + coff, mask=mmask, other=0.0)
comb = mix_c * scale2 + base_c
# Initial row softmax (numerically stabilized) then column normalize.
comb_masked = tl.where(mmask, comb, float("-inf"))
row_max = tl.max(comb_masked, axis=1)
comb = tl.exp(comb - row_max[:, None])
comb = tl.where(mmask, comb, 0.0)
row_sum = tl.sum(comb, axis=1)
comb = comb / row_sum[:, None] + EPS
comb = tl.where(mmask, comb, 0.0)
col_sum = tl.sum(comb, axis=0)
comb = comb / (col_sum[None, :] + EPS)
comb = tl.where(mmask, comb, 0.0)
for _ in tl.static_range(SINKHORN_ITERS - 1):
row_sum = tl.sum(comb, axis=1)
comb = comb / (row_sum[:, None] + EPS)
comb = tl.where(mmask, comb, 0.0)
col_sum = tl.sum(comb, axis=0)
comb = comb / (col_sum[None, :] + EPS)
comb = tl.where(mmask, comb, 0.0)
tl.store(comb_ptr + row * HC * HC + jj * HC + kk, comb, mask=mmask)
def _hc_split_sinkhorn_triton(
mixes: torch.Tensor,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
hc_mult: int = 4,
sinkhorn_iters: int = 20,
eps: float = 1e-6,
):
b, s, _ = mixes.size()
hc = hc_mult
mix_hc = (2 + hc) * hc
n = b * s
flat = mixes.reshape(n, mix_hc).float()
scale = hc_scale.float().contiguous()
base = hc_base.float().contiguous()
pre = mixes.new_empty(n, hc, dtype=torch.float32)
post = mixes.new_empty(n, hc, dtype=torch.float32)
comb = mixes.new_empty(n, hc, hc, dtype=torch.float32)
_hc_split_sinkhorn_triton_kernel[(n,)](
flat,
scale,
base,
pre,
post,
comb,
n,
HC=hc,
MIX_HC=mix_hc,
SINKHORN_ITERS=sinkhorn_iters,
EPS=eps,
BLOCK=triton.next_power_of_2(hc),
)
pre = pre.reshape(b, s, hc).to(mixes.dtype)
post = post.reshape(b, s, hc).to(mixes.dtype)
comb = comb.reshape(b, s, hc, hc).to(mixes.dtype)
return pre, post, comb
def hc_split_sinkhorn(
mixes: torch.Tensor,
hc_scale: torch.Tensor,
@@ -197,6 +362,12 @@ def hc_split_sinkhorn(
sinkhorn_iters: int = 20,
eps: float = 1e-6,
):
if is_gfx1250_supported():
# TileLang's CK-backed addressing doesn't compile on gfx1250; use the
# Triton port. _hc_split_sinkhorn_torch is kept as a reference fallback.
return _hc_split_sinkhorn_triton(
mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, eps
)
b, s, _ = mixes.size()
pre = mixes.new_empty(b, s, hc_mult)
post = mixes.new_empty(b, s, hc_mult)
@@ -33,6 +33,7 @@ from sglang.srt.utils import (
get_device_name,
is_cpu,
is_cuda,
is_gfx1250_supported,
is_hip,
is_musa,
is_xpu,
@@ -45,6 +46,7 @@ _is_hip = is_hip()
_is_cuda = is_cuda()
_is_cpu = is_cpu()
_is_musa = is_musa()
_is_gfx1250 = is_gfx1250_supported()
_is_xpu = is_xpu()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
@@ -1027,6 +1029,109 @@ def _w8a8_block_fp8_matmul(
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _w8a8_block_fp8_matmul_gfx1250(
# Pointers to inputs and output
A,
B,
C,
As,
Bs,
# Shape for matmul
M,
N,
K,
# Block size for block-wise quantization
group_n,
group_k,
# Stride for inputs and output
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_cm,
stride_cn,
stride_As_m,
stride_As_k,
stride_Bs_k,
stride_Bs_n,
# Meta-parameters
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
needs_masking: tl.constexpr,
):
"""
gfx1250 (RDNA4) block-fp8 matmul.
The shared ``_w8a8_block_fp8_matmul`` is unusable on gfx1250.
1. fp8 ``tl.dot`` faults at runtime
2. software pipelining (``num_stages`` > 1) miscompiles and yields NaN.
3. the ``offs % M`` / ``offs % N`` modulo-wrap index trick is
intermittently miscompiled into out-of-bounds addresses.
"""
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + (pid % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
# No modulo-wrap on gfx1250; use explicit masks instead.
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_bn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
m_mask = offs_am < M
n_mask = offs_bn < N
offs_am_c = tl.where(m_mask, offs_am, 0)
offs_bn_c = tl.where(n_mask, offs_bn, 0)
offs_k = tl.arange(0, BLOCK_SIZE_K)
a_ptrs = A + (offs_am_c[:, None] * stride_am + offs_k[None, :] * stride_ak)
b_ptrs = B + (offs_k[:, None] * stride_bk + offs_bn_c[None, :] * stride_bn)
As_ptrs = As + offs_am_c * stride_As_m
offs_bsn = offs_bn_c // group_n
Bs_ptrs = Bs + offs_bsn * stride_Bs_n
n_tiles_k_per_group_k = group_k // BLOCK_SIZE_K
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(0, tl.cdiv(K, BLOCK_SIZE_K)):
k_mask = offs_k < K - k * BLOCK_SIZE_K
a = tl.load(a_ptrs, mask=m_mask[:, None] & k_mask[None, :], other=0.0)
b = tl.load(b_ptrs, mask=k_mask[:, None] & n_mask[None, :], other=0.0)
a_s = tl.load(As_ptrs, mask=m_mask, other=0.0)
b_s = tl.load(Bs_ptrs, mask=n_mask, other=0.0)
scale_step_k = tl.where((k + 1) % n_tiles_k_per_group_k == 0, 1, 0)
# Upcast fp8 to bf16 in-register
accumulator += (
tl.dot(a.to(tl.bfloat16), b.to(tl.bfloat16)) * a_s[:, None] * b_s[None, :]
)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += BLOCK_SIZE_K * stride_bk
As_ptrs += scale_step_k * stride_As_k
Bs_ptrs += scale_step_k * stride_Bs_k
if C.dtype.element_ty == tl.bfloat16:
c = accumulator.to(tl.bfloat16)
elif C.dtype.element_ty == tl.float16:
c = accumulator.to(tl.float16)
else:
c = accumulator.to(tl.float32)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
c_ptrs = C + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _w8a8_block_fp8_matmul_unrolledx4(
# Pointers to inputs and output
@@ -1474,6 +1579,12 @@ def w8a8_block_fp8_matmul_triton(
"num_stages": 3,
}
if _is_gfx1250:
config = {**config, "num_stages": 1}
kernel = _w8a8_block_fp8_matmul_gfx1250
else:
kernel = select_w8a8_block_fp8_matmul_kernel(M, N, config)
needs_masking = bool(K % config["BLOCK_SIZE_K"] != 0)
def grid(META):
@@ -1481,8 +1592,6 @@ def w8a8_block_fp8_matmul_triton(
triton.cdiv(M, META["BLOCK_SIZE_M"]) * triton.cdiv(N, META["BLOCK_SIZE_N"]),
)
kernel = select_w8a8_block_fp8_matmul_kernel(M, N, config)
kernel[grid](
A,
B,
@@ -176,6 +176,24 @@ def handle_attention_backend_compatibility(server_args: Any):
# AMD platforms backends
if resolved_view(server_args).attention_backend == "aiter":
if model_config.context_len > 8192:
# The 0.85 covers the extra non-static workspace aiter reserves for
# long contexts, but it is a heuristic for the auto-derived default
# only. Shrinking a value the user picked can push the static budget
# below the model-weight footprint on a nearly full GPU and break
# KV-cache allocation outright, so an explicit value is honored.
if (getattr(server_args, "_raw_input", None) or {}).get(
"mem_fraction_static"
) is not None:
logger.warning(
"attention_backend=aiter with context_len=%d (>8192) "
"normally scales mem_fraction_static by 0.85, but "
"mem_fraction_static=%.3f was set explicitly and will be "
"used as-is. Ensure enough non-static memory is left for "
"attention workspace and CUDA graphs.",
model_config.context_len,
cfg.mem_fraction_static,
)
else:
declare_resolution(
server_args,
"_handle_attention_backend_compatibility",
@@ -750,7 +750,15 @@ class C4IndexerBackendMixin:
envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() and not use_fp4_indexer
)
_use_aiter = envs.SGLANG_OPT_USE_AITER_INDEXER.get() and not use_fp4_indexer
if _c4sl.dim() == 1 and not _use_tilelang and not _use_aiter:
_use_torch_fn = (
envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get() and not use_fp4_indexer
)
if (
_c4sl.dim() == 1
and not _use_tilelang
and not _use_aiter
and not _use_torch_fn
):
_c4sl = _c4sl.unsqueeze(-1)
nonpaged_plan = self._get_nonpaged_indexer_plan(
c4_indexer=c4_indexer,
+23 -2
View File
@@ -86,6 +86,7 @@ from sglang.srt.utils import (
is_cuda,
is_flashinfer_available,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
is_npu,
)
@@ -96,15 +97,19 @@ _is_sm90_supported = _is_cuda and get_platform().is_sm90
_is_sm100_supported = _is_cuda and get_platform().is_sm100
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and is_hip()
_is_gfx95_supported = is_gfx95_supported()
_is_gfx1250_supported = is_gfx1250_supported()
_is_npu = is_npu()
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
if _use_aiter:
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype as _aiter_fp8_dtype
if _is_gfx1250_supported:
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant
else:
from aiter.ops.rmsnorm import add_rmsnorm_quant as _aiter_add_rmsnorm_quant
from aiter.ops.rmsnorm import rmsnorm_quant as _aiter_rmsnorm_quant
from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype as _aiter_fp8_dtype
if _is_gfx95_supported:
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant
@@ -133,6 +138,22 @@ def _fused_rmsnorm_fp8_per_token_quant(
If residual is None: (out_fp8, scale)
If residual provided: ((out_fp8, scale), residual_out)
"""
if _is_gfx1250_supported:
# per-token quant == group quant with group_size == hidden size, giving
# an (M, 1) scale.
N = hidden_states.shape[-1]
(out_fp8, scale), _out1, _out2, residual_out = fused_rms_fp8_group_quant(
hidden_states,
weight,
epsilon,
group_size=N,
dtype_quant=_aiter_fp8_dtype,
res1=residual,
)
if residual is not None:
return (out_fp8, scale), residual_out
return (out_fp8, scale)
M, N = hidden_states.shape
out_fp8 = torch.empty((M, N), dtype=_aiter_fp8_dtype, device=hidden_states.device)
scale = torch.empty(M, dtype=torch.float32, device=hidden_states.device)
+10
View File
@@ -39,6 +39,7 @@ from sglang.srt.utils import (
is_cpu,
is_cuda,
is_flashinfer_available,
is_gfx1250_supported,
is_hip,
is_musa,
is_npu,
@@ -110,6 +111,15 @@ _has_rocm_triton_gemma_rms_norm = False
if _use_aiter:
import aiter as _aiter
from aiter import layernorm2d_fwd as layer_norm
if is_gfx1250_supported():
from aiter.ops.triton.normalization.rmsnorm import (
rms_norm,
)
from aiter.ops.triton.normalization.rmsnorm import (
rmsnorm2d_fwd_with_add as fused_add_rms_norm,
)
else:
from aiter import rmsnorm2d_fwd as rms_norm
from aiter import rmsnorm2d_fwd_with_add as fused_add_rms_norm
@@ -0,0 +1,210 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from https://github.com/vllm-project/vllm/pull/46516
#
# MXFP4-weight / FP8-activation (W4A8) fused MoE for AMD gfx1250 (RDNA / gfx12).
#
# gfx1250's aiter CK/ASM ``fused_moe`` produces garbage for the GPT-OSS MXFP4
# W4A8 layout, so this path routes through aiter's *triton* ``moe_gemm_a8w4``
# kernel instead (the same kernel gfx950 uses). Two gfx1250-specific quirks are
# handled, mirroring the vLLM enablement:
# 1. The in-kernel TDM gather fails to compile on gfx1250, so we disable the
# TDM routing path and gather activation rows into expert-sorted order in
# torch (passing ``gather_indx=None`` to the GEMM).
# 2. The gfx1250 ``moe_gemm_a8w4`` reads a CDNA4-swizzled MX scale as garbage,
# so the weight scale is kept unswizzled and ``swizzle_mx_scale=None`` is
# passed to the kernel.
from __future__ import annotations
import torch
_TDM_DISABLED = False
def _import_aiter_w4a8():
"""Import the aiter triton routing + a8w4 GEMM entry points.
Returns ``(routing, moe_gemm_a8w4, downcast_to_static_fp8)`` or ``None`` if
the installed aiter build does not expose the triton W4A8 path.
"""
try:
try:
import aiter.ops.triton.moe.moe_routing.routing as _routing_mod
except ImportError:
import aiter.ops.triton.moe_routing.routing as _routing_mod
from aiter.ops.triton.moe.moe_op_gemm_a8w4 import moe_gemm_a8w4
from aiter.ops.triton.moe.quant_moe import downcast_to_static_fp8
except ImportError:
return None
global _TDM_DISABLED
if not _TDM_DISABLED:
# gfx1250: the in-kernel TDM gather emitted by the routing sort / GEMM
# fails to compile (``TDM gather dst must be 2D``). Force the non-TDM
# path; we gather activations manually below.
_routing_mod.is_tdm_avail = lambda: False
_TDM_DISABLED = True
return (
_routing_mod.routing,
moe_gemm_a8w4,
downcast_to_static_fp8,
)
def _interleave_gate_up(t: torch.Tensor) -> torch.Tensor:
"""Convert a SEPARATED ``[gate_0..gate_{I-1}, up_0..up_{I-1}]`` first dim
(after the expert dim) into the INTERLEAVED ``[gate_0, up_0, gate_1, up_1,
...]`` order that ``moe_gemm_a8w4``'s fused SwiGLU expects (gate on the
even lanes, up on the odd lanes)."""
e, two_i = t.shape[0], t.shape[1]
i = two_i // 2
rest = t.shape[2:]
t = t.view(e, 2, i, *rest)
perm = (0, 2, 1) + tuple(range(3, t.dim()))
return t.permute(*perm).reshape(e, two_i, *rest).contiguous()
def prepare_w4a8_gfx1250_weights(
w13_weight: torch.Tensor,
w13_weight_scale: torch.Tensor,
w13_weight_bias: torch.Tensor,
w2_weight: torch.Tensor,
w2_weight_scale: torch.Tensor,
w2_weight_bias: torch.Tensor,
):
"""Reshape SGLang's loaded Quark W4A8 MoE buffers into the ``[E, K, N]``
(contraction-major) packed layout consumed by ``moe_gemm_a8w4``.
Input (SGLang / HF Quark layout, per expert), output-channel major:
w13_weight [E, 2I, H//2] uint8 (2 FP4 packed along H)
w13_weight_scale [E, 2I, H//32] uint8 (e8m0), gate/up SEPARATED
w13_weight_bias [E, 2I] fp32
w2_weight [E, H, I//2] uint8
w2_weight_scale [E, H, I//32] uint8 (e8m0)
w2_weight_bias [E, H] fp32
Output (moe_gemm_a8w4 layout), contraction (K) major, gate/up INTERLEAVED
for w13:
w13 [E, H//2, 2I] w13_scale [E, H//32, 2I] w13_bias [E, 2I]
w2 [E, I//2, H] w2_scale [E, I//32, H] w2_bias [E, H]
"""
# Interleave gate/up on w13 (output dim) so the fused SwiGLU picks gate on
# even lanes and up on odd lanes. The interleave output is contiguous, so
# the subsequent transpose(1, 2) yields a *column-major* [E, K, N] view
# (stride(-2) == 1), which ``moe_gemm_a8w4`` requires for MXFP weights.
w13_weight = _interleave_gate_up(w13_weight)
w13_weight_scale = _interleave_gate_up(w13_weight_scale)
w13_weight_bias = _interleave_gate_up(w13_weight_bias)
# Transpose to contraction-major [E, K(packed), N] *without* making it
# contiguous, so the K dimension stays unit-strided (column-major).
w13_weight = w13_weight.transpose(1, 2)
w13_weight_scale = w13_weight_scale.transpose(1, 2)
w2_weight = w2_weight.contiguous().transpose(1, 2)
w2_weight_scale = w2_weight_scale.contiguous().transpose(1, 2)
return (
w13_weight,
w13_weight_scale,
w13_weight_bias.contiguous(),
w2_weight,
w2_weight_scale,
w2_weight_bias.contiguous(),
)
def aiter_w4a8_gfx1250_forward(
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
topk: int,
w13_weight: torch.Tensor,
w13_weight_scale: torch.Tensor,
w13_weight_bias: torch.Tensor,
a13_scale: torch.Tensor,
w2_weight: torch.Tensor,
w2_weight_scale: torch.Tensor,
w2_weight_bias: torch.Tensor,
a2_scale: torch.Tensor,
gemm1_alpha: float,
gemm1_limit: float,
renormalize: bool = True,
apply_router_weight_on_input: bool = False,
) -> torch.Tensor:
"""MXFP4 W4A8 GPT-OSS MoE forward for gfx1250 via aiter triton
``moe_gemm_a8w4``.
``w*`` / ``w*_scale`` / ``w*_bias`` must already be in the
``moe_gemm_a8w4`` layout produced by :func:`prepare_w4a8_gfx1250_weights`.
``a13_scale`` / ``a2_scale`` are the static per-tensor FP8 activation scales
for gate_up_proj and down_proj respectively.
"""
imported = _import_aiter_w4a8()
if imported is None:
raise RuntimeError(
"aiter triton W4A8 MoE (moe_gemm_a8w4) is required for the gfx1250 "
"GPT-OSS MXFP4 path but was not found in the installed aiter build."
)
routing, moe_gemm_a8w4, downcast_to_static_fp8 = imported
assert hidden_states.dtype == torch.bfloat16
# aiter routing on the raw router logits. renormalize=True (GPT-OSS)
# corresponds to applying softmax to the top-k selection inside the kernel
# (sm_first=False).
routing_data, gather_idx, scatter_idx = routing(
router_logits, topk, sm_first=not renormalize
)
gammas = routing_data.gate_scal
# gfx1250: the in-kernel gather is broken, so we pass gather_indx=None to
# moe_gemm_a8w4 and perform the gather ourselves.
gather_src = gather_idx.to(torch.long) // topk
x = hidden_states[gather_src]
if apply_router_weight_on_input:
# Router weights must be applied in bf16 before quantization.
x = x * gammas[:, None].to(x.dtype)
x_fp8 = downcast_to_static_fp8(x, a13_scale)
# GEMM1: FP8 activations x MXFP4 weights, fused SwiGLU, requantize the
# intermediate to FP8 using the down_proj activation scale (a2_scale) so
# GEMM2 can consume it directly.
intermediate_cache1 = moe_gemm_a8w4(
x_fp8,
w13_weight,
None,
w13_weight_scale,
a13_scale,
a2_scale,
w13_weight_bias,
routing_data,
gather_indx=None,
scatter_indx=None,
gammas=None,
swizzle_mx_scale=None,
out_dtype=x_fp8.dtype,
apply_swiglu=True,
alpha=gemm1_alpha,
limit=gemm1_limit,
)
# GEMM2: down projection, scatter back to token order and apply the router
# weights (gammas) unless they were already applied on the input.
intermediate_cache3 = moe_gemm_a8w4(
intermediate_cache1,
w2_weight,
None,
w2_weight_scale,
a2_scale,
None,
w2_weight_bias,
routing_data,
gather_indx=None,
scatter_indx=scatter_idx,
gammas=None if apply_router_weight_on_input else gammas,
swizzle_mx_scale=None,
out_dtype=torch.bfloat16,
)
return intermediate_cache3.contiguous()
+40 -6
View File
@@ -91,6 +91,7 @@ from sglang.srt.utils import (
is_cuda,
is_flashinfer_available,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
is_musa,
is_npu,
@@ -128,6 +129,10 @@ _mxfp8_to_block_fp8_required = mxfp8_block_convert_required() or get_bool_env_va
_use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip
_use_aiter = envs.SGLANG_USE_AITER.get() and _is_hip
_is_shuffle_moe_mxfp4 = is_gfx95_supported()
_is_gfx1250_supported = is_gfx1250_supported()
# gfx1250 grouped MoE runs the a8w4 (fp8 activation) FlyDSL kernel when
# AITER_FORCE_A8W4 is set; that kernel consumes (16,16)-preshuffled weights.
_use_aiter_a8w4 = get_bool_env_var("AITER_FORCE_A8W4", "false")
def _require_fp4_dtype():
@@ -140,7 +145,12 @@ def _require_fp4_dtype():
if _use_aiter or _use_hip_int4:
from aiter.ops.shuffle import shuffle_scale, shuffle_weight
from aiter.ops.shuffle import (
moe_shuffle_scale,
moe_shuffle_weight,
shuffle_scale,
shuffle_weight,
)
if _use_aiter:
from sglang.srt.layers.quantization.fp8_utils import (
@@ -1524,24 +1534,48 @@ class Fp8MoEMethod(FusedMoEMethodBase):
scale = getattr(layer, scale_name)
num_experts, num_rows, _ = scale.shape
is_w13_scale = scale_name == "w13_weight_scale_inv"
if _is_gfx1250_supported:
scale.data = moe_shuffle_scale(
scale.contiguous(),
experts_cnt=num_experts,
is_guinterleave=gu_intv,
gate_up=is_w13_scale,
)
else:
scale_2d = scale.reshape(-1, scale.shape[-1])
scale.data = shuffle_scale(scale_2d, num_experts, gu_intv, is_w13_scale)
scale.data = shuffle_scale(
scale_2d, num_experts, gu_intv, is_w13_scale
)
layer.w13_weight.data = layer.w13_weight.data.view(fp4_weight_dtype)
layer.w2_weight.data = layer.w2_weight.data.view(fp4_weight_dtype)
is_shuffled = _is_shuffle_moe_mxfp4
if is_shuffled:
layer.w13_weight.data = shuffle_weight(
if _is_gfx1250_supported:
is_shuffled = True
layer.w13_weight.data = moe_shuffle_weight(
layer.w13_weight,
is_guinterleave=gu_intv,
gate_up=True,
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight.data = moe_shuffle_weight(
layer.w2_weight,
is_guinterleave=gu_intv,
gate_up=False,
)
else:
is_shuffled = _is_shuffle_moe_mxfp4 or _use_aiter_a8w4
if is_shuffled:
shuffle_gu_intv = gu_intv and not _use_aiter_a8w4
layer.w13_weight.data = shuffle_weight(
layer.w13_weight,
is_guinterleave=shuffle_gu_intv,
gate_up=True,
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight,
is_guinterleave=shuffle_gu_intv,
gate_up=False,
)
layer.w13_weight.is_shuffled = is_shuffled
layer.w2_weight.is_shuffled = is_shuffled
return
@@ -42,6 +42,7 @@ from sglang.srt.utils import (
is_cuda,
is_flashinfer_available,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
is_musa,
is_xpu,
@@ -57,9 +58,15 @@ _is_cuda = is_cuda()
_is_xpu = is_xpu()
_is_fp8_fnuz = is_fp8_fnuz()
_is_gfx95_supported = is_gfx95_supported()
_is_gfx1250_supported = is_gfx1250_supported()
_is_musa = is_musa()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
# gfx1250 (RDNA4) cannot compile the AITER CK quant/GEMM kernels, and even when
# CK builds it lacks the MFMA/WMMA instructions those kernels rely on. Force the
# pure-triton block-fp8 path on gfx1250.
_use_aiter = (
get_bool_env_var("SGLANG_USE_AITER") and _is_hip and not _is_gfx1250_supported
)
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
# ROCm 7.0 hipcc miscompiles gemm_a8w8_blockscale_bpreshuffle on gfx95 (#23319).
_use_aiter_bpreshuffle_gfx95 = _use_aiter_gfx95 and get_hip_version() >= (7, 2, 0)
@@ -1244,19 +1251,29 @@ def triton_w8a8_block_fp8_linear(
input_scale: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert input_scale is None
if input_scale is not None:
# Pre-quantized input: ``input`` is already fp8 and ``input_scale`` is
# its per-group scale (row-major (M, cdiv(K, 128))). Produced on HIP by
# fused act/rmsnorm+quant ops (e.g. fused_clamp_act_mul) that feed the
# GEMM directly. Skip re-quantization and emit bf16.
q_input = input.view(-1, input.shape[-1])
x_scale = input_scale
output_dtype = torch.bfloat16
output_shape = [*input.shape[:-1], weight.shape[0]]
else:
input_2d = input.view(-1, input.shape[-1])
output_dtype = input_2d.dtype
output_shape = [*input.shape[:-1], weight.shape[0]]
q_input, x_scale = per_token_group_quant_fp8(
input_2d, block_size[1], column_major_scales=False
)
output = w8a8_block_fp8_matmul_triton(
q_input, weight, x_scale, weight_scale, block_size, output_dtype=input_2d.dtype
q_input, weight, x_scale, weight_scale, block_size, output_dtype=output_dtype
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
return output.to(dtype=output_dtype).view(*output_shape)
@lru_cache(maxsize=1)
@@ -22,12 +22,66 @@ from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod
from sglang.srt.layers.quantization.online_quantization import CopyNumelCounter
from sglang.srt.layers.quantization.quark.schemes import QuarkLinearScheme
from sglang.srt.layers.quantization.quark.utils import Nvfp4SourceConfig
from sglang.srt.utils import is_hip
from sglang.srt.utils import get_bool_env_var, is_hip
from sglang.srt.utils.common import direct_register_custom_op, is_gfx95_supported
NVFP4_BLOCK_SIZE = 16
_is_hip = is_hip()
# On GPUs that lack the fp4-activation WMMA scale instruction
# (V_WMMA_SCALE_F32_32X16X128_F4, e.g. gfx1250) the a4w4 (fp4 x fp4) linear GEMM
# cannot run. The MoE path is switched to a8w4 via AITER_FORCE_A8W4=1 (handled
# inside aiter.fused_moe); there is currently no working dense a8w4 GEMM for
# plain nn.Linear on this arch, so under the same flag the (few) MXFP4-quantized
# linear layers dequantize their FP4 weights to bf16 once at load and run a
# plain bf16 GEMM. This trades a little memory for correctness on hardware that
# cannot execute the fp4 kernel at all.
_dequant_linear_to_bf16 = _is_hip and get_bool_env_var("AITER_FORCE_A8W4", "false")
# MXFP4 (OCP MX FP4 / e2m1) decode table, indexed by the 4-bit code.
_MXFP4_VALUES = [
0.0,
0.5,
1.0,
1.5,
2.0,
3.0,
4.0,
6.0,
-0.0,
-0.5,
-1.0,
-1.5,
-2.0,
-3.0,
-4.0,
-6.0,
]
def _dequant_mxfp4_to_bf16(
weight: torch.Tensor, weight_scale: torch.Tensor
) -> torch.Tensor:
"""Dequantize a packed MXFP4 weight ``(N, K//2)`` uint8 + e8m0 group scale
``(N, K//32)`` uint8 into a dense bf16 weight ``(N, K)``."""
N, k_packed = weight.shape
K = k_packed * 2
lut = torch.tensor(_MXFP4_VALUES, device=weight.device, dtype=torch.float32)
lo = (weight & 0xF).long()
hi = (weight >> 4).long()
vals = torch.empty(N, K, device=weight.device, dtype=torch.float32)
vals[:, 0::2] = lut[lo]
vals[:, 1::2] = lut[hi]
# e8m0 byte b decodes to 2^(b-127); 255 is the NaN/Inf sentinel (unused by
# real weights) -> map to 0 so it can't poison the matmul.
scale = torch.exp2(weight_scale.to(torch.float32) - 127.0)
scale = torch.where(weight_scale == 255, torch.zeros_like(scale), scale)
scale = scale.view(N, K // 32, 1)
w = (vals.view(N, K // 32, 32) * scale).view(N, K)
return w.to(torch.bfloat16)
if _is_hip:
from aiter.ops.triton.gemm.fused.fused_gemm_afp4wfp4_split_cat import (
fused_gemm_afp4wfp4_split_cat as _fused_gemm_afp4wfp4_split_cat_orig,
@@ -211,6 +265,13 @@ class QuarkW4A4MXFP4(QuarkLinearScheme):
assert layer.weight.dtype == torch.uint8
assert layer.weight_scale.dtype == torch.uint8
if _dequant_linear_to_bf16:
w_bf16 = _dequant_mxfp4_to_bf16(layer.weight.data, layer.weight_scale.data)
layer.weight = torch.nn.Parameter(w_bf16, requires_grad=False)
# FP4 block scales are folded into the bf16 weight; drop them.
layer.weight_scale = None
layer.dequantized_bf16 = True
def create_weights(
self,
layer: torch.nn.Module,
@@ -620,6 +681,16 @@ class QuarkW4A4MXFP4(QuarkLinearScheme):
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
# bf16 fallback: FP4 weights were dequantized to bf16 at load time
# because this HW cannot run the fp4 GEMM. Run a plain bf16 linear.
# (The fused tuple-input paths below are only used by MLA attention
# projections, which are excluded from quantization for this checkpoint,
# so a plain-tensor activation is what reaches here.)
if getattr(layer, "dequantized_bf16", False):
if isinstance(x, tuple):
x = x[0]
return torch.nn.functional.linear(x, layer.weight, bias)
# Bias will be added after the GEMM if provided
three_d = False
fused_gemm_split_cat = False
@@ -23,10 +23,10 @@ from sglang.srt.layers.quantization.quark.utils import Nvfp4SourceConfig
from sglang.srt.utils import (
get_bool_env_var,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
set_weight_attrs,
)
from sglang.srt.utils.common import is_gfx95_supported
NVFP4_BLOCK_SIZE = 16
@@ -45,9 +45,30 @@ __all__ = ["QuarkW4A4MXFp4MoE"]
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.shuffle import moe_shuffle_scale, moe_shuffle_weight, shuffle_weight
from aiter.utility.fp4_utils import e8m0_shuffle
# gfx1250's grouped MoE GEMM reads weight scales in the n32k4 layout
# (moe_shuffle_scale -> shuffle_scale_n32k4), not the e8m0_shuffle layout used by
# gfx950. Using the wrong layout silently corrupts the dequant scales.
_is_gfx1250 = is_gfx1250_supported()
# The gfx1250 a8w4 grouped MoE kernel consumes (16,16)-preshuffled FP4 weights
# (see aiter op_tests/test_flydsl_grouped_gemm_gfx1250.py, which always does
# shuffle_weight(w, layout=(16,16)); the DSv4 fp8.py path shuffles the same way
# when AITER_FORCE_A8W4 is set). Historically this scheme only shuffled on gfx95
# (_is_shuffle_moe_mxfp4), so on gfx1250 the raw (unshuffled) weight layout was
# fed to a kernel expecting the shuffled one -> garbage. Mirror DSv4: shuffle
# whenever the a8w4 path is forced on gfx1250. SGLANG_MOE_SHUFFLE_GFX1250=false
# reproduces the old (unshuffled) behavior for A/B comparison.
_use_aiter_a8w4 = get_bool_env_var("AITER_FORCE_A8W4", "false")
_shuffle_moe_gfx1250 = (
_is_gfx1250
and _use_aiter_a8w4
and get_bool_env_var("SGLANG_MOE_SHUFFLE_GFX1250", "true")
)
if _is_hip:
from aiter.ops.triton.quant import dynamic_mxfp4_quant
else:
@@ -814,6 +835,19 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
assert layer.w2_weight_scale.dtype == torch.uint8
# Pre-shuffle weight scales
if _is_gfx1250:
# gfx1250 grouped MoE GEMM consumes B-scales in the n32k4 layout.
num_experts = layer.w13_weight_scale.shape[0]
layer.w13_weight_scale.data = moe_shuffle_scale(
layer.w13_weight_scale.contiguous(),
experts_cnt=num_experts,
is_guinterleave=True,
gate_up=True,
)
layer.w2_weight_scale.data = moe_shuffle_scale(
layer.w2_weight_scale.contiguous(), experts_cnt=num_experts
)
else:
s0, s1, _ = layer.w13_weight_scale.shape
w13_weight_scale = layer.w13_weight_scale.view(s0 * s1, -1)
w13_weight_scale = e8m0_shuffle(w13_weight_scale)
@@ -823,9 +857,18 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w2_weight_scale = layer.w2_weight_scale.view(s0 * s1, -1)
w2_weight_scale = e8m0_shuffle(w2_weight_scale)
layer.w2_weight_scale.data = w2_weight_scale.view(s0, s1, -1)
# Pre-shuffle weight
if _is_shuffle_moe_mxfp4:
if _is_gfx1250:
# gfx1250 grouped kernel expects GUGU (gate/up row-interleaved) layout.
# moe_shuffle_weight does interleave_gate_up_rows then tile shuffle,
# which is what grouped_gemm_gfx1250_a8w4 reads.
layer.w13_weight.data = moe_shuffle_weight(
layer.w13_weight.contiguous(), is_guinterleave=True, gate_up=True
)
layer.w2_weight.data = moe_shuffle_weight(layer.w2_weight.contiguous())
layer.w13_weight.is_shuffled = True
layer.w2_weight.is_shuffled = True
elif _is_shuffle_moe_mxfp4:
layer.w13_weight.data = shuffle_weight(
layer.w13_weight.contiguous(), (16, 16)
)
@@ -879,6 +922,13 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w13_weight.is_shuffled = True
w2_weight.is_shuffled = True
if _is_gfx1250:
from aiter.ops.flydsl.moe_common import GateMode
_fused_moe_kwargs = {"gate_mode": GateMode.INTERLEAVE.value}
else:
_fused_moe_kwargs = None
quant_info = AiterMoeQuantInfo(
w13_weight=w13_weight,
w2_weight=w2_weight,
@@ -886,5 +936,6 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w13_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
expert_mask=layer.dispatcher.expert_mask_gpu,
fused_moe_kwargs=_fused_moe_kwargs,
)
return self.runner.run(dispatch_output, quant_info)
@@ -16,6 +16,7 @@ from sglang.srt.layers.quantization.utils import all_close_1d
from sglang.srt.utils import (
get_bool_env_var,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
round_up,
set_weight_attrs,
@@ -35,7 +36,19 @@ __all__ = ["QuarkW4A8MXFp4MoE"]
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
# gfx1250 (RDNA / gfx12): the aiter CK/ASM fused_moe MXFP4 W4A8 path produces
# incorrect results, so route through aiter's triton ``moe_gemm_a8w4`` kernel
# instead. This uses an unpadded, contraction-major weight layout (no aiter
# shuffle) prepared in ``process_weights_after_loading``.
_use_gfx1250_w4a8 = (
_is_hip
and is_gfx1250_supported()
and get_bool_env_var("SGLANG_GFX1250_W4A8_MOE", "true")
)
# Whether to use the aiter CK/ASM shuffled weight layout (gfx950 etc.). Disabled
# on gfx1250, which uses the triton moe_gemm_a8w4 layout.
_use_aiter_shuffle_layout = _use_aiter and not _use_gfx1250_w4a8
if _use_aiter and not _use_gfx1250_w4a8:
from aiter.ops.shuffle import (
shuffle_scale,
shuffle_scale_a16w4,
@@ -101,7 +114,7 @@ class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
self.num_experts = num_experts
self.with_bias = extra_weight_attrs.get("with_bias", False)
if _use_aiter:
if _use_aiter_shuffle_layout:
intermediate_size_per_partition_after_pad = round_up(
intermediate_size_per_partition, 256
)
@@ -112,13 +125,16 @@ class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
- layer.intermediate_size_per_partition
)
else:
# Non-shuffled layout (gfx1250 triton moe_gemm_a8w4 path and the
# generic non-aiter path): keep buffers unpadded so they match the
# HF checkpoint shape exactly.
intermediate_size_per_partition_after_pad = intermediate_size_per_partition
self.hidden_pad = 0
self.intermediate_pad = 0
w13_up_dim, w2_down_dim, weight_padded = get_moe_weight_sizes(
intermediate_size_per_partition_after_pad,
is_aiter_moe=_use_aiter,
is_aiter_moe=_use_aiter_shuffle_layout,
is_concat=True,
is_packed=True,
)
@@ -240,7 +256,61 @@ class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
set_weight_attrs(w13_input_scale, extra_weight_attrs)
set_weight_attrs(w2_input_scale, extra_weight_attrs)
def _process_weights_gfx1250(self, layer: torch.nn.Module) -> None:
"""gfx1250: reshape MXFP4 W4A8 weights into the contraction-major,
gate/up-interleaved layout consumed by aiter's triton ``moe_gemm_a8w4``
(no CK/ASM shuffle)."""
from sglang.srt.layers.moe.fused_moe_triton.aiter_mxfp4_w4a8_moe import (
prepare_w4a8_gfx1250_weights,
)
(
w13_weight,
w13_weight_scale,
w13_weight_bias,
w2_weight,
w2_weight_scale,
w2_weight_bias,
) = prepare_w4a8_gfx1250_weights(
layer.w13_weight.data,
layer.w13_weight_scale.data,
layer.w13_weight_bias.data,
layer.w2_weight.data,
layer.w2_weight_scale.data,
layer.w2_weight_bias.data,
)
layer.w13_weight = torch.nn.Parameter(w13_weight, requires_grad=False)
layer.w13_weight_scale = torch.nn.Parameter(
w13_weight_scale, requires_grad=False
)
layer.w13_weight_bias = torch.nn.Parameter(w13_weight_bias, requires_grad=False)
layer.w2_weight = torch.nn.Parameter(w2_weight, requires_grad=False)
layer.w2_weight_scale = torch.nn.Parameter(w2_weight_scale, requires_grad=False)
layer.w2_weight_bias = torch.nn.Parameter(w2_weight_bias, requires_grad=False)
# Static FP8 MoE kernels consume a single activation scale. Use the
# maximum if expert-local checkpoint scales differ.
if layer.w13_input_scale is None or layer.w2_input_scale is None:
raise ValueError("W4A8 MXFP4-FP8 MoE requires static input scales.")
if not all_close_1d(layer.w13_input_scale) or not all_close_1d(
layer.w2_input_scale
):
logger.warning(
"Found input_scales that are not equal for W4A8 MXFP4-FP8 "
"MoE layer. Using the maximum across experts for each layer."
)
layer.w13_input_scale = torch.nn.Parameter(
layer.w13_input_scale.max().to(torch.float32), requires_grad=False
)
layer.w2_input_scale = torch.nn.Parameter(
layer.w2_input_scale.max().to(torch.float32), requires_grad=False
)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
if _use_gfx1250_w4a8:
self._process_weights_gfx1250(layer)
return
# Mirror native MXFP4 post-load shuffling. The default
# `SGLANG_USE_AITER_MOE_GU_ITLV=1` path uses the gate-up-aware
# a16w4 layout; the `=0` fallback keeps the separated gate/up layout.
@@ -334,6 +404,10 @@ class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
)
self.moe_runner_config = moe_runner_config
if _use_gfx1250_w4a8:
# gfx1250 bypasses the MoeRunner and calls aiter's triton
# moe_gemm_a8w4 directly in ``apply_weights``.
return
moe_runner_backend = get_moe_runner_backend()
if _use_aiter and get_moe_a2a_backend().supports_aiter():
moe_runner_backend = MoeRunnerBackend.AITER
@@ -353,6 +427,40 @@ class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
layer: torch.nn.Module,
dispatch_output: StandardDispatchOutput,
) -> CombineInput:
if _use_gfx1250_w4a8:
from sglang.srt.layers.moe.fused_moe_triton.aiter_mxfp4_w4a8_moe import (
aiter_w4a8_gfx1250_forward,
)
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
topk_weights, topk_ids, router_logits = dispatch_output.topk_output
x = dispatch_output.hidden_states
if x.shape[-1] != self.hidden_size:
x = x[..., : self.hidden_size]
cfg = self.moe_runner_config
alpha = cfg.gemm1_alpha if cfg.gemm1_alpha is not None else 1.702
limit = cfg.gemm1_clamp_limit or cfg.swiglu_limit or 7.0
output = aiter_w4a8_gfx1250_forward(
hidden_states=x,
router_logits=router_logits,
topk=topk_ids.shape[-1],
w13_weight=layer.w13_weight,
w13_weight_scale=layer.w13_weight_scale,
w13_weight_bias=layer.w13_weight_bias,
a13_scale=layer.w13_input_scale,
w2_weight=layer.w2_weight,
w2_weight_scale=layer.w2_weight_scale,
w2_weight_bias=layer.w2_weight_bias,
a2_scale=layer.w2_input_scale,
gemm1_alpha=alpha,
gemm1_limit=limit,
renormalize=True,
apply_router_weight_on_input=cfg.apply_router_weight_on_input,
)
return StandardCombineInput(hidden_states=output)
from sglang.srt.layers.moe.moe_runner.aiter import (
AiterMoeQuantInfo,
AiterQuantType,
@@ -2,6 +2,7 @@
from __future__ import annotations
import functools
import logging
from typing import Any, Dict, Optional, Tuple
@@ -57,6 +58,19 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
from aiter.rotary_embedding import get_rope as aiter_get_rope
@functools.lru_cache(maxsize=1)
def _aiter_rope_unsupported_arch() -> bool:
"""aiter's rope kernels (csrc/kernels/rope/rope_common.h) depend on ck_tile
types that do not build on gfx1250; fall back to sglang's native rope there."""
if not _is_hip:
return False
try:
return "gfx1250" in torch.cuda.get_device_properties(0).gcnArchName
except Exception:
return False
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
@@ -467,7 +481,8 @@ def get_rope_wrapper(
device: Optional[str] = None,
):
if device != "cpu":
wrapper = aiter_get_rope if _use_aiter else get_rope
use_aiter_rope = _use_aiter and not _aiter_rope_unsupported_arch()
wrapper = aiter_get_rope if use_aiter_rope else get_rope
return wrapper(
head_size,
rotary_dim,
+4 -1
View File
@@ -23,6 +23,7 @@ from sglang.srt.utils.async_probe import sanitize_nan_logits
from sglang.srt.utils.common import (
get_bool_env_var,
is_cuda,
is_gfx1250_supported,
is_hip,
is_musa,
is_npu,
@@ -56,7 +57,9 @@ if _use_aiter:
# to an empty string and breaks downstream consumers. Set this to 1 to fall back to
# torch.argmax (which always returns a valid index). Default off so behavior is
# unchanged elsewhere.
_disable_aiter_greedy_sample = get_bool_env_var("SGLANG_DISABLE_AITER_GREEDY_SAMPLE")
_disable_aiter_greedy_sample = (
get_bool_env_var("SGLANG_DISABLE_AITER_GREEDY_SAMPLE") or is_gfx1250_supported()
)
if is_npu():
import torch_npu
@@ -28,6 +28,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
is_in_breakable_cuda_graph,
)
from sglang.srt.runtime_context import get_flags
from sglang.srt.utils import is_gfx1250_supported
# Detect whether the current forward pass is in capture mode.
is_capture_mode = False
@@ -56,7 +57,7 @@ def compile_in_capture_mode(func):
torch.compile during cuda-graph capture without paying the
compilation cost in the eager forward path.
"""
if is_capture_mode:
if is_capture_mode and not is_gfx1250_supported():
return torch.compile(func)
return func
@@ -7,11 +7,13 @@ import triton
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_platform
from sglang.srt.utils import get_bool_env_var, is_gfx95_supported, is_hip
from sglang.srt.utils.common import is_gfx1250_supported
logger = logging.getLogger(__name__)
_is_hip = is_hip()
_is_gfx95_supported = is_gfx95_supported()
_is_gfx1250_supported = is_gfx1250_supported()
_FUSED_HC_POST_PRE_M_THRESHOLD = 64
_FUSED_HC_POST_PRE_CACHE: dict[tuple, dict[str, torch.Tensor]] = {}
@@ -26,6 +28,11 @@ _AITER_MHC_IMPORT_WARNED = False
def _is_fused_mhc_post_pre_enabled() -> bool:
# gfx1250: TileLang doesn't compile; the fused cross-layer path routes
# entirely through the Triton mhc_post_pre (try_fused_hc_post_pre).
# Gate only on SGLANG_OPT_FUSE_MHC_POST_PRE; TileLang switches don't apply.
if _is_gfx1250_supported:
return envs.SGLANG_OPT_FUSE_MHC_POST_PRE.get()
# SM120 disables the standalone TileLang pre path. mhc_fused_post_pre does
# not read that flag and dispatches independently for both small and large
# token batches, so the standalone pre flag must not veto the fused opt-in.
@@ -144,9 +151,11 @@ def try_fused_hc_post_pre(
if (
_TRITON_MHC_POST_PRE_RUNTIME_DISABLED
or not is_gfx95_supported
or not (is_gfx95_supported or _is_gfx1250_supported)
or x.shape[0] == 0
or x.shape[0] > _FUSED_HC_POST_PRE_M_THRESHOLD
# gfx1250 runs the fused cross-layer path for ALL sizes (prefill+decode);
# there is no TileLang fallback available, so don't cap by M there.
or (x.shape[0] > _FUSED_HC_POST_PRE_M_THRESHOLD and not _is_gfx1250_supported)
or x.dim() != 2
or residual.dim() != 3
):
@@ -65,7 +65,7 @@ from sglang.srt.runtime_context import get_exec, get_parallel
from sglang.srt.state_capturer.indexer_topk import (
maybe_capture_indexer_topk,
)
from sglang.srt.utils import BumpAllocator
from sglang.srt.utils import BumpAllocator, get_bool_env_var
logger = logging.getLogger(__name__)
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
@@ -74,6 +74,22 @@ if TYPE_CHECKING:
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
if _use_aiter:
# On gfx1250 the aiter `module_fused_qk_norm_rope_cache_quant_shuffle` kernel
# fails to JIT-build (its `rope_common.h` / `ck_tile/vec_convert.h` are
# incompatible with this image's composable_kernel), which crashes the very
# first MLA forward. This path is a pure RMSNorm (quant_type=No), so under the
# gfx1250 workaround flag (AITER_FORCE_A8W4) substitute a self-contained Triton
# RMSNorm that never touches the aiter fp4 kernel build.
if get_bool_env_var("AITER_FORCE_A8W4", "false"):
if get_bool_env_var("SGLANG_QK_RMSNORM_TORCH", "false"):
from sglang.srt.models.deepseek_common.attention_forward_methods.triton_qk_rmsnorm import (
fused_qk_rmsnorm_torch as fused_qk_rmsnorm_bf16,
)
else:
from sglang.srt.models.deepseek_common.attention_forward_methods.triton_qk_rmsnorm import (
fused_qk_rmsnorm_triton as fused_qk_rmsnorm_bf16,
)
else:
# aiter ROCm/aiter#2958 renamed the public `fused_qk_rmsnorm` in
# `aiter.ops.fused_qk_norm_rope_cache_quant` to a private `_fused_qk_rmsnorm`
# and introduced a unified entry point in `aiter.ops.fused_qk_rmsnorm_group_quant`
@@ -552,7 +568,15 @@ class DeepseekMLARocmForwardMixin:
not _use_aiter
or not _is_gfx95_supported
or self.use_dsa
or self.current_attention_backend == "triton"
# Non-fused, non-specialized attention backends (e.g. Triton) run
# the cat path in forward_absorb_core and need RoPE applied here;
# only the aiter fused MLA path and the specialized MLA backends
# defer RoPE to their own kernels.
or (
self.current_attention_backend
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
and self.current_attention_backend != "aiter"
)
)
):
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
@@ -0,0 +1,108 @@
"""Triton fused q/k RMSNorm for MLA attention.
Drop-in replacement for aiter's ``fused_qk_rmsnorm`` on hardware where the aiter
``module_fused_qk_norm_rope_cache_quant_shuffle`` kernel cannot be built (e.g.
gfx1250, whose composable_kernel version is incompatible with the aiter fork's
``rope_common.h`` / ``ck_tile/vec_convert.h``). Semantics match a plain
RMSNorm (``sglang.srt.layers.layernorm.RMSNorm.forward_native``): compute the
row variance in fp32, scale by ``rsqrt(var + eps)``, multiply by ``weight`` and
cast back to the input dtype. No RoPE, no quantization -- this is only used on
the ``quant_type=No`` path where the fused kernel is a pure RMSNorm.
"""
from typing import Tuple
import torch
import triton
import triton.language as tl
@triton.jit
def _rmsnorm_kernel(
x_ptr,
w_ptr,
out_ptr,
row_stride,
N,
eps,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0)
x_row = x_ptr + row * row_stride
out_row = out_ptr + row * row_stride
cols = tl.arange(0, BLOCK_SIZE)
mask = cols < N
x = tl.load(x_row + cols, mask=mask, other=0.0).to(tl.float32)
var = tl.sum(x * x, axis=0) / N
rstd = 1.0 / tl.sqrt(var + eps)
w = tl.load(w_ptr + cols, mask=mask, other=0.0).to(tl.float32)
y = x * rstd * w
tl.store(out_row + cols, y.to(out_row.dtype.element_ty), mask=mask)
def _rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
orig_shape = x.shape
N = orig_shape[-1]
x2d = x.reshape(-1, N).contiguous()
out = torch.empty_like(x2d)
M = x2d.shape[0]
if M == 0:
return out.reshape(orig_shape)
BLOCK_SIZE = triton.next_power_of_2(N)
num_warps = min(max(BLOCK_SIZE // 256, 1), 16)
_rmsnorm_kernel[(M,)](
x2d,
weight,
out,
x2d.stride(0),
N,
float(eps),
BLOCK_SIZE=BLOCK_SIZE,
num_warps=num_warps,
)
return out.reshape(orig_shape)
def fused_qk_rmsnorm_triton(
q: torch.Tensor,
q_weight: torch.Tensor,
q_eps: float,
k: torch.Tensor,
k_weight: torch.Tensor,
k_eps: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""RMSNorm ``q`` (with ``q_weight``/``q_eps``) and ``k`` (with
``k_weight``/``k_eps``) independently. Matches the signature and return
convention of the aiter ``fused_qk_rmsnorm`` shim used in ``forward_mla``.
"""
q_out = _rmsnorm(q, q_weight, q_eps)
k_out = _rmsnorm(k, k_weight, k_eps)
return q_out, k_out
def _rmsnorm_torch(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
orig_dtype = x.dtype
xf = x.to(torch.float32)
var = xf.pow(2).mean(dim=-1, keepdim=True)
xf = xf * torch.rsqrt(var + eps)
return (xf * weight).to(orig_dtype)
def fused_qk_rmsnorm_torch(
q: torch.Tensor,
q_weight: torch.Tensor,
q_eps: float,
k: torch.Tensor,
k_weight: torch.Tensor,
k_eps: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Pure-torch reference equivalent of ``fused_qk_rmsnorm_triton`` (bisect
aid for the decode degeneration investigation)."""
return _rmsnorm_torch(q, q_weight, q_eps), _rmsnorm_torch(k, k_weight, k_eps)
@@ -29,6 +29,7 @@ from sglang.srt.utils import (
is_cpu,
is_cuda,
is_gfx95_supported,
is_gfx1250_supported,
is_hip,
is_musa,
is_npu,
@@ -47,10 +48,14 @@ _is_cpu = is_cpu()
_is_xpu = is_xpu()
_device_sm = get_device_sm()
_is_gfx95_supported = is_gfx95_supported()
# gfx1250 reuses the gfx95 (CDNA4) code paths for MXFP4 q/k-norm kernels, but its
# aiter rope kernels (ck_tile) do not build, so it runs sglang's native rope which
# lacks the separate cos_cache/sin_cache buffers the gfx95 fused-rope decode path
# expects. This flag lets gfx1250 carve out of that fused-rope path.
_is_gfx1250_supported = is_gfx1250_supported()
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
_use_aiter_bpreshuffle_gfx95 = _use_aiter_gfx95 and get_hip_version() >= (7, 2, 0)
_is_cublas_ge_129 = is_nvidia_cublas_version_ge_12_9()
logger = logging.getLogger(__name__)
+11 -16
View File
@@ -175,6 +175,7 @@ from sglang.srt.utils import (
get_bool_env_var,
is_gfx95_supported,
is_gfx942_supported,
is_gfx1250_supported,
log_info_on_rank0,
make_layers,
)
@@ -303,9 +304,10 @@ _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_SHARED_EXPERT_LOCAL = get_bool_env_var("SGLANG_DP_SHARED_EXPERT_LOCAL")
_is_gfx95_supported = is_gfx95_supported()
_is_gfx942_supported = is_gfx942_supported()
_is_gfx1250_supported = is_gfx1250_supported()
if _use_aiter:
if _is_gfx95_supported:
if _is_gfx95_supported or _is_gfx1250_supported:
from aiter.ops.triton.fused_fp8_quant import fused_rms_fp8_group_quant
@@ -1223,7 +1225,7 @@ class MQALayer(MqaAttentionBase):
qkv_a = None
if self.use_fused_qk_norm_rope:
if _is_gfx95_supported:
if _is_gfx95_supported or _is_gfx1250_supported:
q_for_wqb, q_lora = _fused_rmsnorm_fp8_quant(
q_lora,
self.q_norm.weight,
@@ -1340,7 +1342,7 @@ class MQALayer(MqaAttentionBase):
)
if do_fused_qk_norm_rope:
if _is_gfx95_supported:
if _is_gfx95_supported or _is_gfx1250_supported:
q_for_wqb, q_lora = _fused_rmsnorm_fp8_quant(
q_lora,
self.q_norm.weight,
@@ -2063,6 +2065,8 @@ class DeepseekV4DecoderLayer(nn.Module):
use_fused = self.use_fused_mhc_post_pre
if prev_residual is not None and use_fused:
# Dispatch cascade: aiter HIP (gfx95) -> Triton (gfx95 small-batch
# <=64 tokens, or gfx1250 all sizes) -> TileLang -> None.
input_norm_weight = (
self._input_layernorm_weight_bf16
if self._input_layernorm_weight_bf16 is not None
@@ -2088,11 +2092,9 @@ class DeepseekV4DecoderLayer(nn.Module):
if fused is not None:
residual, hidden_states, post, comb, norm_fused = fused
if not norm_fused:
# The Triton fused post+pre returns the layer input WITHOUT
# the input layernorm applied (norm_fused=False). Apply it
# (fp8-quant on aiter gfx95) before attention, exactly as the
# unfused hc_pre path below does; otherwise unnormalized
# activations reach self_attn.
# Triton fused post+pre (gfx95 small-batch or gfx1250) returns
# norm_fused=False — the input layernorm is NOT folded.
# gfx95 takes the fp8-quant path; gfx1250 takes plain layernorm.
if _use_aiter and _is_gfx95_supported:
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
hidden_states,
@@ -2105,10 +2107,6 @@ class DeepseekV4DecoderLayer(nn.Module):
else:
x_quant = None
else:
# Fused dispatch declined: close the previous layer's deferred
# mHC post (prev_residual/prev_post/prev_comb) before opening this
# layer's pre. Skipping hc_post here would drop the previous-layer
# post state and corrupt all subsequent layers.
hidden_states = self.hc_post(
hidden_states, prev_residual, prev_post, prev_comb
)
@@ -2190,9 +2188,6 @@ class DeepseekV4DecoderLayer(nn.Module):
if fused is not None:
residual, hidden_states, post, comb, norm_fused = fused
if not norm_fused:
# The Triton fused post+pre skips the post-attention
# layernorm (norm_fused=False); apply it before the MoE,
# matching the unfused hc_pre path below.
hidden_states = self.post_attention_layernorm(hidden_states)
else:
hidden_states = self.hc_post(hidden_states, residual, post, comb)
@@ -2417,7 +2412,7 @@ class DeepseekV4DecoderLayer(nn.Module):
forward_batch=forward_batch,
)
if not norm_fused:
if _use_aiter and _is_gfx95_supported:
if _use_aiter and (_is_gfx95_supported or _is_gfx1250_supported):
x_quant, hidden_states = _fused_rmsnorm_fp8_quant(
hidden_states,
self.input_layernorm.weight,
+12
View File
@@ -1052,6 +1052,18 @@ def is_gfx942_supported():
return False
@lru_cache(maxsize=1)
def is_gfx1250_supported():
"""
Returns whether the current platform is AMD RDNA4 (gfx1250).
"""
if torch.version.hip:
gcn_arch = torch.cuda.get_device_properties(0).gcnArchName
return any(gfx in gcn_arch for gfx in ["gfx1250"])
else:
return False
def get_hip_version():
if torch.version.hip:
return tuple(map(int, torch.version.hip.split("-")[0].split(".")))
+22 -7
View File
@@ -50,7 +50,7 @@ fi
IMAGE_TORCH_VERSION=$(docker exec ci_sglang python3 -c 'import torch; print(torch.__version__)')
IMAGE_HIP_VERSION=$(docker exec ci_sglang python3 -c 'import torch; print(torch.version.hip or "")')
IMAGE_GPU_ARCH=$(docker exec ci_sglang printenv GPU_ARCH 2>/dev/null || true)
if [[ "${IMAGE_GPU_ARCH}" =~ ^(gfx942|gfx950)(-rocm720|-rocm724|-rocm1000)?$ ]]; then
if [[ "${IMAGE_GPU_ARCH}" =~ ^(gfx942|gfx950|gfx1250)(-rocm720|-rocm724|-rocm1000)?$ ]]; then
echo "[CI-IMAGE] Image GPU_ARCH=${IMAGE_GPU_ARCH}"
case "${IMAGE_GPU_ARCH}" in
*-rocm724) IMAGE_BASE_ARG_SUFFIX="_ROCM724"; IMAGE_STAGE_SUFFIX="-rocm724" ;;
@@ -259,6 +259,11 @@ if docker exec ci_sglang test -d /sgl-workspace/mori; then
MORI_REPO=$(grep -E '^[[:space:]]*ARG[[:space:]]+MORI_REPO=' docker/rocm.Dockerfile | head -n1 | sed 's/.*MORI_REPO="\([^"]*\)".*/\1/')
MORI_COMMIT=$(grep -E '^[[:space:]]*ARG[[:space:]]+MORI_COMMIT=' docker/rocm.Dockerfile | head -n1 | sed 's/.*MORI_COMMIT="\([^"]*\)".*/\1/')
if [[ -z "${MORI_COMMIT}" ]]; then
echo "[MORI] ERROR: Failed to extract MORI_COMMIT from Dockerfile"
exit 1
fi
if [[ "${GPU_ARCH}" == "mi35x" ]]; then
MORI_GPU_ARCHS="gfx950"
else
@@ -290,8 +295,18 @@ if docker exec ci_sglang test -d /sgl-workspace/mori; then
fi
# The pip ROCm SDK vendors NUMA and libdrm under rocm_sysdeps, outside the
# default compiler, CMake, and linker search paths used by MORI.
# gfx1250 additionally needs the SDK's own cmake trees on the prefix path:
# that is what lets hsakmt-config.cmake resolve find_dependency(NUMA)
# without patching MORI's CMakeLists. Same split as docker/rocm.Dockerfile.
ROCM_SYSDEPS="\${ROCM_HOME:-/opt/rocm}/lib/rocm_sysdeps"
if [ '${IMAGE_STAGE_SUFFIX}' = '-rocm1000' ] && [ -d "\${ROCM_SYSDEPS}" ]; then
if [ '${IMAGE_GFX}' = 'gfx1250' ] && [ -d "\${ROCM_SYSDEPS}" ]; then
export PATH="\${ROCM_HOME}/bin:\${PATH}"
export CMAKE_PREFIX_PATH="\${ROCM_SYSDEPS}/lib/cmake:\${ROCM_SYSDEPS}:\${ROCM_HOME}/lib/cmake:\${ROCM_HOME}\${CMAKE_PREFIX_PATH:+:\${CMAKE_PREFIX_PATH}}"
export CPATH="\${ROCM_SYSDEPS}/include\${CPATH:+:\${CPATH}}"
export LIBRARY_PATH="\${ROCM_SYSDEPS}/lib\${LIBRARY_PATH:+:\${LIBRARY_PATH}}"
echo "\${ROCM_SYSDEPS}/lib" > /etc/ld.so.conf.d/rocm-sysdeps.conf
ldconfig
elif [ '${IMAGE_STAGE_SUFFIX}' = '-rocm1000' ] && [ -d "\${ROCM_SYSDEPS}" ]; then
export CMAKE_PREFIX_PATH="\${ROCM_SYSDEPS}\${CMAKE_PREFIX_PATH:+:\${CMAKE_PREFIX_PATH}}"
export CPATH="\${ROCM_SYSDEPS}/include\${CPATH:+:\${CPATH}}"
export LIBRARY_PATH="\${ROCM_SYSDEPS}/lib\${LIBRARY_PATH:+:\${LIBRARY_PATH}}"
@@ -331,11 +346,11 @@ echo "[CI-AITER-CHECK] Runner GPU_ARCH=${GPU_ARCH}"
# 1. Extract AITER_COMMIT from the Dockerfile stage that built this image, as
# identified near the top of this script.
#############################################
if [[ "${IMAGE_GFX}" == "gfx950" ]]; then
_from_line="FROM \$BASE_IMAGE_950${IMAGE_BASE_ARG_SUFFIX} AS gfx950${IMAGE_STAGE_SUFFIX}"
else
_from_line="FROM \$BASE_IMAGE_942${IMAGE_BASE_ARG_SUFFIX} AS gfx942${IMAGE_STAGE_SUFFIX}"
fi
case "${IMAGE_GFX}" in
gfx950) _from_line="FROM \$BASE_IMAGE_950${IMAGE_BASE_ARG_SUFFIX} AS gfx950${IMAGE_STAGE_SUFFIX}" ;;
gfx1250) _from_line="FROM \$BASE_IMAGE_1250${IMAGE_BASE_ARG_SUFFIX} AS gfx1250${IMAGE_STAGE_SUFFIX}" ;;
*) _from_line="FROM \$BASE_IMAGE_942${IMAGE_BASE_ARG_SUFFIX} AS gfx942${IMAGE_STAGE_SUFFIX}" ;;
esac
echo "[CI-AITER-CHECK] Using ${_from_line} from Dockerfile..."
REPO_AITER_COMMIT=$(grep -F -A20 "${_from_line}" docker/rocm.Dockerfile \
| grep 'AITER_COMMIT_DEFAULT=' \
@@ -0,0 +1,259 @@
"""MI45x DeepSeek-R1-0528-MXFP4 GSM8K Completion Evaluation Test (1-GPU)
Tests amd/DeepSeek-R1-0528-MXFP4 quantized model with triton attention
using few-shot completion benchmark on MI45x.
Registry: nightly-amd-1-gpu-mi45x-deepseek-r1-0528-mxfp4 suite
"""
import ast
import os
import re
import time
import unittest
from dataclasses import dataclass
from typing import List, Optional, Tuple
import numpy as np
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
is_in_ci,
popen_launch_server,
write_github_step_summary,
)
from sglang.utils import download_and_cache_file, read_jsonl
# Register for AMD CI - MI45x DeepSeek-R1-0528-MXFP4 accuracy test (~60 min)
register_amd_ci(
est_time=3600,
suite="nightly-amd-1-gpu-mi45x-deepseek-r1-0528-mxfp4",
nightly=True,
)
INVALID = -9999999
@dataclass
class ModelConfig:
"""Configuration for a model to test."""
model_path: str
tp_size: int = 1
accuracy_threshold: float = 0.50
other_args: Optional[List[str]] = None
env_vars: Optional[dict] = None
timeout: Optional[int] = None
def __post_init__(self):
if self.other_args is None:
self.other_args = []
if self.env_vars is None:
self.env_vars = {}
MI45X_DEEPSEEK_R1_0528_MXFP4_MODELS = [
ModelConfig(
model_path="amd/DeepSeek-R1-0528-MXFP4",
tp_size=1,
accuracy_threshold=0.50,
timeout=3600,
other_args=[
"--host",
"0.0.0.0",
"--mem-fraction-static",
"0.90",
"--chunked-prefill-size",
"16384",
"--attention-backend",
"triton",
"--max-running-requests",
"32",
"--kv-cache-dtype",
"auto",
"--page-size",
"64",
"--trust-remote-code",
],
env_vars={
"HSA_ENABLE_COREDUMP": "0",
"HSA_COREDUMP_PATTERN": "/dev/null",
"AMD_COREDUMP": "0",
"ENABLE_CK": "0",
"SGLANG_USE_AITER": "1",
"AITER_FORCE_A8W4": "1",
"AITER_GROUPED_FORCE_SPLIT_K1": "1",
"SGLANG_MOE_SHUFFLE_GFX1250": "1",
"ROCM_QUICK_REDUCE_QUANTIZATION": "NONE",
"SGLANG_AITER_FP8_PREFILL_ATTN": "0",
"SGLANG_AITER_MLA_PERSIST": "0",
"SGLANG_INT4_WEIGHT": "0",
"SGLANG_MOE_PADDING": "1",
"SGLANG_SET_CPU_AFFINITY": "1",
"SGLANG_ROCM_FUSED_DECODE_MLA": "0",
"SGLANG_USE_ROCM700A": "1",
},
),
]
def get_one_example(lines, i, include_answer):
"""Format a single GSM8K example."""
ret = "Question: " + lines[i]["question"] + "\nAnswer:"
if include_answer:
ret += " " + lines[i]["answer"]
return ret
def get_few_shot_examples(lines, k):
"""Get k few-shot examples for prompting."""
ret = ""
for i in range(k):
ret += get_one_example(lines, i, True) + "\n\n"
return ret
def get_answer_value(answer_str):
"""Extract numerical answer from response."""
answer_str = answer_str.replace(",", "")
numbers = re.findall(r"\d+", answer_str)
if len(numbers) < 1:
return INVALID
try:
return ast.literal_eval(numbers[-1])
except SyntaxError:
return INVALID
def run_gsm8k_benchmark(
base_url: str,
num_questions: int = 200,
num_shots: int = 5,
parallel: int = 64,
) -> Tuple[float, float, float]:
"""Run GSM8K few-shot completion benchmark."""
import sglang as sgl
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl"
data_path = download_and_cache_file(url)
lines = list(read_jsonl(data_path))
few_shot_examples = get_few_shot_examples(lines, num_shots)
questions = []
labels = []
for i in range(len(lines[:num_questions])):
questions.append(get_one_example(lines, i, False))
labels.append(get_answer_value(lines[i]["answer"]))
assert all(l != INVALID for l in labels)
arguments = [{"question": q} for q in questions]
@sgl.function
def few_shot_gsm8k(s, question):
s += few_shot_examples + question
s += sgl.gen(
"answer", max_tokens=512, stop=["Question", "Assistant:", "<|separator|>"]
)
backend = RuntimeEndpoint(base_url)
sgl.set_default_backend(backend)
tic = time.perf_counter()
states = few_shot_gsm8k.run_batch(
arguments, temperature=0, num_threads=parallel, progress_bar=True
)
latency = time.perf_counter() - tic
preds = [get_answer_value(states[i]["answer"]) for i in range(len(states))]
acc = np.mean(np.array(preds) == np.array(labels))
invalid = np.mean(np.array(preds) == INVALID)
return float(acc), float(invalid), float(latency)
class TestDeepSeekR10528MXFP4EvalMI45x(unittest.TestCase):
"""DeepSeek-R1-0528-MXFP4 GSM8K Completion Evaluation Test for AMD MI45x."""
@classmethod
def setUpClass(cls):
cls.models = MI45X_DEEPSEEK_R1_0528_MXFP4_MODELS
cls.base_url = DEFAULT_URL_FOR_TEST
cls.num_questions = int(os.environ.get("GSM8K_NUM_QUESTIONS", "200"))
def test_deepseek_r1_0528_mxfp4_accuracy(self):
"""Test DeepSeek-R1-0528-MXFP4 with GSM8K completion benchmark."""
all_results = []
summary = "### DeepSeek-R1-0528-MXFP4 Models (MI45x)\n\n"
summary += "| Model | TP | Accuracy | Threshold | Status |\n"
summary += "| ----- | -- | -------- | --------- | ------ |\n"
for config in self.models:
with self.subTest(model=config.model_path):
print(f"\n{'='*60}")
print(f"Testing: {config.model_path}")
print(f"{'='*60}")
env = os.environ.copy()
for key, value in config.env_vars.items():
env[key] = value
other_args = list(config.other_args)
other_args.extend(["--tensor-parallel-size", str(config.tp_size)])
timeout = config.timeout or DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
try:
process = popen_launch_server(
model=config.model_path,
base_url=self.base_url,
timeout=timeout,
other_args=other_args,
env=env,
)
try:
acc, invalid, latency = run_gsm8k_benchmark(
self.base_url, num_questions=self.num_questions
)
passed = acc >= config.accuracy_threshold
status = "✅ PASS" if passed else "❌ FAIL"
print(
f" accuracy={acc:.3f} threshold={config.accuracy_threshold} {status}"
)
all_results.append(
{
"model": config.model_path,
"accuracy": acc,
"passed": passed,
}
)
summary += f"| {config.model_path} | {config.tp_size} | {acc:.3f} | {config.accuracy_threshold} | {status} |\n"
finally:
kill_process_tree(process.pid)
except Exception as e:
summary += f"| {config.model_path} | {config.tp_size} | N/A | {config.accuracy_threshold} | ❌ ERROR |\n"
all_results.append(
{
"model": config.model_path,
"accuracy": None,
"passed": False,
"error": str(e),
}
)
if is_in_ci():
write_github_step_summary(summary)
failed = [r for r in all_results if not r["passed"]]
if failed:
raise AssertionError(f"Failed models: {[r['model'] for r in failed]}")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,262 @@
"""MI45x DeepSeek-V4-Flash GSM8K Completion Evaluation Test (1-GPU)
Tests deepseek-ai/DeepSeek-V4-Flash with DSV4 attention backend
using few-shot completion benchmark on MI45x.
Registry: nightly-amd-1-gpu-mi45x-deepseek-v4-flash suite
"""
import ast
import os
import re
import time
import unittest
from dataclasses import dataclass
from typing import List, Optional, Tuple
import numpy as np
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
is_in_ci,
popen_launch_server,
write_github_step_summary,
)
from sglang.utils import download_and_cache_file, read_jsonl
# Register for AMD CI - MI45x DeepSeek-V4-Flash accuracy test (~60 min)
register_amd_ci(
est_time=3600,
suite="nightly-amd-1-gpu-mi45x-deepseek-v4-flash",
nightly=True,
)
INVALID = -9999999
@dataclass
class ModelConfig:
"""Configuration for a model to test."""
model_path: str
tp_size: int = 1
accuracy_threshold: float = 0.50
other_args: Optional[List[str]] = None
env_vars: Optional[dict] = None
timeout: Optional[int] = None
def __post_init__(self):
if self.other_args is None:
self.other_args = []
if self.env_vars is None:
self.env_vars = {}
MI45X_DEEPSEEK_V4_FLASH_MODELS = [
ModelConfig(
model_path="deepseek-ai/DeepSeek-V4-Flash",
tp_size=1,
accuracy_threshold=0.50,
timeout=3600,
other_args=[
"--attention-backend",
"dsv4",
"--page-size",
"256",
"--mem-fraction-static",
"0.60",
"--swa-full-tokens-ratio",
"0.15",
"--disable-shared-experts-fusion",
"--tool-call-parser",
"deepseekv4",
"--reasoning-parser",
"deepseek-v4",
"--chunked-prefill-size",
"8192",
"--cuda-graph-max-bs",
"256",
"--max-running-requests",
"256",
"--disable-radix-cache",
"--kv-cache-dtype",
"fp8_e4m3",
"--trust-remote-code",
],
env_vars={
"SGLANG_DEFAULT_THINKING": "1",
"SGLANG_DSV4_REASONING_EFFORT": "max",
"SGLANG_USE_ROCM700A": "0",
"SGLANG_HACK_FLASHMLA_BACKEND": "unified_kv_triton",
"AITER_BF16_FP8_MOE_BOUND": "0",
"AITER_FORCE_A8W4": "1",
"SGLANG_USE_AITER_MOE_GU_ITLV": "0",
"SGLANG_OPT_FUSE_MHC_POST_PRE": "0",
"ENABLE_CK": "0",
"SGLANG_USE_AITER": "1",
"AITER_GROUPED_FORCE_SPLIT_K1": "1",
},
),
]
def get_one_example(lines, i, include_answer):
"""Format a single GSM8K example."""
ret = "Question: " + lines[i]["question"] + "\nAnswer:"
if include_answer:
ret += " " + lines[i]["answer"]
return ret
def get_few_shot_examples(lines, k):
"""Get k few-shot examples for prompting."""
ret = ""
for i in range(k):
ret += get_one_example(lines, i, True) + "\n\n"
return ret
def get_answer_value(answer_str):
"""Extract numerical answer from response."""
answer_str = answer_str.replace(",", "")
numbers = re.findall(r"\d+", answer_str)
if len(numbers) < 1:
return INVALID
try:
return ast.literal_eval(numbers[-1])
except SyntaxError:
return INVALID
def run_gsm8k_benchmark(
base_url: str,
num_questions: int = 200,
num_shots: int = 5,
parallel: int = 64,
) -> Tuple[float, float, float]:
"""Run GSM8K few-shot completion benchmark."""
import sglang as sgl
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl"
data_path = download_and_cache_file(url)
lines = list(read_jsonl(data_path))
few_shot_examples = get_few_shot_examples(lines, num_shots)
questions = []
labels = []
for i in range(len(lines[:num_questions])):
questions.append(get_one_example(lines, i, False))
labels.append(get_answer_value(lines[i]["answer"]))
assert all(l != INVALID for l in labels)
arguments = [{"question": q} for q in questions]
@sgl.function
def few_shot_gsm8k(s, question):
s += few_shot_examples + question
s += sgl.gen(
"answer", max_tokens=512, stop=["Question", "Assistant:", "<|separator|>"]
)
backend = RuntimeEndpoint(base_url)
sgl.set_default_backend(backend)
tic = time.perf_counter()
states = few_shot_gsm8k.run_batch(
arguments, temperature=0, num_threads=parallel, progress_bar=True
)
latency = time.perf_counter() - tic
preds = [get_answer_value(states[i]["answer"]) for i in range(len(states))]
acc = np.mean(np.array(preds) == np.array(labels))
invalid = np.mean(np.array(preds) == INVALID)
return float(acc), float(invalid), float(latency)
class TestDeepSeekV4FlashEvalMI45x(unittest.TestCase):
"""DeepSeek-V4-Flash GSM8K Completion Evaluation Test for AMD MI45x."""
@classmethod
def setUpClass(cls):
cls.models = MI45X_DEEPSEEK_V4_FLASH_MODELS
cls.base_url = DEFAULT_URL_FOR_TEST
cls.num_questions = int(os.environ.get("GSM8K_NUM_QUESTIONS", "200"))
def test_deepseek_v4_flash_accuracy(self):
"""Test DeepSeek-V4-Flash with GSM8K completion benchmark."""
all_results = []
summary = "### DeepSeek-V4-Flash Models (MI45x)\n\n"
summary += "| Model | TP | Accuracy | Threshold | Status |\n"
summary += "| ----- | -- | -------- | --------- | ------ |\n"
for config in self.models:
with self.subTest(model=config.model_path):
print(f"\n{'='*60}")
print(f"Testing: {config.model_path}")
print(f"{'='*60}")
env = os.environ.copy()
for key, value in config.env_vars.items():
env[key] = value
other_args = list(config.other_args)
other_args.extend(["--tp", str(config.tp_size)])
timeout = config.timeout or DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
try:
process = popen_launch_server(
model=config.model_path,
base_url=self.base_url,
timeout=timeout,
other_args=other_args,
env=env,
)
try:
acc, invalid, latency = run_gsm8k_benchmark(
self.base_url, num_questions=self.num_questions
)
passed = acc >= config.accuracy_threshold
status = "✅ PASS" if passed else "❌ FAIL"
print(
f" accuracy={acc:.3f} threshold={config.accuracy_threshold} {status}"
)
all_results.append(
{
"model": config.model_path,
"accuracy": acc,
"passed": passed,
}
)
summary += f"| {config.model_path} | {config.tp_size} | {acc:.3f} | {config.accuracy_threshold} | {status} |\n"
finally:
kill_process_tree(process.pid)
except Exception as e:
summary += f"| {config.model_path} | {config.tp_size} | N/A | {config.accuracy_threshold} | ❌ ERROR |\n"
all_results.append(
{
"model": config.model_path,
"accuracy": None,
"passed": False,
"error": str(e),
}
)
if is_in_ci():
write_github_step_summary(summary)
failed = [r for r in all_results if not r["passed"]]
if failed:
raise AssertionError(f"Failed models: {[r['model'] for r in failed]}")
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,248 @@
"""MI45x GPT-OSS W4A8 MXFP4-FP8 GSM8K Completion Evaluation Test (1-GPU)
Tests the AMD Quark `gpt-oss-120b-w-mxfp4-a-fp8` checkpoint (MXFP4
weights + static per-tensor FP8 activations) using few-shot completion
benchmark on MI45x.
Registry: nightly-amd-1-gpu-mi45x suite
"""
import ast
import os
import re
import time
import unittest
from dataclasses import dataclass
from typing import List, Optional, Tuple
import numpy as np
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
is_in_ci,
popen_launch_server,
write_github_step_summary,
)
from sglang.utils import download_and_cache_file, read_jsonl
# Register for AMD CI - MI45x GPT-OSS W4A8 MXFP4-FP8 accuracy test (~15 min)
register_amd_ci(est_time=900, suite="nightly-amd-1-gpu-mi45x", nightly=True)
INVALID = -9999999
@dataclass
class ModelConfig:
"""Configuration for a model to test."""
model_path: str
tp_size: int = 1
accuracy_threshold: float = 0.50
other_args: Optional[List[str]] = None
env_vars: Optional[dict] = None
timeout: Optional[int] = None
def __post_init__(self):
if self.other_args is None:
self.other_args = []
if self.env_vars is None:
self.env_vars = {}
# GPT-OSS Quark W4A8 MXFP4-FP8 models for MI45x
MI45X_GPT_OSS_W4A8_MXFP4_MODELS = [
ModelConfig(
model_path="amd/gpt-oss-120b-w-mxfp4-a-fp8",
tp_size=1,
accuracy_threshold=0.79,
timeout=900,
other_args=[
"--prefill-attention-backend",
"triton",
"--decode-attention-backend",
"aiter",
"--max-running-requests",
"128",
"--mem-fraction-static",
"0.9",
"--disable-radix-cache",
"--page-size",
"64",
"--trust-remote-code",
],
# AITER MXFP4 fused-MoE for gpt-oss uses the SEPARATED gate/up tile
# layout (matches `gptoss_fp4_tuned_fmoe.csv` flydsl entries and the
# Mxfp4MoEMethod weight shuffle). Other AITER MXFP4 callers default
# to INTERLEAVE, so opt out explicitly here.
env_vars={
"AITER_FORCE_A8W4": "1",
"ENABLE_CK": "0",
"HSA_COREDUMP_PATTERN": "/dev/null",
"SGLANG_USE_AITER": "1",
"SGLANG_USE_AITER_MOE_GU_ITLV": "1",
"SGLANG_USE_AITER_UNIFIED_ATTN": "1",
},
),
]
def get_one_example(lines, i, include_answer):
"""Format a single GSM8K example."""
ret = "Question: " + lines[i]["question"] + "\nAnswer:"
if include_answer:
ret += " " + lines[i]["answer"]
return ret
def get_few_shot_examples(lines, k):
"""Get k few-shot examples for prompting."""
ret = ""
for i in range(k):
ret += get_one_example(lines, i, True) + "\n\n"
return ret
def get_answer_value(answer_str):
"""Extract numerical answer from response."""
answer_str = answer_str.replace(",", "")
numbers = re.findall(r"\d+", answer_str)
if len(numbers) < 1:
return INVALID
try:
return ast.literal_eval(numbers[-1])
except SyntaxError:
return INVALID
def run_gsm8k_benchmark(
base_url: str,
num_questions: int = 200,
num_shots: int = 5,
parallel: int = 64,
) -> Tuple[float, float, float]:
"""Run GSM8K few-shot completion benchmark."""
import sglang as sgl
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl"
data_path = download_and_cache_file(url)
lines = list(read_jsonl(data_path))
few_shot_examples = get_few_shot_examples(lines, num_shots)
questions = []
labels = []
for i in range(len(lines[:num_questions])):
questions.append(get_one_example(lines, i, False))
labels.append(get_answer_value(lines[i]["answer"]))
assert all(l != INVALID for l in labels)
arguments = [{"question": q} for q in questions]
@sgl.function
def few_shot_gsm8k(s, question):
s += few_shot_examples + question
s += sgl.gen(
"answer", max_tokens=512, stop=["Question", "Assistant:", "<|separator|>"]
)
backend = RuntimeEndpoint(base_url)
sgl.set_default_backend(backend)
tic = time.perf_counter()
states = few_shot_gsm8k.run_batch(
arguments, temperature=0, num_threads=parallel, progress_bar=True
)
latency = time.perf_counter() - tic
preds = [get_answer_value(states[i]["answer"]) for i in range(len(states))]
acc = np.mean(np.array(preds) == np.array(labels))
invalid = np.mean(np.array(preds) == INVALID)
return float(acc), float(invalid), float(latency)
class TestGptOssW4A8Mxfp4EvalMI45x(unittest.TestCase):
"""GPT-OSS W4A8 MXFP4-FP8 GSM8K Completion Evaluation Test for AMD MI45x."""
@classmethod
def setUpClass(cls):
cls.models = MI45X_GPT_OSS_W4A8_MXFP4_MODELS
cls.base_url = DEFAULT_URL_FOR_TEST
cls.num_questions = int(os.environ.get("GSM8K_NUM_QUESTIONS", "200"))
def test_gpt_oss_accuracy(self):
"""Test GPT-OSS W4A8 MXFP4-FP8 models with GSM8K completion benchmark."""
all_results = []
summary = "### GPT-OSS W4A8 MXFP4-FP8 Models (MI45x)\n\n"
summary += "| Model | TP | Accuracy | Threshold | Status |\n"
summary += "| ----- | -- | -------- | --------- | ------ |\n"
for config in self.models:
with self.subTest(model=config.model_path):
print(f"\n{'='*60}")
print(f"Testing: {config.model_path}")
print(f"{'='*60}")
env = os.environ.copy()
for key, value in config.env_vars.items():
env[key] = value
other_args = list(config.other_args)
other_args.extend(["--tp", str(config.tp_size)])
timeout = config.timeout or DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
try:
process = popen_launch_server(
model=config.model_path,
base_url=self.base_url,
timeout=timeout,
other_args=other_args,
env=env,
)
try:
acc, invalid, latency = run_gsm8k_benchmark(
self.base_url, num_questions=self.num_questions
)
passed = acc >= config.accuracy_threshold
status = "✅ PASS" if passed else "❌ FAIL"
print(
f" accuracy={acc:.3f} threshold={config.accuracy_threshold} {status}"
)
all_results.append(
{
"model": config.model_path,
"accuracy": acc,
"passed": passed,
}
)
summary += f"| {config.model_path} | {config.tp_size} | {acc:.3f} | {config.accuracy_threshold} | {status} |\n"
finally:
kill_process_tree(process.pid)
except Exception as e:
summary += f"| {config.model_path} | {config.tp_size} | N/A | {config.accuracy_threshold} | ❌ ERROR |\n"
all_results.append(
{
"model": config.model_path,
"accuracy": None,
"passed": False,
"error": str(e),
}
)
if is_in_ci():
write_github_step_summary(summary)
failed = [r for r in all_results if not r["passed"]]
if failed:
raise AssertionError(f"Failed models: {[r['model'] for r in failed]}")
if __name__ == "__main__":
unittest.main()