[New Model][ROCm] Add AMD support for DeepSeek V4 (#40871)

Signed-off-by: ganyi <[email protected]>
Signed-off-by: whx-sjtu <[email protected]>
Signed-off-by: tjtanaa <[email protected]>
Signed-off-by: tjtanaavllm <[email protected]>
Co-authored-by: ganyi <[email protected]>
Co-authored-by: tjtanaa <[email protected]>
Co-authored-by: tjtanaavllm <[email protected]>
This commit is contained in:
Hexiang Wang
2026-05-05 08:55:37 -07:00
committed by GitHub
co-authored by ganyi tjtanaa tjtanaavllm
parent 2228fe6868
commit 628c436301
22 changed files with 939 additions and 134 deletions
+6 -6
View File
@@ -307,12 +307,12 @@ set(VLLM_EXT_SRC
"csrc/quantization/activation_kernels.cu"
"csrc/cuda_utils_kernels.cu"
"csrc/custom_all_reduce.cu"
"csrc/torch_bindings.cpp")
"csrc/torch_bindings.cpp"
"csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu")
if(VLLM_GPU_LANG STREQUAL "CUDA")
list(APPEND VLLM_EXT_SRC
"csrc/minimax_reduce_rms_kernel.cu"
"csrc/fused_deepseek_v4_qnorm_rope_kv_insert_kernel.cu")
"csrc/minimax_reduce_rms_kernel.cu")
SET(CUTLASS_ENABLE_HEADERS_ONLY ON CACHE BOOL "Enable only the header library")
@@ -1047,13 +1047,13 @@ endif()
set(VLLM_MOE_EXT_SRC
"csrc/moe/torch_bindings.cpp"
"csrc/moe/moe_align_sum_kernels.cu"
"csrc/moe/topk_softmax_kernels.cu")
"csrc/moe/topk_softmax_kernels.cu"
"csrc/moe/topk_softplus_sqrt_kernels.cu")
if(VLLM_GPU_LANG STREQUAL "CUDA")
list(APPEND VLLM_MOE_EXT_SRC
"csrc/moe/moe_wna16.cu"
"csrc/moe/grouped_topk_kernels.cu"
"csrc/moe/topk_softplus_sqrt_kernels.cu")
"csrc/moe/grouped_topk_kernels.cu")
endif()
if(VLLM_GPU_LANG STREQUAL "CUDA")
@@ -29,7 +29,11 @@
*/
#include <cmath>
#include <cuda_fp8.h>
#ifndef USE_ROCM
#include <cuda_fp8.h>
#else
#include <hip/hip_fp8.h>
#endif
#include <cuda_runtime.h>
#include <type_traits>
@@ -42,7 +46,23 @@
#include "type_convert.cuh"
#ifndef FINAL_MASK
#define FINAL_MASK 0xffffffffu
#ifdef USE_ROCM
#define FINAL_MASK 0xffffffffffffffffULL
#else
#define FINAL_MASK 0xffffffffu
#endif
#endif
#ifdef USE_ROCM
// ROCm-compatible FP8 conversion helpers
__device__ __forceinline__ uint8_t rocm_cvt_float_to_fp8_e4m3(float val) {
#if defined(HIP_FP8_TYPE_OCP)
__hip_fp8_e4m3 fp8_val(val);
#else
__hip_fp8_e4m3_fnuz fp8_val(val);
#endif
return reinterpret_cast<uint8_t&>(fp8_val);
}
#endif
namespace vllm {
@@ -314,9 +334,13 @@ __global__ void fusedDeepseekV4QNormRopeKVRopeQuantInsertKernel(
for (int i = 0; i < kElemsPerLane; i++) {
float scaled = elements[i] * inv_scale;
scaled = fminf(fmaxf(scaled, -kFp8Max), kFp8Max);
#ifndef USE_ROCM
__nv_fp8_storage_t s =
__nv_cvt_float_to_fp8(scaled, __NV_SATFINITE, __NV_E4M3);
out_bytes[i] = static_cast<uint8_t>(s);
#else
out_bytes[i] = rocm_cvt_float_to_fp8_e4m3(scaled);
#endif
}
// One 16-byte STG per lane.
*reinterpret_cast<uint4*>(token_fp8_ptr + dim_base) =
@@ -384,6 +408,7 @@ void launchFusedDeepseekV4QNormRopeKVRopeQuantInsert(
// PDL: enable programmatic stream serialization whenever the hardware
// supports it (SM90+). On pre-Hopper GPUs the attribute is unavailable,
// so leave numAttrs = 0 and launch as a regular kernel.
#ifndef USE_ROCM
static int const sm_version = getSMVersion();
// Host-side guard: the device kernel body is compiled as a no-op for
// bf16 on pre-Ampere (sm_70/sm_75) because _typeConvert<BFloat16> is
@@ -410,6 +435,15 @@ void launchFusedDeepseekV4QNormRopeKVRopeQuantInsert(
q_inout, kv_in, k_cache, slot_mapping, position_ids, cos_sin_cache, eps,
num_tokens_full, num_tokens_insert, num_heads_q, cache_block_size,
kv_block_stride);
#else
// ROCm: use standard kernel launch syntax (no PDL/stream serialization)
// clang-format off
fusedDeepseekV4QNormRopeKVRopeQuantInsertKernel<scalar_t_in>
<<<grid, kBlockSize, 0, stream>>>(
q_inout, kv_in, k_cache, slot_mapping, position_ids, cos_sin_cache,
eps, num_tokens_full, num_tokens_insert, num_heads_q,
cache_block_size, kv_block_stride);
#endif
}
} // namespace deepseek_v4_fused_ops
+32 -21
View File
@@ -60,15 +60,6 @@ __device__ __forceinline__ float toFloat(T value) {
}
}
#define FINAL_MASK 0xffffffff
template <typename T>
__inline__ __device__ T warpReduceSum(T val) {
#pragma unroll
for (int mask = 16; mask > 0; mask >>= 1)
val += __shfl_xor_sync(FINAL_MASK, val, mask, 32);
return val;
}
// ====================== TopK softplus_sqrt things
// ===============================
@@ -272,8 +263,14 @@ __launch_bounds__(WARPS_PER_CTA* WARP_SIZE_PARAM) __global__
}
}
// Compute per-thread scale (using warp reduction when renormalizing).
// THREADS_PER_ROW-parameterized butterfly works for both warp sizes (32
// on CUDA, 64 on ROCm CDNA) and any THREADS_PER_ROW the dispatch picks.
if (renormalize) {
selected_sum = warpReduceSum(selected_sum);
#pragma unroll
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
selected_sum +=
VLLM_SHFL_XOR_SYNC_WIDTH(selected_sum, mask, THREADS_PER_ROW);
}
}
float scale = static_cast<float>(routed_scaling_factor);
if (renormalize) {
@@ -544,7 +541,6 @@ void topkGatingSoftplusSqrtKernelLauncher(
const IndType* tid2eid, cudaStream_t stream) {
static constexpr int WARPS_PER_TB = 4;
static constexpr int BYTES_PER_LDG_POWER_OF_2 = 16;
#ifndef USE_ROCM
// for bfloat16 dtype, we need 4 bytes loading to make sure num_experts
// elements can be loaded by a warp
static constexpr int BYTES_PER_LDG_MULTIPLE_64 =
@@ -552,6 +548,19 @@ void topkGatingSoftplusSqrtKernelLauncher(
std::is_same_v<InputType, __half>)
? 4
: 8;
// Narrower LDG (ELTS_PER_LDG=1) used by 192/320/448/576 on ROCm WARP_SIZE=64
// where ELTS_PER_LDG=2 fails the EXPERTS%(ELTS_PER_LDG*WARP_SIZE)==0 check.
// On CUDA WARP_SIZE=32 the wider LDG already aligns, so the alias collapses
// back to BYTES_PER_LDG_MULTIPLE_64 — no behavioral change for CUDA.
#ifdef USE_ROCM
static constexpr int BYTES_PER_LDG_MULTIPLE_64_NARROW =
(std::is_same_v<InputType, __nv_bfloat16> ||
std::is_same_v<InputType, __half>)
? 2
: 4;
#else
static constexpr int BYTES_PER_LDG_MULTIPLE_64_NARROW =
BYTES_PER_LDG_MULTIPLE_64;
#endif
switch (num_experts) {
case 1:
@@ -584,27 +593,29 @@ void topkGatingSoftplusSqrtKernelLauncher(
case 512:
LAUNCH_SOFTPLUS_SQRT(512, WARPS_PER_TB, BYTES_PER_LDG_POWER_OF_2);
break;
// (CUDA only) support multiples of 64 when num_experts is not power of 2.
// ROCm uses WARP_SIZE 64 so 8 bytes loading won't fit for some of
// num_experts, alternatively we can test 4 bytes loading and enable it in
// future.
#ifndef USE_ROCM
// Multiples of 64 that are not powers of 2. The kernel requires
// EXPERTS % (ELTS_PER_LDG * WARP_SIZE) == 0. With ELTS_PER_LDG=2
// (BYTES_PER_LDG_MULTIPLE_64), this holds for all five values on CUDA
// WARP_SIZE=32 but only for 384 on ROCm WARP_SIZE=64. The other four
// use BYTES_PER_LDG_MULTIPLE_64_NARROW (ELTS_PER_LDG=1), which
// satisfies the assertion for any multiple of 64 on either backend;
// on CUDA the narrow alias collapses back to the wider load, so CUDA
// behavior is unchanged.
case 192:
LAUNCH_SOFTPLUS_SQRT(192, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64);
LAUNCH_SOFTPLUS_SQRT(192, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW);
break;
case 320:
LAUNCH_SOFTPLUS_SQRT(320, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64);
LAUNCH_SOFTPLUS_SQRT(320, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW);
break;
case 384:
LAUNCH_SOFTPLUS_SQRT(384, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64);
break;
case 448:
LAUNCH_SOFTPLUS_SQRT(448, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64);
LAUNCH_SOFTPLUS_SQRT(448, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW);
break;
case 576:
LAUNCH_SOFTPLUS_SQRT(576, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64);
LAUNCH_SOFTPLUS_SQRT(576, WARPS_PER_TB, BYTES_PER_LDG_MULTIPLE_64_NARROW);
break;
#endif
default: {
TORCH_CHECK(false, "Unsupported expert number: ", num_experts);
}
+1 -2
View File
@@ -16,14 +16,13 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, m) {
"bias) -> ()");
m.impl("topk_sigmoid", torch::kCUDA, &topk_sigmoid);
#ifndef USE_ROCM
m.def(
"topk_softplus_sqrt(Tensor! topk_weights, Tensor! topk_indices, Tensor! "
"token_expert_indices, Tensor gating_output, bool renormalize, float "
"routed_scaling_factor, Tensor? "
"bias, Tensor? input_ids, Tensor? tid2eid) -> ()");
m.impl("topk_softplus_sqrt", torch::kCUDA, &topk_softplus_sqrt);
#endif
// Calculate the result of moe by summing up the partial results
// from all selected experts.
m.def("moe_sum(Tensor input, Tensor! output) -> ()");
-2
View File
@@ -183,7 +183,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
"int forced_token_heads_per_warp=-1) -> ()");
ops.impl("fused_qk_norm_rope", torch::kCUDA, &fused_qk_norm_rope);
#ifndef USE_ROCM
// Horizontally-fused DeepseekV4-MLA: per-head RMSNorm + GPT-J RoPE for Q, and
// GPT-J RoPE + UE8M0 FP8 quant + paged cache insert for KV, all in one
// kernel launch.
@@ -194,7 +193,6 @@ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
"float eps, int cache_block_size) -> ()");
ops.impl("fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert", torch::kCUDA,
&fused_deepseek_v4_qnorm_rope_kv_rope_quant_insert);
#endif
// Apply repetition penalties to logits in-place
ops.def(
+3
View File
@@ -21,3 +21,6 @@ timm>=1.0.17
# amd-quark: required for Quark quantization on ROCm
# To be consistent with test_quark.py
amd-quark>=0.8.99
# tilelang has to be installed for mhc module to be
# imported correctly.
tilelang==0.1.9
+4 -2
View File
@@ -70,7 +70,8 @@ def test_sqrtsoftplus_bias_uses_deepseek_v4_routing_method():
@pytest.mark.skipif(
not current_platform.is_cuda(), reason="This test is skipped on non-CUDA platform."
not current_platform.is_cuda_alike(),
reason="This test is skipped on non-CUDA platform.",
)
@pytest.mark.parametrize("num_tokens", [1, 33, 128])
@pytest.mark.parametrize("hidden_size", [1024, 2048])
@@ -125,7 +126,8 @@ def test_fused_topk_softplus_sqrt(
@pytest.mark.skipif(
not current_platform.is_cuda(), reason="This test is skipped on non-CUDA platform."
not current_platform.is_cuda_alike(),
reason="This test is skipped on non-CUDA platform.",
)
@pytest.mark.parametrize("num_tokens", [1, 33, 128])
@pytest.mark.parametrize("hidden_size", [1024, 2048])
+2
View File
@@ -119,6 +119,7 @@ MoEBackend = Literal[
"flashinfer_cutedsl",
"marlin",
"humming",
"triton_unfused",
"aiter",
"emulation",
]
@@ -150,6 +151,7 @@ class KernelConfig:
- "flashinfer_cutedsl": Use FlashInfer with CuteDSL kernels (FP4 only)
- "marlin": Use Marlin kernels (weight-only quantization)
- "humming": Use Humming Mixed Precision kernels
- "triton_unfused": Use Triton unfused MoE kernels
- "aiter": Use AMD AITer kernels (ROCm only)
- "emulation": use BF16/FP16 GEMM, dequantizing weights and
running QDQ on activations.
@@ -312,6 +312,21 @@ class AiterFp8BlockScaledMMKernel(Fp8BlockScaledMMLinearKernel):
As: torch.Tensor,
Bs: torch.Tensor,
) -> torch.Tensor:
if As.dtype != Bs.dtype:
from vllm.model_executor.layers.quantization.utils.fp8_utils import (
_upcast_e8m0_to_fp32,
)
if As.dtype == torch.float8_e8m0fnu:
As = _upcast_e8m0_to_fp32(As).contiguous()
else:
As = As.to(torch.float32)
if Bs.dtype == torch.float8_e8m0fnu:
Bs = _upcast_e8m0_to_fp32(Bs).contiguous()
else:
Bs = Bs.to(torch.float32)
out_dtype = self.config.out_dtype
if self.use_triton:
gemm_a8w8_blockscale_op = rocm_aiter_ops.triton_gemm_a8w8_blockscale
+3 -1
View File
@@ -169,7 +169,9 @@ class SiluAndMulWithClamp(CustomOp):
def __init__(self, swiglu_limit: float, *, compile_native: bool = True):
super().__init__(compile_native=compile_native)
self.swiglu_limit = float(swiglu_limit)
if current_platform.is_cuda_alike() or current_platform.is_xpu():
if current_platform.is_rocm():
self._forward_method = self.forward_native
elif current_platform.is_cuda_alike() or current_platform.is_xpu():
self.op = torch.ops._C.silu_and_mul_with_clamp
elif current_platform.is_cpu():
self._forward_method = self.forward_native
@@ -300,6 +300,7 @@ class DeepseekCompressor(nn.Module):
state_cache = self.state_cache.kv_cache
# kv_state stored in first half, score_state stored in second half
state_width = state_cache.shape[-1] // 2
pdl_kwargs = {} if current_platform.is_rocm() else {"launch_pdl": False}
# Store the KV and score (with fused APE addition) in the state.
# NOTE: PDL is disabled — both this kernel and _fused_kernel below
@@ -324,7 +325,7 @@ class DeepseekCompressor(nn.Module):
TRITON_BLOCK_SIZE=triton.next_power_of_2(kv.shape[-1]),
STATE_WIDTH=state_width,
COMPRESS_RATIO=self.compress_ratio,
launch_pdl=False,
**pdl_kwargs,
)
# Fused: compress → RMSNorm → RoPE → FP8 quant → KV cache write.
@@ -373,7 +374,7 @@ class DeepseekCompressor(nn.Module):
SCALE_DIM=self._scale_dim,
KV_BLOCK_STRIDE=kv_cache.stride(0),
num_warps=self._num_warps,
launch_pdl=False,
**pdl_kwargs,
)
@@ -28,6 +28,11 @@ from vllm.v1.attention.ops.deepseek_v4_ops import (
fused_inv_rope_fp8_quant,
fused_q_kv_rmsnorm,
)
from vllm.v1.attention.ops.rocm_aiter_mla_sparse import (
rocm_forward_decode_fallback,
rocm_inv_rope_einsum,
rocm_sparse_attn_prefill,
)
if TYPE_CHECKING:
from vllm.v1.attention.backends.mla.sparse_swa import (
@@ -53,6 +58,7 @@ from vllm.model_executor.layers.quantization.input_quant_fp8 import (
from vllm.model_executor.layers.quantization.utils.quant_utils import (
GroupShape,
)
from vllm.platforms import current_platform
from vllm.utils.multi_stream_utils import (
execute_in_parallel,
maybe_execute_in_parallel,
@@ -198,8 +204,6 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
# Pick fp8_einsum recipe based on GPU arch:
# SM90: FP32 block scales stay [g, r/128, d/128] → sfb_gran_mn=128
# SM100: INT32 packed scales become [g, r, ...] → sfb_gran_mn=1
from vllm.platforms import current_platform
cap = current_platform.get_device_capability()
assert cap is not None, "DeepseekV4 attention requires a CUDA device"
self._einsum_recipe = (1, 128, 128) if cap.major <= 9 else (1, 1, 128)
@@ -222,6 +226,7 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
+ 1 # 1B pad
)
# Will be None on ROCm for now.
self.aux_stream_list = mla_modules.aux_stream_list
# [0]: GEMM start / post-GEMM event0. [1..3]: GEMM done events;
# [1] doubles as post-GEMM event1. Reuse is safe: GEMM fully joins
@@ -303,6 +308,19 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
)
o = o_padded[:, : self.n_local_heads, :]
# Keep ROCm on the BF16 reference wo_a path util kernel ready.
if current_platform.is_rocm():
z = rocm_inv_rope_einsum(
self.rotary_emb,
o,
positions,
self.rope_head_dim,
self.n_local_groups,
self.o_lora_rank,
self.wo_a,
)
return self.wo_b(z.flatten(1))
# O projection: inverse RoPE + FP8 quant + einsum + wo_b
o_fp8, o_scale = fused_inv_rope_fp8_quant(
o,
@@ -336,12 +354,15 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
return self.wo_b(z.flatten(1))
def attn_gemm_parallel_execute(self, hidden_states) -> tuple[Any, ...]:
assert self.aux_stream_list is not None
assert len(self.aux_stream_list) >= 3
aux_streams = self.aux_stream_list
if aux_streams is not None:
assert len(aux_streams) >= 3
aux_streams = aux_streams[:3]
# fused_wqa_wkv (heaviest) on default; the three lighter input GEMMs
# on aux streams 0..2 when their owning module exists. ln_events[0]
# is the fan-out start event; ln_events[1..3] are per-aux done events.
# On ROCm, aux_streams is None and execute_in_parallel runs serially.
aux_fns: list[Callable[[], Any] | None] = [None, None, None]
if self.compressor is not None:
@@ -385,7 +406,7 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
aux_fns,
self.ln_events[0],
self.ln_events[1:4],
self.aux_stream_list[:3],
aux_streams,
enable=hidden_states.shape[0]
<= envs.VLLM_MULTI_STREAM_GEMM_TOKEN_THRESHOLD,
)
@@ -419,8 +440,9 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
# downstream reads q on default). Indexer/compressor go on aux for
# overlap with default's GEMM + cache write.
if self.indexer is not None:
assert self.aux_stream_list is not None
aux_stream = self.aux_stream_list[0]
aux_stream = (
self.aux_stream_list[0] if self.aux_stream_list is not None else None
)
indexer = self.indexer
# Local ref so the closure keeps a non-None type for mypy.
assert self.compressor is not None
@@ -448,8 +470,9 @@ class DeepseekV4MultiHeadLatentAttentionWrapper(PluggableLayer):
)
elif self.compressor is not None:
# wq_b + kv_insert on default, compressor on aux.
assert self.aux_stream_list is not None
aux_stream = self.aux_stream_list[0]
aux_stream = (
self.aux_stream_list[0] if self.aux_stream_list is not None else None
)
compressor = self.compressor
def wq_b_kv_insert() -> torch.Tensor:
@@ -668,7 +691,7 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase):
vllm_config.scheduler_config.max_num_batched_tokens
)
self.max_model_len = vllm_config.model_config.max_model_len
# DeepseekV4 only supports fp8 kv-cache format for now
# DeepseekV4 only supports fp8 kv-cache format for now.
kv_cache_dtype = cache_config.cache_dtype if cache_config is not None else "fp8"
assert kv_cache_dtype.startswith("fp8"), (
@@ -816,6 +839,25 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase):
swa_indices = swa_metadata.decode_swa_indices
swa_lens = swa_metadata.decode_swa_lens
if current_platform.is_rocm():
rocm_forward_decode_fallback(
q=q,
kv_cache=kv_cache,
swa_k_cache=self.swa_cache_layer.kv_cache,
swa_only=swa_only,
topk_indices=topk_indices,
topk_lens=topk_lens,
swa_indices=swa_indices,
swa_lens=swa_lens,
attn_sink=self.attn_sink,
scale=self.scale,
head_dim=self.head_dim,
nope_head_dim=self.nope_head_dim,
rope_head_dim=self.rope_head_dim,
output=output,
)
return
# We treat queries in the same seq as different queries
# and later we only attend by generated indices.
# q arrives pre-padded to self.padded_heads by the outer wrapper.
@@ -980,15 +1022,27 @@ class DeepseekV4MLAAttention(nn.Module, AttentionLayerBase):
N,
)
output_chunk, _, _ = flash_mla_sparse_fwd(
q=q[query_start:query_end],
kv=kv.view(-1, 1, q.shape[-1]),
indices=combined_indices.unsqueeze(1),
sm_scale=self.scale,
attn_sink=self.attn_sink,
topk_length=combined_lens,
out=output[query_start:query_end],
)
if current_platform.is_rocm():
rocm_sparse_attn_prefill(
q=q[query_start:query_end],
kv=kv.view(-1, 1, q.shape[-1]),
indices=combined_indices.unsqueeze(1),
topk_length=combined_lens,
scale=self.scale,
head_dim=self.head_dim,
attn_sink=self.attn_sink,
output=output[query_start:query_end],
)
else:
output_chunk, _, _ = flash_mla_sparse_fwd(
q=q[query_start:query_end],
kv=kv.view(-1, 1, q.shape[-1]),
indices=combined_indices.unsqueeze(1),
sm_scale=self.scale,
attn_sink=self.attn_sink,
topk_length=combined_lens,
out=output[query_start:query_end],
)
class DeepseekV4IndexerCache(torch.nn.Module, AttentionLayerBase):
@@ -18,6 +18,7 @@ from vllm.model_executor.layers.fused_moe.all2all_utils import (
from vllm.model_executor.layers.fused_moe.config import (
FusedMoEQuantConfig,
FusedMoEQuantDesc,
RoutingMethodType,
mxfp4_mxfp8_moe_quant_config,
mxfp4_w4a8_moe_quant_config,
mxfp4_w4a16_moe_quant_config,
@@ -64,6 +65,8 @@ class Mxfp4MoeBackend(Enum):
MARLIN = "MARLIN"
# ROCm AITER backends
AITER_MXFP4_BF16 = "AITER_MXFP4_BF16" # W4A16: CK kernel
# Keep the legacy name as an alias while the ROCm split backend rename settles.
AITER = "AITER_MXFP4_BF16"
AITER_MXFP4_FP8 = "AITER_MXFP4_FP8" # W4A8: triton kernel
# Triton
TRITON = "TRITON"
@@ -253,6 +256,8 @@ def _get_priority_backends() -> list[Mxfp4MoeBackend]:
TRTLLM MXFP8; SM90 falls through to Triton_unfused or Marlin (the
backend-level ``is_supported_config`` check filters by device capability).
"""
if current_platform.is_rocm():
return [Mxfp4MoeBackend.AITER_MXFP4_BF16]
_AVAILABLE_BACKENDS = [
Mxfp4MoeBackend.FLASHINFER_TRTLLM_MXFP4_MXFP8,
Mxfp4MoeBackend.DEEPGEMM_MXFP4,
@@ -543,8 +548,22 @@ def select_deepseek_v4_mxfp4_moe_backend(
activation_format,
)
# DeepSeek-V4 on ROCm is more accurate with the unfused Triton MXFP4 path
# than the default AITER path. Prefer Triton-unfused for this routing mode,
# while keeping AITER as a fallback if Triton-unfused rejects the config.
if (
current_platform.is_rocm()
and config.routing_method == RoutingMethodType.DeepseekV4
):
priority_backends = [
Mxfp4MoeBackend.TRITON_UNFUSED,
Mxfp4MoeBackend.AITER_MXFP4_BF16,
]
else:
priority_backends = _get_priority_backends()
# Iterate priority backends: TRTLLM MXFP8, then Triton.
for backend in _get_priority_backends():
for backend in priority_backends:
activation_key = _backend_activation_key(backend)
for k_cls in backend_to_kernel_cls(backend):
supported, reason = k_cls.is_supported_config(
@@ -1252,6 +1271,64 @@ def convert_weight_to_mxfp4_moe_kernel_format(
w2_bias,
)
elif mxfp4_backend == Mxfp4MoeBackend.AITER_MXFP4_BF16:
from vllm._aiter_ops import rocm_aiter_ops
if w13_bias is not None:
w13_bias = w13_bias.data.to(torch.float32)
if w2_bias is not None:
w2_bias = w2_bias.data.to(torch.float32)
e, n, k = w13_weight.shape
w13_weight.view(torch.uint8).copy_(
w13_weight.data.view(torch.uint8)
.view(e, n // 2, 2, k)
.permute(0, 2, 1, 3)
.contiguous()
.view(e, n, k)
)
w13_weight_scale.data = (
w13_weight_scale.data.view(e, n // 2, 2, -1)
.permute(0, 2, 1, 3)
.contiguous()
.view(e, n, -1)
)
w13_weight.data = w13_weight.data.view(torch.float4_e2m1fn_x2)
w2_weight.data = w2_weight.data.view(torch.float4_e2m1fn_x2)
w13_weight.data = rocm_aiter_ops.shuffle_weight_a16w4(w13_weight, 16, True)
shuffled_w13_scale = rocm_aiter_ops.shuffle_scale_a16w4(
w13_weight_scale.view(-1, w13_weight_scale.shape[-1]),
num_experts,
True,
)
w2_weight.data = rocm_aiter_ops.shuffle_weight_a16w4(w2_weight, 16, False)
shuffled_w2_scale = rocm_aiter_ops.shuffle_scale_a16w4(
w2_weight_scale.view(-1, w2_weight_scale.shape[-1]),
num_experts,
False,
)
if w13_bias is not None:
w13_bias = (
w13_bias.data.view(-1, n // 2, 2)
.permute(0, 2, 1)
.contiguous()
.view(-1, n)
)
return (
w13_weight,
w2_weight,
shuffled_w13_scale,
shuffled_w2_scale,
w13_bias,
w2_bias,
)
elif mxfp4_backend in TRITON_BACKENDS:
from triton_kernels.matmul_ogs import FlexCtx, PrecisionConfig
@@ -1307,7 +1384,7 @@ def convert_weight_to_mxfp4_moe_kernel_format(
else:
raise ValueError(
f"Unsupported mxfp4_backend for Mxfp4MoEMethod: {mxfp4_backend}. "
f"Expected TRTLLM or Triton backend."
f"Expected TRTLLM, Triton, or AITER backend."
)
+105 -2
View File
@@ -234,6 +234,39 @@ def mhc_pre(
num_tokens = residual_flat.shape[0]
fn_flat = fn
if current_platform.is_rocm():
x = residual_flat.view(num_tokens, hc_mult * hidden_size).to(torch.float32)
mixes = torch.matmul(x, fn_flat.t())
sqrsum = x.square().sum(dim=-1, keepdim=True)
mixes = mixes * torch.rsqrt(sqrsum / (hc_mult * hidden_size) + rms_eps)
pre_logits = mixes[:, :hc_mult] * hc_scale[0] + hc_base[:hc_mult]
pre_mix = torch.sigmoid(pre_logits) + hc_pre_eps
post_logits = (
mixes[:, hc_mult : 2 * hc_mult] * hc_scale[1]
+ hc_base[hc_mult : 2 * hc_mult]
)
post_mix = torch.sigmoid(post_logits) * hc_post_mult_value
comb_logits = mixes[:, 2 * hc_mult :].view(
num_tokens, hc_mult, hc_mult
) * hc_scale[2] + hc_base[2 * hc_mult :].view(1, hc_mult, hc_mult)
comb_mix = torch.softmax(comb_logits, dim=-1) + hc_sinkhorn_eps
comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + hc_sinkhorn_eps)
for _ in range(sinkhorn_repeat - 1):
comb_mix = comb_mix / (comb_mix.sum(dim=-1, keepdim=True) + hc_sinkhorn_eps)
comb_mix = comb_mix / (comb_mix.sum(dim=-2, keepdim=True) + hc_sinkhorn_eps)
layer_input = torch.sum(
pre_mix.unsqueeze(-1) * residual_flat.to(torch.float32), dim=1
).to(torch.bfloat16)
return (
post_mix.view(*outer_shape, hc_mult, 1),
comb_mix.view(*outer_shape, hc_mult, hc_mult),
layer_input.view(*outer_shape, hidden_size),
)
# these number are from deepgemm kernel impl
block_k = 64
block_m = 64
@@ -414,6 +447,14 @@ def mhc_post(
post_layer_mix: torch.Tensor,
comb_res_mix: torch.Tensor,
) -> torch.Tensor:
if current_platform.is_rocm():
mixed_residual = torch.einsum(
"...ij,...ih->...jh",
comb_res_mix.to(torch.float32),
residual.to(torch.float32),
)
post_term = post_layer_mix.to(torch.float32) * x.unsqueeze(-2).to(torch.float32)
return (mixed_residual + post_term).to(residual.dtype)
out = torch.empty_like(residual)
mhc_post_tilelang(
comb_res_mix,
@@ -551,6 +592,49 @@ def hc_head_fuse_tilelang(
T.pdl_trigger()
def _hc_head_fused_reference(
hs_flat: torch.Tensor,
fn: torch.Tensor,
hc_scale: torch.Tensor,
hc_base: torch.Tensor,
out: torch.Tensor,
hidden_size: int,
rms_eps: float,
hc_eps: float,
hc_mult: int,
) -> None:
"""Pure-PyTorch reference for `hc_head_fuse_tilelang`.
Used on platforms where the tilelang HIP/CUDA backend is not available
(e.g. ROCm builds shipping a tilelang wheel without `target.build.tilelang_hip`).
Mirrors the math of the tilelang kernel exactly:
x = hs_flat.flatten(-2, -1) # (T, hc_mult * H), fp32
mixes = x @ fn.T # (T, hc_mult)
rsqrt = 1 / sqrt(||x||^2 / (hc_mult * H) + rms_eps)
pre[m] = sigmoid(mixes[m] * rsqrt * hc_scale[0] + hc_base[m]) + hc_eps
out = sum_m pre[m] * hs_flat[:, m, :] # cast back to bf16
`out` is mutated in place to keep the same op contract
(`mutates_args=["out"]`).
"""
num_tokens = hs_flat.shape[0]
if num_tokens == 0:
return
x = hs_flat.reshape(num_tokens, hc_mult * hidden_size).to(torch.float32)
# fn: (hc_mult, hc_mult * hidden_size) → mixes: (T, hc_mult)
mixes = torch.matmul(x, fn.t())
sqrsum = x.square().sum(dim=-1, keepdim=True)
rsqrt = torch.rsqrt(sqrsum / (hc_mult * hidden_size) + rms_eps)
# hc_scale has shape (1,); hc_base has shape (hc_mult,)
pre_mix = torch.sigmoid(mixes * rsqrt * hc_scale[0] + hc_base) + hc_eps
# weighted sum over the hc_mult channel dim
result = torch.sum(pre_mix.unsqueeze(-1) * hs_flat.to(torch.float32), dim=1).to(
out.dtype
)
out.copy_(result)
def _hc_head_fused_kernel(
hs_flat: torch.Tensor,
fn: torch.Tensor,
@@ -563,8 +647,15 @@ def _hc_head_fused_kernel(
hc_mult: int,
) -> None:
"""Fill pre-allocated `out` (T, H) in-place with the hc_head result."""
if hs_flat.shape[0] > 0:
hc_head_fuse_tilelang(
if hs_flat.shape[0] == 0:
return
if current_platform.is_rocm():
# tilelang ships only the CUDA codegen in upstream wheels, so the HIP
# FFI target (`target.build.tilelang_hip`) is missing and the JIT call
# would raise `ValueError: Cannot find global function ...`. Use a
# numerically equivalent torch fallback instead. `mhc_pre` and
# `mhc_post` already follow this same pattern above.
_hc_head_fused_reference(
hs_flat,
fn,
hc_scale,
@@ -575,6 +666,18 @@ def _hc_head_fused_kernel(
hc_eps,
hc_mult,
)
return
hc_head_fuse_tilelang(
hs_flat,
fn,
hc_scale,
hc_base,
out,
hidden_size,
rms_eps,
hc_eps,
hc_mult,
)
direct_register_custom_op(
@@ -843,6 +843,15 @@ def w8a8_triton_block_scaled_mm(
assert len(block_size) == 2
block_n, block_k = block_size[0], block_size[1]
# Triton cannot currently bind E8M0 scale tensors directly. On ROCm,
# DeepSeek-V4 checkpoints store block scales in exponent-only E8M0 format,
# so decode them to fp32 before launching the kernel.
if current_platform.is_rocm():
if As.dtype == torch.float8_e8m0fnu:
As = _upcast_e8m0_to_fp32(As).contiguous()
if Bs.dtype == torch.float8_e8m0fnu:
Bs = _upcast_e8m0_to_fp32(Bs).contiguous()
assert A.shape[-1] == B.shape[-1]
assert A.shape[:-1] == As.shape[:-1] and A.is_contiguous()
assert triton.cdiv(A.shape[-1], block_k) == As.shape[-1]
@@ -499,13 +499,31 @@ class SparseAttnIndexer(CustomOp):
k: torch.Tensor,
weights: torch.Tensor,
):
assert not self.skip_k_cache_insert, (
"AMD platform doesn't support skip cache insert yet"
)
assert not self.use_fp4_cache, "AMD platform doesn't support fp4 cache yet"
assert isinstance(q_quant, torch.Tensor), (
"AMD sparse_attn_indexer expects a single FP8 q_quant tensor"
)
if self.skip_k_cache_insert or not rocm_aiter_ops.is_enabled():
from vllm.v1.attention.ops.rocm_aiter_mla_sparse import (
rocm_aiter_sparse_attn_indexer_native,
)
return rocm_aiter_sparse_attn_indexer_native(
hidden_states,
_encode_layer_name(self.k_cache.prefix),
self.k_cache.kv_cache,
q_quant,
k,
weights,
self.quant_block_size,
self.scale_fmt,
self.topk_tokens,
self.head_dim,
self.max_model_len,
self.max_total_seq_len,
self.topk_indices_buffer,
skip_k_cache_insert=self.skip_k_cache_insert,
)
if rocm_aiter_ops.is_enabled():
return torch.ops.vllm.rocm_aiter_sparse_attn_indexer(
hidden_states,
@@ -522,8 +540,4 @@ class SparseAttnIndexer(CustomOp):
self.max_total_seq_len,
self.topk_indices_buffer,
)
else:
raise RuntimeError(
"Sparse attention indexer ROCm custom op requires ROCm "
"Aiter ops to be enabled."
)
raise RuntimeError("Sparse attention indexer ROCm path could not be selected.")
+6 -1
View File
@@ -1245,7 +1245,12 @@ class DeepseekV4Model(nn.Module):
# DeepseekV4MultiHeadLatentAttentionWrapper.attn_gemm_parallel_execute
# (compressor kv_score, indexer.weights_proj, indexer.compressor
# kv_score). fused_wqa_wkv stays on the default stream.
aux_stream_list = [torch.cuda.Stream() for _ in range(3)]
# Disable them on ROCm because of hang issues.
aux_stream_list = (
None
if current_platform.is_rocm()
else [torch.cuda.Stream() for _ in range(3)]
)
self.device = current_platform.device_type
# Reserved topk indices buffer for all Indexer layers to reuse.
@@ -167,8 +167,12 @@ class DeepSeekV4MultiTokenPredictor(nn.Module):
)
# Three aux streams shared across all MTP layers, mirroring
# DeepseekV4Model.
aux_stream_list = [torch.cuda.Stream() for _ in range(3)]
# DeepseekV4Model. ROCm runs the same work serially for now.
aux_stream_list = (
None
if current_platform.is_rocm()
else [torch.cuda.Stream() for _ in range(3)]
)
# to map the exact layer index from weights
self.layers = torch.nn.ModuleDict(
+1
View File
@@ -409,6 +409,7 @@ class RocmPlatform(Platform):
"gptq",
"gptq_marlin", # will be overwritten with gptq
"fp8",
"deepseek_v4_fp8",
"compressed-tensors",
"fbgemm_fp8",
"gguf",
+2 -1
View File
@@ -7,6 +7,7 @@ import torch
from vllm.config import CacheConfig, VllmConfig, get_current_vllm_config
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.platforms import current_platform
from vllm.triton_utils import tl, triton
from vllm.v1.attention.backend import (
AttentionBackend,
@@ -360,7 +361,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder):
_LAYER_TYPE_C4A: None,
_LAYER_TYPE_C128A: None,
}
if num_decode_tokens == 0:
if num_decode_tokens == 0 or current_platform.is_rocm():
return out
for layer_type in self._layer_types:
# get_mla_metadata() is the official FlashMLA entry point that
@@ -9,6 +9,7 @@ INT32-packed UE8M0 on SM100) so fp8_einsum skips transform_sf_into_required_layo
import torch
from vllm.platforms import current_platform
from vllm.triton_utils import tl, triton
from vllm.utils.torch_utils import direct_register_custom_op
@@ -242,6 +243,7 @@ def _fused_inv_rope_fp8_quant_kernel_impl(
(scale_inner * tma_aligned_T, 1, tma_aligned_T),
)
grid = (tma_aligned_T, n_groups * heads_per_group)
pdl_kwargs = {} if current_platform.is_rocm() else {"launch_pdl": False}
_fused_inv_rope_fp8_quant_per_head[grid](
o,
positions,
@@ -265,7 +267,7 @@ def _fused_inv_rope_fp8_quant_kernel_impl(
HALF_ROPE=half_rope,
TMA_ALIGNED_SCALES=tma_aligned_scales,
num_stages=1,
launch_pdl=False,
**pdl_kwargs,
num_warps=1,
)
return fp8_buf, scale_buf
+528 -60
View File
@@ -2,9 +2,11 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import functools
import importlib
import math
from importlib.util import find_spec
import torch
import torch.nn.functional as F
from vllm.forward_context import get_forward_context
from vllm.platforms import current_platform
@@ -13,6 +15,11 @@ from vllm.utils.torch_utils import LayerNameType
from vllm.v1.attention.backends.mla.indexer import DeepseekV32IndexerMetadata
from vllm.v1.attention.ops.common import pack_seq_triton, unpack_seq_triton
if current_platform.is_rocm():
from vllm.platforms.rocm import _ON_GFX942
else:
_ON_GFX942 = False
@triton.jit
def _indexer_k_quant_and_cache_kernel(
@@ -230,6 +237,43 @@ def fp8_paged_mqa_logits_torch(
fp8_dtype = current_platform.fp8_dtype()
batch_size, next_n, _, dim = q.size()
if next_n == 1:
block_size = kv_cache.shape[1]
logits = torch.full(
[batch_size, max_model_len],
float("-inf"),
device=q.device,
dtype=torch.float32,
)
if context_lens.dim() > 1:
context_lens = context_lens.squeeze(-1)
kv_cache_flat = kv_cache.view(-1, block_size * (dim + 4))
for i in range(batch_size):
q_i = q[i, 0].to(torch.float32)
q_scale = weights[i]
seq_len = int(context_lens[i].item())
assert seq_len <= max_model_len
num_pages = cdiv(seq_len, block_size)
padded_seq_len = num_pages * block_size
pages = block_tables[i, :num_pages]
cache = kv_cache_flat[pages]
scale_offset = block_size * dim
cache_value = (
cache[..., :scale_offset].view(dtype=fp8_dtype).to(torch.float32)
)
cache_scale = (
cache[..., scale_offset:].view(dtype=torch.float32).contiguous()
)
cache_value = cache_value.view(padded_seq_len, dim)
cache_scale = cache_scale.view(padded_seq_len)
score = F.linear(cache_value, q_i)
score = F.relu(score)
score *= q_scale[None, :]
score = score.sum(dim=1)
score *= cache_scale
logits[i, :seq_len] = score[:seq_len]
return logits
kv_cache, scale = kv_cache[..., :dim], kv_cache[..., dim:]
scale = scale.contiguous().view(torch.float)
q = q.float()
@@ -241,20 +285,30 @@ def fp8_paged_mqa_logits_torch(
device=q.device,
dtype=torch.float32,
)
context_lens = context_lens.tolist()
for i in range(batch_size):
context_len = context_lens[i]
q_offsets = torch.arange(context_len - next_n, context_len, device="cuda")
if context_len.ndim == 0:
context_len_i = int(context_len.item())
q_offsets = torch.arange(
context_len_i - next_n, context_len_i, device=q.device
)
context_limit = torch.full(
(next_n,), context_len_i, dtype=torch.int32, device=q.device
)
else:
context_limit = context_len.to(device=q.device, dtype=torch.int32)
q_offsets = context_limit - 1
weight_slice = (
weights[i * next_n : (i + 1) * next_n, :].transpose(0, 1).contiguous()
)
for block_rk in range(cdiv(context_len, block_size)):
max_context_len = int(context_limit.max().item())
for block_rk in range(cdiv(max_context_len, block_size)):
block_idx = block_tables[i][block_rk]
qx, kx = q[i], kv_cache[block_idx]
k_offsets = torch.arange(
block_rk * block_size, (block_rk + 1) * block_size, device="cuda"
block_rk * block_size, (block_rk + 1) * block_size, device=q.device
)
mask = (k_offsets[None, :] < context_len) & (
mask = (k_offsets[None, :] < context_limit[:, None]) & (
k_offsets[None, :] <= q_offsets[:, None]
)
s = torch.where(
@@ -331,30 +385,52 @@ def rocm_fp8_paged_mqa_logits(
aiter_paged_mqa_logits_module = paged_mqa_logits_module()
if aiter_paged_mqa_logits_module is not None:
deepgemm_fp8_paged_mqa_logits = (
aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits
if _ON_GFX942:
deepgemm_fp8_paged_mqa_logits = (
aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits
)
batch_size, next_n, heads, _ = q_fp8.shape
out_logits = torch.full(
[batch_size * next_n, max_model_len],
float("-inf"),
device="cuda",
dtype=torch.float32,
)
deepgemm_fp8_paged_mqa_logits(
q_fp8,
kv_cache_fp8,
weights,
out_logits,
context_lens,
block_tables,
max_model_len,
ChunkK=256,
Preshuffle=block_size == 64,
KVBlockSize=block_size,
WavePerEU=2,
)
return out_logits
deepgemm_fp8_paged_mqa_logits_stage1 = (
aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits_stage1
)
batch_size, next_n, heads, _ = q_fp8.shape
out_logits = torch.full(
[batch_size * next_n, max_model_len],
out_qk = torch.full(
(heads, batch_size * next_n, max_model_len),
float("-inf"),
device="cuda",
dtype=torch.float32,
)
deepgemm_fp8_paged_mqa_logits(
deepgemm_fp8_paged_mqa_logits_stage1(
q_fp8,
kv_cache_fp8,
weights,
out_logits,
out_qk,
context_lens,
block_tables,
max_model_len,
ChunkK=256,
Preshuffle=block_size == 64,
KVBlockSize=block_size,
WavePerEU=2,
ChunkQ=heads,
)
return out_logits
return out_qk.sum(dim=0)
else:
return fp8_paged_mqa_logits_torch(
q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len
@@ -464,6 +540,27 @@ def rocm_fp8_mqa_logits(
return fp8_mqa_logits_torch(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke)
def _topk_indices_torch(logits: torch.Tensor, topk_tokens: int) -> torch.Tensor:
k = min(topk_tokens, logits.shape[-1])
values, indices = torch.topk(logits, k=k, dim=-1)
indices = indices.to(torch.int32)
indices = torch.where(
values == float("-inf"),
torch.full_like(indices, -1, dtype=torch.int32),
indices,
)
if k == topk_tokens:
return indices
padded = torch.full(
(logits.shape[0], topk_tokens),
-1,
dtype=torch.int32,
device=logits.device,
)
padded[:, :k] = indices
return padded
def rocm_aiter_sparse_attn_indexer_fake(
hidden_states: torch.Tensor,
k_cache_prefix: LayerNameType,
@@ -482,8 +579,9 @@ def rocm_aiter_sparse_attn_indexer_fake(
# profile run
# NOTE(Chen): create the max possible flattened_kv. So that
# profile_run can get correct memory usage.
device = hidden_states.device if k is None else k.device
_flattened_kv = torch.empty(
[total_seq_lens, head_dim + 4], device=k.device, dtype=torch.uint8
[total_seq_lens, head_dim + 4], device=device, dtype=torch.uint8
)
fp8_dtype = current_platform.fp8_dtype()
_k_fp8 = _flattened_kv[..., :head_dim].view(fp8_dtype).contiguous()
@@ -491,7 +589,7 @@ def rocm_aiter_sparse_attn_indexer_fake(
return topk_indices_buffer
def rocm_aiter_sparse_attn_indexer(
def rocm_aiter_sparse_attn_indexer_native(
hidden_states: torch.Tensor,
k_cache_prefix: LayerNameType,
kv_cache: torch.Tensor,
@@ -505,10 +603,12 @@ def rocm_aiter_sparse_attn_indexer(
max_model_len: int,
total_seq_lens: int,
topk_indices_buffer: torch.Tensor | None,
skip_k_cache_insert: bool = False,
) -> torch.Tensor:
# careful! this will be None in dummy run
attn_metadata = get_forward_context().attn_metadata
fp8_dtype = current_platform.fp8_dtype()
from vllm import _custom_ops as ops
from vllm.utils.torch_utils import _resolve_layer_name
k_cache_prefix = _resolve_layer_name(k_cache_prefix)
@@ -537,19 +637,33 @@ def rocm_aiter_sparse_attn_indexer(
has_decode = layer_attn_metadata.num_decodes > 0
has_prefill = layer_attn_metadata.num_prefills > 0
num_decode_tokens = layer_attn_metadata.num_decode_tokens
device = hidden_states.device if k is None else k.device
# during speculative decoding, k may be padded to the CUDA graph batch
# size while slot_mapping only covers actual tokens.
num_tokens = slot_mapping.shape[0]
k = k[:num_tokens]
if k is not None:
k = k[:num_tokens]
elif not skip_k_cache_insert:
raise ValueError("k must be provided when skip_k_cache_insert is False")
indexer_k_quant_and_cache_triton(
k,
kv_cache,
slot_mapping,
quant_block_size,
scale_fmt,
)
if not skip_k_cache_insert:
if _ON_GFX942:
ops.indexer_k_quant_and_cache(
k,
kv_cache,
slot_mapping,
quant_block_size,
scale_fmt,
)
else:
indexer_k_quant_and_cache_triton(
k,
kv_cache,
slot_mapping,
quant_block_size,
scale_fmt,
)
topk_indices_buffer[: hidden_states.shape[0]] = -1
if has_prefill:
@@ -558,22 +672,31 @@ def rocm_aiter_sparse_attn_indexer(
for chunk in prefill_metadata.chunks:
k_fp8 = torch.empty(
[chunk.total_seq_lens, head_dim],
device=k.device,
device=device,
dtype=fp8_dtype,
)
k_scale = torch.empty(
[chunk.total_seq_lens, 4],
device=k.device,
device=device,
dtype=torch.uint8,
)
cp_gather_indexer_k_quant_cache_triton(
kv_cache,
k_fp8,
k_scale,
chunk.block_table,
chunk.cu_seq_lens,
chunk.token_to_seq,
)
if _ON_GFX942:
ops.cp_gather_indexer_k_quant_cache(
kv_cache,
k_fp8,
k_scale,
chunk.block_table,
chunk.cu_seq_lens,
)
else:
cp_gather_indexer_k_quant_cache_triton(
kv_cache,
k_fp8,
k_scale,
chunk.block_table,
chunk.cu_seq_lens,
token_to_seq=chunk.token_to_seq,
)
logits = rocm_fp8_mqa_logits(
q_fp8[chunk.token_start : chunk.token_end],
@@ -582,21 +705,10 @@ def rocm_aiter_sparse_attn_indexer(
chunk.cu_seqlen_ks,
chunk.cu_seqlen_ke,
)
num_rows = logits.shape[0]
assert topk_tokens == 2048, "top_k_per_row assumes size 2048"
topk_indices = topk_indices_buffer[
chunk.token_start : chunk.token_end, :topk_tokens
]
torch.ops._C.top_k_per_row_prefill(
logits,
chunk.cu_seqlen_ks,
chunk.cu_seqlen_ke,
topk_indices,
num_rows,
logits.stride(0),
logits.stride(1),
topk_tokens,
)
topk_indices.copy_(_topk_indices_torch(logits, topk_tokens))
if has_decode:
decode_metadata = layer_attn_metadata.decode
@@ -633,19 +745,8 @@ def rocm_aiter_sparse_attn_indexer(
max_model_len=max_model_len,
)
num_rows = logits.shape[0]
assert topk_tokens == 2048, "top_k_per_row assumes size 2048"
topk_indices = topk_indices_buffer[:num_decode_tokens, :topk_tokens]
torch.ops._C.top_k_per_row_decode(
logits,
next_n,
decode_metadata.seq_lens,
topk_indices,
num_rows,
logits.stride(0),
logits.stride(1),
topk_tokens,
)
topk_indices.copy_(_topk_indices_torch(logits, topk_tokens)[:num_decode_tokens])
if decode_metadata.requires_padding:
# if padded, we need to unpack
@@ -659,3 +760,370 @@ def rocm_aiter_sparse_attn_indexer(
)
return topk_indices_buffer
def rocm_aiter_sparse_attn_indexer(
hidden_states: torch.Tensor,
k_cache_prefix: LayerNameType,
kv_cache: torch.Tensor,
q_fp8: torch.Tensor,
k: torch.Tensor,
weights: torch.Tensor,
quant_block_size: int,
scale_fmt: str | None,
topk_tokens: int,
head_dim: int,
max_model_len: int,
total_seq_lens: int,
topk_indices_buffer: torch.Tensor | None,
) -> torch.Tensor:
return rocm_aiter_sparse_attn_indexer_native(
hidden_states,
k_cache_prefix,
kv_cache,
q_fp8,
k,
weights,
quant_block_size,
scale_fmt,
topk_tokens,
head_dim,
max_model_len,
total_seq_lens,
topk_indices_buffer,
skip_k_cache_insert=False,
)
def _decode_e8m0_scales(scale: torch.Tensor) -> torch.Tensor:
if scale.dtype == torch.float8_e8m0fnu:
from vllm.model_executor.layers.quantization.utils.fp8_utils import (
_upcast_e8m0_to_fp32,
)
return _upcast_e8m0_to_fp32(scale).contiguous()
return scale.to(torch.float32)
def _expand_2d_block_scales(
scale: torch.Tensor,
rows: int,
cols: int,
) -> torch.Tensor:
scale = _decode_e8m0_scales(scale)
row_blocks, col_blocks = scale.shape[-2:]
row_block = math.ceil(rows / row_blocks)
col_block = math.ceil(cols / col_blocks)
scale = torch.repeat_interleave(scale, row_block, dim=-2)[..., :rows, :]
scale = torch.repeat_interleave(scale, col_block, dim=-1)[..., :, :cols]
return scale
def _apply_gptj_inv_rope_ref(
x: torch.Tensor,
positions: torch.Tensor,
cos_sin_cache: torch.Tensor,
rope_dim: int,
) -> torch.Tensor:
if rope_dim == 0 or x.numel() == 0:
return x
half_rot = rope_dim // 2
nope_dim = x.shape[-1] - rope_dim
dtype = x.dtype
x = x.to(torch.float32)
cache = cos_sin_cache.index_select(0, positions.to(torch.long))
cos = cache[:, :half_rot].to(torch.float32)
sin = cache[:, half_rot : 2 * half_rot].to(torch.float32)
view_shape = (positions.shape[0],) + (1,) * (x.dim() - 2) + (half_rot,)
cos = cos.view(view_shape)
sin = sin.view(view_shape)
rope = x[..., nope_dim:]
y_even = rope[..., 0::2]
y_odd = rope[..., 1::2]
rope_out = torch.stack(
(y_even * cos + y_odd * sin, y_odd * cos - y_even * sin),
dim=-1,
).flatten(-2)
x = x.clone()
x[..., nope_dim:] = rope_out
return x.to(dtype)
def _apply_inv_rope_ref(
rotary_emb: torch.nn.Module,
x: torch.Tensor,
positions: torch.Tensor,
rope_dim: int,
) -> torch.Tensor:
if hasattr(rotary_emb, "forward_native"):
try:
query, _ = rotary_emb.forward_native(
positions,
x.clone(),
None,
inverse=True,
)
return query
except TypeError:
pass
return _apply_gptj_inv_rope_ref(x, positions, rotary_emb.cos_sin_cache, rope_dim)
def rocm_inv_rope_einsum(
rotary_emb: torch.nn.Module,
o: torch.Tensor,
positions: torch.Tensor,
rope_head_dim: int,
n_local_groups: int,
o_lora_rank: int,
wo_a: torch.nn.Module,
) -> torch.Tensor:
"""Reference inverse-RoPE + WO_A einsum path used on ROCm."""
o_ref = _apply_inv_rope_ref(rotary_emb, o, positions, rope_head_dim).to(
torch.bfloat16
)
o_ref = o_ref.view(o.shape[0], n_local_groups, -1)
hidden_dim = o_ref.shape[-1]
if hasattr(wo_a, "weight_scale_inv"):
wo_a_weight = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to(
torch.float32
)
wo_a_scale = _expand_2d_block_scales(
wo_a.weight_scale_inv.view(
n_local_groups, -1, wo_a.weight_scale_inv.shape[-1]
),
o_lora_rank,
hidden_dim,
)
wo_a_weight = (wo_a_weight * wo_a_scale).to(torch.bfloat16)
else:
wo_a_weight = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to(
torch.bfloat16
)
return torch.einsum("tgd,grd->tgr", o_ref, wo_a_weight)
def rocm_ref_sparse_attn_prefill(
q: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
topk_length: torch.Tensor | None,
scale: float,
head_dim: int,
attn_sink: torch.Tensor | None,
) -> torch.Tensor:
indices = indices.clone().squeeze(1)
s_q, h_q, d_qk = q.shape
topk = indices.shape[-1]
s_kv = kv.shape[0]
if topk_length is not None:
mask = torch.arange(topk, device=indices.device).unsqueeze(
0
) >= topk_length.unsqueeze(1)
indices[mask] = -1
invalid_mask = (indices < 0) | (indices >= s_kv)
indices[invalid_mask] = 0
qf = q.float()
gathered_kv = kv.index_select(0, indices.flatten()).reshape(s_q, topk, d_qk).float()
scores = qf @ gathered_kv.transpose(1, 2)
scores *= scale
scores[invalid_mask.unsqueeze(1).expand_as(scores)] = float("-inf")
orig_lse = torch.logsumexp(scores, dim=-1)
lse_for_o = orig_lse
if attn_sink is not None:
lse_for_o = torch.logsumexp(
torch.stack(
[orig_lse, attn_sink[:h_q].view(1, h_q).expand_as(orig_lse)],
dim=0,
),
dim=0,
)
lse_for_o = lse_for_o.clone()
lse_for_o[lse_for_o == float("-inf")] = float("+inf")
probs = torch.exp(scores - lse_for_o.unsqueeze(-1))
out = probs @ gathered_kv[..., :head_dim]
lonely_q_mask = orig_lse == float("-inf")
out[lonely_q_mask.unsqueeze(-1).expand_as(out)] = 0.0
return out.to(torch.bfloat16)
def rocm_sparse_attn_prefill(
q: torch.Tensor,
kv: torch.Tensor,
indices: torch.Tensor,
topk_length: torch.Tensor | None,
scale: float,
head_dim: int,
attn_sink: torch.Tensor | None,
output: torch.Tensor,
) -> None:
output_chunk = rocm_ref_sparse_attn_prefill(
q=q,
kv=kv,
indices=indices,
topk_length=topk_length,
scale=scale,
head_dim=head_dim,
attn_sink=attn_sink,
)
output.copy_(output_chunk.to(output.dtype))
def rocm_dequantize_blocked_k_cache(
quant_k_cache: torch.Tensor,
head_dim: int,
nope_head_dim: int,
rope_head_dim: int,
) -> torch.Tensor:
fp8_dtype = current_platform.fp8_dtype()
tile_size = 64
num_tiles = nope_head_dim // tile_size
num_blocks, block_size, _ = quant_k_cache.shape
quant_k_cache = quant_k_cache.view(num_blocks, -1)
input_nope_rope = quant_k_cache[
:, : block_size * (nope_head_dim + 2 * rope_head_dim)
].view(num_blocks, block_size, nope_head_dim + 2 * rope_head_dim)
input_nope = input_nope_rope[:, :, :nope_head_dim].view(fp8_dtype)
input_rope = input_nope_rope[:, :, nope_head_dim:].view(torch.bfloat16)
input_scale = (
quant_k_cache[:, block_size * (nope_head_dim + 2 * rope_head_dim) :]
.view(num_blocks, block_size, 8)[:, :, :num_tiles]
.view(torch.float8_e8m0fnu)
)
result = torch.empty(
(num_blocks, block_size, 1, head_dim),
dtype=torch.bfloat16,
device=quant_k_cache.device,
)
result[..., nope_head_dim:] = input_rope.unsqueeze(2)
for tile_idx in range(num_tiles):
cur_nope = input_nope[
..., tile_idx * tile_size : (tile_idx + 1) * tile_size
].to(torch.bfloat16)
cur_scales = input_scale[:, :, tile_idx].to(torch.bfloat16).unsqueeze(-1)
result[..., tile_idx * tile_size : (tile_idx + 1) * tile_size] = (
cur_nope * cur_scales
).unsqueeze(2)
return result
def rocm_ref_sparse_attn_decode(
q: torch.Tensor,
blocked_k: torch.Tensor,
indices_in_kvcache: torch.Tensor,
topk_length: torch.Tensor | None,
scale: float,
head_dim: int,
attn_sink: torch.Tensor | None,
extra_blocked_k: torch.Tensor | None = None,
extra_indices_in_kvcache: torch.Tensor | None = None,
extra_topk_length: torch.Tensor | None = None,
) -> torch.Tensor:
b, s_q, h_q, d_qk = q.shape
def process_scope(
cur_blocked_k: torch.Tensor,
cur_indices: torch.Tensor,
cur_topk_length: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
cur_indices = cur_indices.reshape(b, s_q, -1)
topk = cur_indices.size(-1)
fixed_indices = torch.clamp_min(cur_indices, 0)
gathered_kv = (
cur_blocked_k.view(-1, d_qk)
.index_select(0, fixed_indices.view(-1))
.view(b, s_q, topk, d_qk)
)
invalid_mask = cur_indices == -1
if cur_topk_length is not None:
cur_topk_length = cur_topk_length.reshape(b)
invalid_mask |= torch.arange(0, topk, device=invalid_mask.device).view(
1, 1, topk
) >= cur_topk_length.view(b, 1, 1)
return gathered_kv, invalid_mask
gathered_kv, invalid_mask = process_scope(
blocked_k, indices_in_kvcache, topk_length
)
if extra_blocked_k is not None:
assert extra_indices_in_kvcache is not None
gathered_kv1, invalid_mask1 = process_scope(
extra_blocked_k, extra_indices_in_kvcache, extra_topk_length
)
gathered_kv = torch.cat([gathered_kv, gathered_kv1], dim=2)
invalid_mask = torch.cat([invalid_mask, invalid_mask1], dim=2)
gathered_kv = gathered_kv.view(b * s_q, -1, d_qk).float()
gathered_kv[gathered_kv != gathered_kv] = 0.0
qf = q.float().view(b * s_q, h_q, d_qk)
attn_weight = qf @ gathered_kv.transpose(-1, -2)
attn_weight *= scale
attn_weight[
invalid_mask.view(b * s_q, 1, -1).expand(b * s_q, h_q, invalid_mask.size(-1))
] = float("-inf")
lse = attn_weight.logsumexp(dim=-1)
attn_weight = torch.exp(attn_weight - lse.unsqueeze(-1))
output = attn_weight @ gathered_kv[..., :head_dim]
output = output.view(b, s_q, h_q, head_dim)
lse = lse.view(b, s_q, h_q)
if attn_sink is not None:
output *= (1.0 / (1.0 + torch.exp(attn_sink.view(1, 1, h_q) - lse))).unsqueeze(
-1
)
lonely_q_mask = lse == float("-inf")
output[lonely_q_mask.unsqueeze(-1).expand_as(output)] = 0.0
return output.squeeze(1).to(torch.bfloat16)
def rocm_forward_decode_fallback(
q: torch.Tensor,
kv_cache: torch.Tensor | None,
swa_k_cache: torch.Tensor,
swa_only: bool,
topk_indices: torch.Tensor | None,
topk_lens: torch.Tensor | None,
swa_indices: torch.Tensor,
swa_lens: torch.Tensor,
attn_sink: torch.Tensor | None,
scale: float,
head_dim: int,
nope_head_dim: int,
rope_head_dim: int,
output: torch.Tensor,
) -> None:
blocked_swa = rocm_dequantize_blocked_k_cache(
swa_k_cache,
head_dim=head_dim,
nope_head_dim=nope_head_dim,
rope_head_dim=rope_head_dim,
)
blocked_extra = None
if not swa_only:
assert kv_cache is not None
blocked_extra = rocm_dequantize_blocked_k_cache(
kv_cache,
head_dim=head_dim,
nope_head_dim=nope_head_dim,
rope_head_dim=rope_head_dim,
)
attn_out = rocm_ref_sparse_attn_decode(
q=q.unsqueeze(1),
blocked_k=blocked_swa,
indices_in_kvcache=swa_indices.unsqueeze(1),
topk_length=swa_lens,
scale=scale,
head_dim=head_dim,
attn_sink=attn_sink[: q.shape[1]] if attn_sink is not None else None,
extra_blocked_k=blocked_extra,
extra_indices_in_kvcache=topk_indices,
extra_topk_length=topk_lens,
)
output.copy_(attn_out.to(output.dtype))