Amd/deepseek v4 rebase main 0509 (#24933)

Co-authored-by: root <root@smci355-ccs-aus-m12-33.cs-aus.dcgpu>
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: HaiShaw <hixiao@gmail.com>
Co-authored-by: amd-danli103 <danli103@amd.com>
Co-authored-by: Lin, Soga <soga.lin@amd.com>
Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
Co-authored-by: Hubert Lu <55214931+hubertlu-tw@users.noreply.github.com>
Co-authored-by: yichiche@amd.com <jacky.cheng>
Co-authored-by: yctseng0211 <yctseng@amd.com>
Co-authored-by: Bingxu Chen <bingxche@amd.com>
This commit is contained in:
kk
2026-05-18 09:15:07 -07:00
committed by GitHub
co-authored by root wunhuang Thomas Wang Xinyi Song HaiShaw amd-danli103 Lin, Soga Raiden-Makoto Hubert Lu yichiche@amd.com yctseng0211 Bingxu Chen
parent 110bbdcad7
commit 866793c502
17 changed files with 3677 additions and 69 deletions
+26
View File
@@ -13,6 +13,13 @@ from sglang.jit_kernel.utils import (
make_cpp_args,
)
from sglang.srt.environ import envs
from sglang.srt.utils import get_bool_env_var, is_hip
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _use_aiter:
from aiter.tuned_gemm import tgemm
if TYPE_CHECKING:
from tvm_ffi.module import Module
@@ -644,6 +651,23 @@ def fused_rope(
positions: torch.Tensor,
inverse: bool = False,
) -> None:
"""Apply rotary embeddings to both Q and K in a single fused CUDA kernel.
Args:
q: [batch_size, num_q_heads, rope_dim] bfloat16
k: [batch_size, num_k_heads, rope_dim] bfloat16 or None
freqs_cis: [max_seq_len, rope_dim // 2] complex64 (full table)
positions: [batch_size] int32 or int64, indices into freqs_cis
inverse: if True, apply inverse rotation (conjugate freqs)
"""
if _is_hip:
from sglang.srt.layers.deepseek_v4_rope import apply_rotary_emb_triton
apply_rotary_emb_triton(q, freqs_cis, positions=positions, inverse=inverse)
if k is not None:
apply_rotary_emb_triton(k, freqs_cis, positions=positions, inverse=inverse)
return
freqs_real = torch.view_as_real(freqs_cis).flatten(-2).contiguous()
module = _jit_fused_rope_module()
module.forward(q, k, freqs_real, positions, inverse)
@@ -1029,5 +1053,7 @@ def _dispatch_bf16_fp32_backend(
z = x.new_empty(x.size(0), y.size(0), dtype=torch.float32)
deep_gemm.bf16_gemm_nt(x, y, z)
return z
elif _use_aiter:
return tgemm.mm(x, y, otype=torch.float32)
else:
return torch.nn.functional.linear(x.float(), y.float())
+7
View File
@@ -571,6 +571,13 @@ class Envs:
# ====================================================================
# DeepSeek V4
SGLANG_OPT_DPSK_V4_RADIX = EnvBool(True)
SGLANG_OPT_USE_OLD_COMPRESSOR = EnvBool(False)
SGLANG_OPT_USE_TRITON_SWA_PREPARE = EnvBool(True)
SGLANG_OPT_USE_AITER_MHC_PRE = EnvBool(True)
SGLANG_OPT_USE_AITER_MHC_POST = EnvBool(True)
SGLANG_OPT_USE_FUSED_COMPRESS = EnvBool(False)
SGLANG_FIX_MTP_HC_HIDDEN = EnvBool(False)
# ====================================================================
# Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint.
@@ -105,11 +105,24 @@ def create_nsa_backend(runner):
@register_attention_backend("dsv4")
def create_dsv4_backend(runner):
from sglang.srt.layers.attention.deepseek_v4_backend import (
DeepseekV4AttnBackend,
)
from sglang.srt.utils import is_hip
return DeepseekV4AttnBackend(runner)
if is_hip():
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
DeepseekV4HipRadixBackend,
)
logger.info(
"Using DeepseekV4HipRadixBackend for compressed attention backend (HIP)."
)
return DeepseekV4HipRadixBackend(runner)
else:
from sglang.srt.layers.attention.deepseek_v4_backend import (
DeepseekV4AttnBackend,
)
logger.info("Using DeepseekV4AttnBackend for dsv4 attention backend (CUDA).")
return DeepseekV4AttnBackend(runner)
@register_attention_backend("triton")
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,455 @@
from __future__ import annotations
import os
from functools import cached_property
from typing import TYPE_CHECKING, Any
import torch
import torch.nn as nn
import triton
import triton.language as tl
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsv4.compressor import Compressor as _CompressorBase
from sglang.srt.layers.attention.nsa.nsa_indexer import rotate_activation
from sglang.srt.layers.deepseek_v4_rope import (
apply_rotary_emb_triton,
fused_norm_rope_inplace_triton,
)
from sglang.srt.mem_cache.deepseek_v4_compress_state import (
CompressStatePool,
KVAndScore,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
if TYPE_CHECKING:
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
DeepseekV4HipRadixBackend,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@triton.jit
def _rms_normalize_kernel(
x_ptr,
weight_ptr,
eps,
stride_row,
dim,
BLOCK_SIZE: tl.constexpr,
HAS_WEIGHT: tl.constexpr,
):
pid = tl.program_id(0)
offs = tl.arange(0, BLOCK_SIZE)
mask = offs < dim
base = pid * stride_row
x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32)
mean_sq = tl.sum(x * x, axis=0) / dim
rms_inv = tl.rsqrt(mean_sq + eps)
out = x * rms_inv
if HAS_WEIGHT:
weight = tl.load(weight_ptr + offs, mask=mask, other=0.0)
out = out * weight
tl.store(x_ptr + base + offs, out, mask=mask)
def rms_normalize_triton(
x: torch.Tensor, eps: float, weight: torch.Tensor = None
) -> torch.Tensor:
dim = x.shape[-1]
x_flat = x.view(-1, dim)
num_rows = x_flat.shape[0]
BLOCK_SIZE = triton.next_power_of_2(dim)
grid = (num_rows,)
_rms_normalize_kernel[grid](
x_flat,
weight,
eps,
x_flat.stride(0),
dim,
BLOCK_SIZE=BLOCK_SIZE,
HAS_WEIGHT=(weight is not None),
)
return x
class DeepseekRefRMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32))
def forward(self, x: torch.Tensor):
return rms_normalize_triton(x, self.eps, self.weight)
class CompressorHip(_CompressorBase):
"""HIP (ROCm) specific Compressor implementation."""
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.norm = DeepseekRefRMSNorm(self.head_dim, eps=self.norm.variance_epsilon)
@cached_property
def use_fused_compress(self) -> bool:
return False
@cached_property
def use_hip_fused_compress(self) -> bool:
return envs.SGLANG_OPT_USE_FUSED_COMPRESS.get()
def _get_states(self, forward_batch: ForwardBatch) -> KVAndScore:
token_to_kv_pool = forward_batch.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if self.is_in_indexer:
return token_to_kv_pool.get_indexer_compress_states(self.layer_id)
else:
return token_to_kv_pool.get_attention_compress_states(self.layer_id)
def _get_state_pool(self, forward_batch: ForwardBatch) -> CompressStatePool:
token_to_kv_pool = forward_batch.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
if self.is_in_indexer:
ret = token_to_kv_pool.get_indexer_compress_states(self.layer_id)
else:
ret = token_to_kv_pool.get_attention_compress_states(self.layer_id)
assert isinstance(ret, CompressStatePool)
return ret
def overlap_transform(self, tensor: torch.Tensor, fill_value: Any) -> torch.Tensor:
assert tensor.dim() == 3
assert tensor.shape[1:] == (self.ratio, 2 * self.head_dim)
s, r, d = tensor.size(0), self.ratio, self.head_dim
new_tensor = tensor.new_full((s, 2 * r, d), fill_value)
new_tensor[:, r:] = tensor[:, :, d:]
new_tensor[1:, :r] = tensor[:-1, :, :d]
return new_tensor
def overlap_transform_decode(self, tensor: torch.Tensor) -> torch.Tensor:
assert tensor.dim() == 3
assert tensor.shape[1:] == (2 * self.ratio, 2 * self.head_dim)
r, d = self.ratio, self.head_dim
ret = torch.cat((tensor[:, :r, :d], tensor[:, r:, d:]), dim=1)
return ret
@staticmethod
def compute_state_len(seq_len: int, ratio: int):
return seq_len % ratio + (ratio == 4) * ratio
@staticmethod
def compute_state_len_indices(seq_len: int, ratio: int):
state_len = seq_len % ratio + (ratio == 4) * ratio
return torch.arange(seq_len - state_len, seq_len).clamp(min=-1)
def print_tensor(self, y: torch.Tensor, name: str):
enable = int(os.environ.get("SGLANG_ENABLE_PRINT_TENSOR", 0))
if enable:
print(f"[sgl] {name}: shape={y.shape}, dtype={y.dtype}, device={y.device}")
print(f"{y.flatten()[:10]}...{y.flatten()[-10:]}")
def compress_extend_paged(
self,
kv_and_scores: KVAndScore,
forward_batch: ForwardBatch,
):
backend = forward_batch.attn_backend
if TYPE_CHECKING:
assert isinstance(backend, DeepseekV4HipRadixBackend)
token_to_kv_pool = forward_batch.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
state_pool = self._get_state_pool(forward_batch)
prefix_lens = forward_batch.extend_prefix_lens_cpu
extend_lens = forward_batch.extend_seq_lens_cpu
req_pool_indices = forward_batch.req_pool_indices
req_to_token = forward_batch.req_to_token_pool.req_to_token
assert not self.forward_mode.is_target_verify()
assert extend_lens is not None and prefix_lens is not None
device = kv_and_scores.kv.device
assert kv_and_scores.kv.shape[-1] == self.head_dim * self.coff
compressed_kv_output = torch.full(
(kv_and_scores.kv.size(0), self.head_dim),
fill_value=10000.0,
dtype=kv_and_scores.kv.dtype,
device=device,
)
bs = forward_batch.batch_size
pt = 0
for i in range(bs):
kv_and_score = kv_and_scores[pt : pt + extend_lens[i]]
pre_state_indices = self.compute_state_len_indices(
seq_len=prefix_lens[i], ratio=self.ratio
).to(device)
raw_loc = torch.where(
pre_state_indices < 0,
-1,
req_to_token[req_pool_indices[i], pre_state_indices],
)
swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(raw_loc)
state_loc = state_pool.translate_from_swa_loc_to_state_loc(swa_loc)
pre_kv_state = state_pool.get_state_by_state_loc(state_loc)
kv_and_score_buffer = KVAndScore.cat([pre_kv_state, kv_and_score], dim=0)
valid_kv_len = kv_and_score_buffer.kv.size(0)
post_state_indices = self.compute_state_len_indices(
seq_len=prefix_lens[i] + extend_lens[i], ratio=self.ratio
).to(device)
post_state_len = post_state_indices.size(0)
assert post_state_len <= valid_kv_len
post_raw_loc = torch.where(
post_state_indices < 0,
-1,
req_to_token[req_pool_indices[i], post_state_indices],
)
post_swa_loc = token_to_kv_pool.translate_loc_from_full_to_swa(post_raw_loc)
post_state_loc = state_pool.translate_from_swa_loc_to_state_loc(
post_swa_loc
)
post_state_to_set = kv_and_score_buffer[valid_kv_len - post_state_len :]
state_pool.set_state_by_state_loc(post_state_loc, post_state_to_set)
compress_len = valid_kv_len // self.ratio * self.ratio
if compress_len == 0:
pt += extend_lens[i]
continue
kv_and_score_to_compress = kv_and_score_buffer[:compress_len].view(
compress_len // self.ratio, self.ratio, -1
)
kv_and_score_to_compress.score.add_(self.ape.unsqueeze(0))
if self.overlap:
new_kv = self.overlap_transform(
kv_and_score_to_compress.kv, fill_value=0
)
new_score = self.overlap_transform(
kv_and_score_to_compress.score, fill_value=float("-inf")
)
kv_and_score_to_compress = KVAndScore.from_kv_score(
kv=new_kv, score=new_score
)
del new_kv, new_score
kv_and_score_to_compress = kv_and_score_to_compress[1:]
if kv_and_score_to_compress.kv.size(0) == 0:
pt += extend_lens[i]
continue
kv_compressed = (
kv_and_score_to_compress.kv
* kv_and_score_to_compress.score.softmax(dim=1)
).sum(dim=1)
assert kv_compressed.dtype == torch.float32
beg_idx = prefix_lens[i] // self.ratio * self.ratio
end_idx = (prefix_lens[i] + extend_lens[i]) // self.ratio * self.ratio
freqs_cis = self.freqs_cis[beg_idx : end_idx : self.ratio]
assert freqs_cis.size(0) == kv_compressed.size(
0
), f"{freqs_cis.shape=} {kv_compressed.shape=}"
if self.use_hip_fused_compress:
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
else:
kv_compressed = self.norm(kv_compressed)
apply_rotary_emb_triton(
kv_compressed[..., -self.rope_head_dim :], freqs_cis
)
del beg_idx, end_idx
if self.rotate:
kv_compressed = rotate_activation(kv_compressed)
start = prefix_lens[i]
start = start + self.ratio - 1 - start % self.ratio
indices_in_seq = torch.arange(
start,
prefix_lens[i] + extend_lens[i],
self.ratio,
device=kv_and_scores.kv.device,
)
assert indices_in_seq.size(0) == kv_compressed.size(0)
compressed_kv_output[indices_in_seq - prefix_lens[i] + pt] = kv_compressed
pt += extend_lens[i]
return compressed_kv_output
def compress_decode_paged(
self,
kv_and_scores: KVAndScore,
forward_batch: ForwardBatch,
):
"""Paged and cudagraph compatible version of compress_decode"""
assert self.ape_converted
state_pool = self._get_state_pool(forward_batch)
token_to_kv_pool = forward_batch.token_to_kv_pool
assert isinstance(token_to_kv_pool, DeepSeekV4TokenToKVPool)
req_pool_indices = forward_batch.req_pool_indices
req_to_token = forward_batch.req_to_token_pool.req_to_token
seq_lens = forward_batch.seq_lens
if forward_batch.forward_mode.is_target_verify():
draft_tokens = forward_batch.attn_backend.speculative_num_draft_tokens
offsets = torch.arange(1, draft_tokens + 1, device=seq_lens.device)
seq_lens_2d = seq_lens[:, None] + offsets[None, :]
seq_lens = seq_lens_2d.view(-1)
req_pool_indices = req_pool_indices.repeat_interleave(draft_tokens)
raw_locs = req_to_token[req_pool_indices, seq_lens - 1]
swa_locs = token_to_kv_pool.translate_loc_from_full_to_swa(raw_locs)
state_locs = state_pool.translate_from_swa_loc_to_state_loc(swa_locs)
state_pool.set_state_by_state_loc(state_locs, kv_and_scores)
compress_bulk_len = self.ratio * self.coff
compress_indices = seq_lens[:, None] + torch.arange(
-compress_bulk_len, 0, device=seq_lens.device
)
compress_indices.clamp_(min=-1)
compress_indices_raw = torch.where(
compress_indices < 0,
-1,
req_to_token[req_pool_indices[:, None], compress_indices],
)
compress_indices_swa = token_to_kv_pool.translate_loc_from_full_to_swa(
compress_indices_raw
)
compress_indices_state = state_pool.translate_from_swa_loc_to_state_loc(
compress_indices_swa
)
kv_and_score_to_compress = state_pool.get_state_by_state_loc(
compress_indices_state.view(-1)
).view(-1, self.ratio, self.coff * self.head_dim)
kv_and_score_to_compress.score.add_(self.ape.unsqueeze(0))
bs = seq_lens.size(0)
if self.overlap:
kv_and_score_to_compress = kv_and_score_to_compress.view(
bs, self.coff * self.ratio, self.coff * self.head_dim
)
kv_and_score_to_compress = KVAndScore.from_kv_score(
kv=self.overlap_transform_decode(kv_and_score_to_compress.kv),
score=self.overlap_transform_decode(kv_and_score_to_compress.score),
)
self.print_tensor(kv_and_score_to_compress.kv, "kv_to_compress")
self.print_tensor(kv_and_score_to_compress.score, "score_to_compress")
kv_and_score_to_compress = kv_and_score_to_compress.view(
bs, self.ratio * self.coff, self.head_dim
)
kv_compressed = (
kv_and_score_to_compress.kv * kv_and_score_to_compress.score.softmax(dim=1)
).sum(dim=1)
self.print_tensor(kv_compressed, "kv_before_norm")
if self.use_hip_fused_compress:
freqs_cis = self._init_freqs_cis_per_decode_step(forward_batch, seq_lens)
fused_norm_rope_inplace_triton(
kv_compressed, self.norm.weight, self.norm.eps, freqs_cis
)
else:
kv_compressed = self.norm(kv_compressed)
self.print_tensor(kv_compressed, "kv_after_norm")
freqs_cis = self.freqs_cis[(seq_lens - 1) // self.ratio * self.ratio]
self.print_tensor(freqs_cis, "freqs_cis")
apply_rotary_emb_triton(
kv_compressed[..., -self.rope_head_dim :], freqs_cis
)
self.print_tensor(kv_compressed, "kv_after_rope")
if self.rotate:
kv_compressed = rotate_activation(kv_compressed)
self.print_tensor(kv_compressed, "compressed_kv_output")
return kv_compressed
def compress_fused(
self,
kv_score: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
backend = forward_batch.attn_backend
if TYPE_CHECKING:
assert isinstance(backend, DeepseekV4HipRadixBackend)
kv_score_buffer = self._get_state_pool(forward_batch)
kv_score_buffer = kv_score_buffer.kv_score_buffer.kv_score
return backend.forward_compress(
kv_score_buffer=kv_score_buffer,
kv_score_input=kv_score,
ape=self.ape.view(-1, self.head_dim),
head_dim=self.head_dim,
norm=self.norm,
freqs_cis_cache=self.freqs_cis,
rotate=self.rotate,
compress_ratio=self.ratio,
forward_batch=forward_batch,
is_paged=True,
)
def compress_dispatch(
self,
kv_score: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
if self.use_fused_compress:
return self.compress_fused(kv_score, forward_batch)
self.compress_decode = self.compress_decode_paged
self.compress_extend = self.compress_extend_paged
kv_and_scores = KVAndScore(kv_score)
if TYPE_CHECKING:
assert isinstance(kv_and_scores, KVAndScore)
if (
forward_batch.forward_mode.is_decode()
or forward_batch.forward_mode.is_target_verify()
):
result = self.compress_decode(
kv_and_scores=kv_and_scores,
forward_batch=forward_batch,
)
elif forward_batch.forward_mode.is_extend():
result = self.compress_extend(
kv_and_scores=kv_and_scores,
forward_batch=forward_batch,
)
else:
msg = f"Forward mode {forward_batch.forward_mode} not supported in Compressor."
raise NotImplementedError(msg)
return result
def _init_freqs_cis_per_decode_step(
self,
forward_batch: ForwardBatch,
seq_lens: torch.Tensor,
) -> torch.Tensor:
attr = f"freqs_cis_c{self.ratio}"
cached = getattr(forward_batch, attr, None)
if cached is not None:
return cached
decoded = self.freqs_cis[(seq_lens - 1) // self.ratio * self.ratio]
setattr(forward_batch, attr, decoded)
return decoded
def forward(self, x: torch.Tensor, forward_batch: ForwardBatch) -> torch.Tensor:
if forward_batch.forward_mode.is_idle():
assert x.shape[0] == 0
return x.new_empty(0, self.head_dim)
kv_score = self.compute_kv_score(x, forward_batch)
self.forward_mode = forward_batch.forward_mode
return self.compress_dispatch(kv_score, forward_batch)
@@ -24,12 +24,16 @@ from sglang.srt.layers.dp_attention import get_attention_cp_size
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import ReplicatedLinear
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
from sglang.srt.mem_cache.deepseek_v4_compress_state import (
CompressStatePool,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.models.deepseek_v2 import _is_hip
from sglang.srt.utils import add_prefix
if TYPE_CHECKING:
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
from sglang.srt.layers.rotary_embedding import RotaryEmbedding
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -293,6 +297,7 @@ class Compressor(nn.Module):
head_dim: int,
rotate: bool = False,
prefix: str = "",
rotary_emb: Optional[RotaryEmbedding] = None,
) -> None:
super().__init__()
self.layer_id = layer_id
@@ -304,7 +309,7 @@ class Compressor(nn.Module):
self.ratio = compress_ratio
self.overlap = self.ratio == 4
self.rotate = rotate
coff = 1 + self.overlap
self.coff = coff = 1 + self.overlap
self.ape = nn.Parameter(
torch.empty(self.ratio, coff * self.head_dim, dtype=torch.float32)
@@ -321,6 +326,7 @@ class Compressor(nn.Module):
self.norm = RMSNorm(
self.head_dim, eps=config.rms_norm_eps, weight_dtype=torch.float32
)
self.rotary_emb = rotary_emb
self.freqs_cis = freqs_cis
self.ape_converted = False
@@ -350,6 +356,8 @@ class Compressor(nn.Module):
# NOTE: used by v2 compressor backend
def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch):
kv_score = linear_bf16_fp32(x, self.wkv_gate.weight)
# CUDA path: delegate to backend
if nsa_use_prefill_cp(forward_batch):
kv_score = cp_all_gather_rerange_output(
kv_score,
@@ -383,3 +391,9 @@ class Compressor(nn.Module):
forward_batch=forward_batch,
is_paged=True,
)
if _is_hip:
from sglang.srt.layers.attention.dsv4.compress_hip import ( # noqa: F811
CompressorHip as Compressor,
)
@@ -22,7 +22,6 @@ from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
from sglang.srt.utils import add_prefix, is_hip
if TYPE_CHECKING:
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
from sglang.srt.layers.attention.dsv4.compressor import (
CompressorBackendMixin,
)
@@ -100,7 +99,7 @@ def topk_transform_512_pytorch_vectorized(
out_raw_indices: Optional[torch.Tensor] = None,
) -> None:
TOPK = 512
TOPK = out_page_indices.shape[1]
batch_size = scores.shape[0]
max_seq_len = scores.shape[1]
device = scores.device
@@ -332,11 +331,6 @@ class C4IndexerBackendMixin:
indexer_metadata = metadata.indexer_metadata
core_metadata = metadata.core_metadata
from sglang.srt.layers.attention.deepseek_v4_backend import (
DSV4AttnMetadata,
)
assert isinstance(core_metadata, DSV4AttnMetadata)
assert isinstance(indexer_metadata, PagedIndexerMetadata)
if enable_multi_stream:
@@ -374,7 +368,7 @@ class C4IndexerBackendMixin:
assert len(weights.shape) == 3
weights = weights.squeeze(2)
if envs.SGLANG_OPT_USE_TILELANG_INDEXER.get():
from sglang.srt.layers.attention.dsv4.tilelang_kernel import (
from sglang.srt.layers.attention.nsa.tilelang_kernel import (
tilelang_fp8_paged_mqa_logits as fn,
)
elif envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get():
@@ -383,7 +377,8 @@ class C4IndexerBackendMixin:
from deep_gemm import fp8_paged_mqa_logits as fn
_c4sl = indexer_metadata.c4_seq_lens
if _c4sl.dim() == 1:
_use_tilelang = envs.SGLANG_OPT_USE_TILELANG_INDEXER.get()
if _c4sl.dim() == 1 and not _use_tilelang:
_c4sl = _c4sl.unsqueeze(-1)
logits = fn(
q_fp8,
@@ -479,6 +474,7 @@ class C4Indexer(nn.Module):
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
alt_streams: Optional[List[torch.cuda.Stream]] = None,
rotary_emb=None,
):
super().__init__()
self.layer_id = layer_id
@@ -486,6 +482,7 @@ class C4Indexer(nn.Module):
self.n_heads = config.index_n_heads
self.head_dim = config.index_head_dim
self.rope_head_dim = config.qk_rope_head_dim
self.index_topk = config.index_topk
self.q_lora_rank = config.q_lora_rank
self.softmax_scale = self.head_dim**-0.5
self.n_local_heads = self.n_heads
@@ -514,7 +511,9 @@ class C4Indexer(nn.Module):
head_dim=self.head_dim,
rotate=True,
prefix=add_prefix("compressor", prefix),
rotary_emb=rotary_emb,
)
self.rotary_emb = rotary_emb
self.freqs_cis = freqs_cis
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
self.alt_streams = alt_streams
@@ -545,8 +544,6 @@ class C4Indexer(nn.Module):
enable_multi_stream: bool = False,
q_lora_ready: Optional[torch.cuda.Event] = None,
) -> None:
if TYPE_CHECKING:
assert isinstance(forward_batch.attn_backend, DeepseekV4AttnBackend)
return forward_batch.attn_backend.forward_c4_indexer(
x=x,
q_lora=q_lora,
@@ -0,0 +1,197 @@
from typing import Any, Optional
import torch
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
from sglang.srt.utils import is_hip
FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn
def flash_mla_with_kvcache_entrypoint(backend: str, **kwargs):
if is_hip():
import os
from sglang.srt.layers.attention.nsa.tilelang_kernel import (
dpsk_v4_fp8_attention_fwd,
)
backend = os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "tilelang")
else:
import flash_mla
if backend == "comparison":
pack_ref, pack_fast_via_tester = flash_mla_with_kvcache_entrypoint(
backend="torch", **kwargs
)
pack_fast_via_api = flash_mla_with_kvcache_entrypoint(
backend="kernel", **kwargs
)
_assert_close(pack_ref=pack_fast_via_tester, pack_fast=pack_fast_via_api)
_assert_close(pack_ref=pack_ref, pack_fast=pack_fast_via_tester)
_assert_close(pack_ref=pack_ref, pack_fast=pack_fast_via_api)
return pack_ref
if backend == "torch":
return flash_mla_with_kvcache_torch(**kwargs)
if backend == "tilelang":
return dpsk_v4_fp8_attention_fwd(**kwargs)
if backend == "kernel":
return flash_mla.flash_mla_with_kvcache(**kwargs)
raise NotImplementedError(f"unknown backend: {backend!r}")
def flash_mla_with_kvcache_torch(
q: torch.Tensor,
k_cache: torch.Tensor,
block_table: Optional[torch.Tensor],
cache_seqlens: Optional[torch.Tensor],
head_dim_v: int,
tile_scheduler_metadata: Any,
num_splits: None = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
is_fp8_kvcache: bool = False,
indices: Optional[torch.Tensor] = None,
attn_sink: Optional[torch.Tensor] = None,
extra_k_cache: Optional[torch.Tensor] = None,
extra_indices_in_kvcache: Optional[torch.Tensor] = None,
topk_length: Optional[torch.Tensor] = None,
extra_topk_length: Optional[torch.Tensor] = None,
):
from sglang.srt.flashmla_tests import quant as flashmla_quant
from sglang.srt.flashmla_tests.lib import (
ExtraTestParamForDecode,
KVScope,
TestcaseForDecode,
TestParam,
)
from sglang.srt.flashmla_tests.ref import ref_sparse_attn_decode
assert block_table is None
assert cache_seqlens is None
assert is_fp8_kvcache
b, s_q, h_q, d_qk = q.shape
d_v = head_dim_v
fp8_layout = flashmla_quant.FP8KVCacheLayout.MODEL1_FP8Sparse
p = TestParam(
s_q=s_q,
s_kv="unused",
topk="unused",
h_q=h_q,
h_kv=1,
d_qk=d_qk,
d_v=d_v,
decode=ExtraTestParamForDecode(
b=b,
is_varlen="unused",
have_zero_seqlen_k="unused",
extra_s_k="unused",
extra_topk="unused",
extra_block_size="unused",
have_extra_topk_length="unused",
),
# unused?
seed=-1,
check_correctness=True,
is_all_indices_invalid=False,
num_runs=10,
have_attn_sink=True,
have_topk_length=True,
)
blocked_k_quantized = k_cache
blocked_k = flashmla_quant.dequantize_k_cache(
blocked_k_quantized.view(FP8_DTYPE), fp8_layout
)
# blocked_k_requantized = flashmla_quant.quantize_k_cache(blocked_k, fp8_layout)
# assert torch.testing.assert_allclose(blocked_k_requantized.byte(), blocked_k_quantized.byte())
kv_scope = KVScope(
t="unused",
cache_seqlens="unused",
block_table="unused",
blocked_k=blocked_k,
blocked_k_quantized=blocked_k_quantized,
abs_indices="unused",
indices_in_kvcache=indices,
topk_length=topk_length,
)
extra_kv_scope = None
if extra_k_cache is not None:
extra_blocked_k_quantized = extra_k_cache
extra_blocked_k = flashmla_quant.dequantize_k_cache(
extra_blocked_k_quantized.view(FP8_DTYPE), fp8_layout
)
# extra_blocked_k_requantized = flashmla_quant.quantize_k_cache(extra_blocked_k, fp8_layout)
# assert torch.testing.assert_allclose(extra_blocked_k_requantized.byte(), extra_blocked_k_quantized.byte())
extra_kv_scope = KVScope(
t="unused",
cache_seqlens="unused",
block_table="unused",
blocked_k=extra_blocked_k,
blocked_k_quantized=extra_blocked_k_quantized,
abs_indices="unused",
indices_in_kvcache=extra_indices_in_kvcache,
topk_length=extra_topk_length,
)
t = TestcaseForDecode(
p="unused",
q=q,
attn_sink=attn_sink,
sm_scale=softmax_scale,
kv_scope=kv_scope,
extra_kv_scope=extra_kv_scope,
)
# print(f"hi {p=} {t=}")
# print(
# f"hi info "
# f"{get_tensor_info(t.kv_scope.blocked_k)=} "
# f"{get_tensor_info(t.kv_scope.blocked_k_quantized)=} "
# f"{get_tensor_info(t.extra_kv_scope.blocked_k) if t.extra_kv_scope is not None else None=} "
# f"{get_tensor_info(t.extra_kv_scope.blocked_k_quantized) if t.extra_kv_scope is not None else None=} "
# )
pack_ref = ref_sparse_attn_decode(p, t)
# tile_scheduler_metadata, _ = flash_mla.get_mla_metadata()
# pack_fast_via_tester = flashmla_lib.run_flash_mla_decode(
# p, t, tile_scheduler_metadata, num_splits=None
# )
# return pack_ref, pack_fast_via_tester
return pack_ref
def _assert_close(pack_ref, pack_fast):
import sglang.srt.flashmla_tests.kernelkit as kk
out_ref, lse_ref = pack_ref
out_fast, lse_fast = pack_fast
# the copied threshold is too strict, not checked why
# copied from: test_flash_mla_sparse_decoding.py
# is_out_correct = kk.check_is_allclose(
# "out", out_fast, out_ref, abs_tol=1e-3, rel_tol=2.01 / 128, cos_diff_tol=5e-6
# )
# is_lse_correct = kk.check_is_allclose(
# "lse", lse_fast, lse_ref, abs_tol=1e-6, rel_tol=8.01 / 65536
# )
# loosen thresh
is_out_correct = kk.check_is_allclose(
"out", out_fast, out_ref, abs_tol=1e-2, rel_tol=10.0, cos_diff_tol=5e-6
)
is_lse_correct = kk.check_is_allclose(
"lse", lse_fast, lse_ref, abs_tol=1e-6, rel_tol=8.01 / 65536
)
assert is_out_correct and is_lse_correct, f"{is_out_correct=} {is_lse_correct=}"
@@ -432,10 +432,6 @@ def _set_k_and_s_triton(
assert (
page_size % 16 == 0
), f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}"
else:
assert (
page_size == 1
), f"HIP legacy NSA path requires page_size == 1, got {page_size}"
else:
assert page_size == 64
File diff suppressed because it is too large Load Diff
@@ -177,3 +177,171 @@ def apply_rotary_emb_triton(
)
return x
@triton.jit
def _fused_norm_rope_kernel(
x_ptr,
weight_ptr,
freqs_real_ptr,
positions_ptr,
eps,
stride_x_row,
stride_freq_row,
HEAD_DIM: tl.constexpr,
ROPE_DIM: tl.constexpr,
HEAD_BLOCK: tl.constexpr,
ROPE_PAIR_BLOCK: tl.constexpr,
HAS_WEIGHT: tl.constexpr,
USE_POS: tl.constexpr,
):
# NOTE: avoids store-then-reload on the same kernel: rope-segment values
# are loaded a 2nd time as (real, imag) pairs straight from the input,
# rms_inv/weight applied in register, and all stores happen at the end.
pid = tl.program_id(0)
base = pid.to(tl.int64) * stride_x_row
offs = tl.arange(0, HEAD_BLOCK)
mask = offs < HEAD_DIM
x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32)
sum_sq = tl.sum(x * x, axis=0)
rms_inv = tl.rsqrt(sum_sq / HEAD_DIM + eps)
if HAS_WEIGHT:
w = tl.load(weight_ptr + offs, mask=mask, other=0.0).to(tl.float32)
x_normed = x * rms_inv * w
else:
x_normed = x * rms_inv
rope_start = HEAD_DIM - ROPE_DIM
pair_offs = tl.arange(0, ROPE_PAIR_BLOCK)
pair_mask = pair_offs < (ROPE_DIM // 2)
x_real = tl.load(
x_ptr + base + rope_start + 2 * pair_offs,
mask=pair_mask,
other=0.0,
).to(tl.float32)
x_imag = tl.load(
x_ptr + base + rope_start + 2 * pair_offs + 1,
mask=pair_mask,
other=0.0,
).to(tl.float32)
if HAS_WEIGHT:
w_real = tl.load(
weight_ptr + rope_start + 2 * pair_offs,
mask=pair_mask,
other=1.0,
).to(tl.float32)
w_imag = tl.load(
weight_ptr + rope_start + 2 * pair_offs + 1,
mask=pair_mask,
other=1.0,
).to(tl.float32)
x_real = x_real * rms_inv * w_real
x_imag = x_imag * rms_inv * w_imag
else:
x_real = x_real * rms_inv
x_imag = x_imag * rms_inv
if USE_POS:
position = tl.load(positions_ptr + pid).to(tl.int64)
else:
position = pid.to(tl.int64)
freq_base = position * stride_freq_row
f_real = tl.load(
freqs_real_ptr + freq_base + 2 * pair_offs,
mask=pair_mask,
other=0.0,
).to(tl.float32)
f_imag = tl.load(
freqs_real_ptr + freq_base + 2 * pair_offs + 1,
mask=pair_mask,
other=0.0,
).to(tl.float32)
out_real = x_real * f_real - x_imag * f_imag
out_imag = x_real * f_imag + x_imag * f_real
is_non_rope = offs < rope_start
tl.store(
x_ptr + base + offs,
x_normed.to(x_ptr.dtype.element_ty),
mask=mask & is_non_rope,
)
tl.store(
x_ptr + base + rope_start + 2 * pair_offs,
out_real.to(x_ptr.dtype.element_ty),
mask=pair_mask,
)
tl.store(
x_ptr + base + rope_start + 2 * pair_offs + 1,
out_imag.to(x_ptr.dtype.element_ty),
mask=pair_mask,
)
def fused_norm_rope_inplace_triton(
kv: torch.Tensor,
weight: Optional[torch.Tensor],
eps: float,
freqs_cis: torch.Tensor,
positions: Optional[torch.Tensor] = None,
) -> None:
"""Fused RMSNorm (over head_dim) + RoPE (on last rope_dim of head_dim), in-place.
Equivalent to::
kv = rms_normalize(kv, eps, weight)
apply_rotary_emb_triton(kv[..., -rope_dim:], freqs_cis, positions=positions)
Args:
kv: [M, head_dim], any float dtype, contiguous along last dim. Modified in-place.
weight: [head_dim] or None.
eps: RMSNorm epsilon.
freqs_cis: complex tensor.
- If ``positions`` is None: shape [M, rope_dim // 2], one freq per token.
- Else: shape [max_seq, rope_dim // 2], full table; indexed by ``positions``.
positions: optional [M] int tensor, absolute positions to index into ``freqs_cis``.
"""
assert kv.dim() == 2 and kv.stride(-1) == 1
M, head_dim = kv.shape
freqs_real = torch.view_as_real(freqs_cis).flatten(-2)
rope_dim = freqs_real.shape[-1]
assert head_dim >= rope_dim and rope_dim % 2 == 0
if weight is not None:
assert weight.shape == (head_dim,)
if positions is None:
assert (
freqs_real.shape[0] == M
), f"freqs_cis row count {freqs_real.shape[0]} != M={M}"
else:
assert positions.shape == (M,) and positions.dim() == 1
if M == 0:
return
HEAD_BLOCK = triton.next_power_of_2(head_dim)
ROPE_PAIR_BLOCK = max(triton.next_power_of_2(rope_dim // 2), 1)
grid = (M,)
_fused_norm_rope_kernel[grid](
kv,
weight,
freqs_real,
positions,
eps,
kv.stride(0),
freqs_real.stride(0),
HEAD_DIM=head_dim,
ROPE_DIM=rope_dim,
HEAD_BLOCK=HEAD_BLOCK,
ROPE_PAIR_BLOCK=ROPE_PAIR_BLOCK,
HAS_WEIGHT=(weight is not None),
USE_POS=(positions is not None),
)
+3 -1
View File
@@ -1407,7 +1407,9 @@ def select_experts(
scoring_func=scoring_func,
)
elif custom_routing_function is None:
assert not apply_routed_scaling_factor_on_output, "Not implemented"
if scoring_func != "sqrtsoftplus":
assert not apply_routed_scaling_factor_on_output, "Not implemented"
if scoring_func == "sqrtsoftplus":
_biased_topk = (
biased_topk_jit_kernel_impl
+143 -16
View File
@@ -84,6 +84,7 @@ from sglang.srt.utils import (
get_bool_env_var,
is_cpu,
is_cuda,
is_gfx95_supported,
is_hip,
is_musa,
is_npu,
@@ -111,9 +112,21 @@ _is_cpu = is_cpu()
_is_fp8_fnuz = is_fp8_fnuz()
_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()
def _require_fp4_dtype():
fp4_dtype = getattr(torch, "float4_e2m1fn_x2", None)
if fp4_dtype is None:
raise RuntimeError(
"DeepSeek-V4 FP4 experts require torch.float4_e2m1fn_x2 support."
)
return fp4_dtype
if _use_aiter or _use_hip_int4:
from aiter.ops.shuffle import shuffle_weight
from aiter.utility.fp4_utils import e8m0_shuffle
if _use_aiter:
from sglang.srt.layers.quantization.fp8_utils import (
@@ -998,12 +1011,13 @@ class Fp8MoEMethod(FusedMoEMethodBase):
# WEIGHT_SCALES
if self.is_fp4_expert:
fp4_block_k = 32
fp4_scale_dtype = torch.float8_e8m0fnu if _use_aiter else torch.float32
w13_weight_scale = torch.nn.Parameter(
torch.ones(
num_experts,
2 * intermediate_size_per_partition,
hidden_size // fp4_block_k,
dtype=torch.float32,
dtype=fp4_scale_dtype,
),
requires_grad=False,
)
@@ -1012,7 +1026,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
num_experts,
hidden_size,
intermediate_size_per_partition // fp4_block_k,
dtype=torch.float32,
dtype=fp4_scale_dtype,
),
requires_grad=False,
)
@@ -1123,6 +1137,105 @@ class Fp8MoEMethod(FusedMoEMethodBase):
layer.w2_input_scale = None
def process_weights_after_loading_block_quant(self, layer: Module) -> None:
# AMD FP4 experts: use aiter's native MXFP4 MoE path
if _use_aiter and self.is_fp4_expert:
fp4_weight_dtype = _require_fp4_dtype()
# CK FP4 MoE kernel requires K_packed divisible by 128
# (i.e., K_logical divisible by 256).
# Pad intermediate_size_per_partition if needed.
fp4_k_align = 256
E, w13_N, w13_K_packed = layer.w13_weight.shape
_, w2_N, w2_K_packed = layer.w2_weight.shape
inter_per_part = w13_N // 2
padded_inter = (
(inter_per_part + fp4_k_align - 1) // fp4_k_align * fp4_k_align
)
if padded_inter != inter_per_part:
pad_amount = padded_inter - inter_per_part
fp4_block_k = 32
# Pad w13_weight: (E, 2*inter, K_packed) → (E, 2*padded, K_packed)
old_w13 = layer.w13_weight.data
new_w13 = torch.zeros(
E,
2 * padded_inter,
w13_K_packed,
dtype=old_w13.dtype,
device=old_w13.device,
)
new_w13[:, :inter_per_part, :] = old_w13[:, :inter_per_part, :]
new_w13[:, padded_inter : padded_inter + inter_per_part, :] = old_w13[
:, inter_per_part:, :
]
layer.w13_weight = torch.nn.Parameter(new_w13, requires_grad=False)
# Pad w2_weight: (E, N, inter_packed) → (E, N, padded_packed)
old_w2 = layer.w2_weight.data
new_w2 = torch.zeros(
E,
w2_N,
padded_inter // 2,
dtype=old_w2.dtype,
device=old_w2.device,
)
new_w2[:, :, :w2_K_packed] = old_w2
layer.w2_weight = torch.nn.Parameter(new_w2, requires_grad=False)
# Pad w13 scale: (E, 2*inter, K/block_k) → (E, 2*padded, K/block_k)
old_s13 = layer.w13_weight_scale_inv.data
_, _, s13_K = old_s13.shape
new_s13 = torch.zeros(
E,
2 * padded_inter,
s13_K,
dtype=old_s13.dtype,
device=old_s13.device,
)
new_s13[:, :inter_per_part, :] = old_s13[:, :inter_per_part, :]
new_s13[:, padded_inter : padded_inter + inter_per_part, :] = old_s13[
:, inter_per_part:, :
]
layer.w13_weight_scale_inv = torch.nn.Parameter(
new_s13, requires_grad=False
)
# Pad w2 scale: (E, N, inter/block_k) → (E, N, padded/block_k)
old_s2 = layer.w2_weight_scale_inv.data
new_s2 = torch.zeros(
E,
w2_N,
padded_inter // fp4_block_k,
dtype=old_s2.dtype,
device=old_s2.device,
)
new_s2[:, :, : old_s2.shape[2]] = old_s2
layer.w2_weight_scale_inv = torch.nn.Parameter(
new_s2, requires_grad=False
)
for scale_name in ("w13_weight_scale_inv", "w2_weight_scale_inv"):
scale = getattr(layer, scale_name)
num_experts, num_rows, _ = scale.shape
scale.data = e8m0_shuffle(scale.view(num_experts * num_rows, -1)).view(
num_experts, num_rows, -1
)
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(
layer.w13_weight.contiguous(), (16, 16)
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight.contiguous(), (16, 16)
)
layer.w13_weight.is_shuffled = is_shuffled
layer.w2_weight.is_shuffled = is_shuffled
return
# If ROCm, normalize the weights and scales to e4m3fnuz
if _is_fp8_fnuz:
# activation_scheme: dynamic
@@ -1148,8 +1261,6 @@ class Fp8MoEMethod(FusedMoEMethodBase):
)
layer.w2_input_scale = None
if _use_aiter:
# add this section for MI300
# Pre-shuffle weights
layer.w13_weight.data = shuffle_weight(
layer.w13_weight.contiguous(), (16, 16)
)
@@ -1158,12 +1269,12 @@ class Fp8MoEMethod(FusedMoEMethodBase):
)
elif _use_aiter:
# Pre-shuffle weights
layer.w13_weight.data = shuffle_weight(
layer.w13_weight.contiguous(), (16, 16)
)
layer.w2_weight.data = shuffle_weight(
layer.w2_weight.contiguous(), (16, 16)
)
t = shuffle_weight(layer.w13_weight, (16, 16))
layer.w13_weight.copy_(t)
del t
t = shuffle_weight(layer.w2_weight, (16, 16))
layer.w2_weight.copy_(t)
del t
elif _is_cpu:
assert (
_is_cpu_amx_available
@@ -1190,8 +1301,9 @@ class Fp8MoEMethod(FusedMoEMethodBase):
layer.w2_weight.data = layer.w2_weight.data.view(torch.int8)
return
layer.w13_weight.data = layer.w13_weight.data.view(torch.int8)
layer.w2_weight.data = layer.w2_weight.data.view(torch.int8)
fp4_weight_dtype = _require_fp4_dtype() if _use_aiter else torch.int8
layer.w13_weight.data = layer.w13_weight.data.view(fp4_weight_dtype)
layer.w2_weight.data = layer.w2_weight.data.view(fp4_weight_dtype)
if get_moe_a2a_backend().is_megamoe():
from sglang.srt.layers.moe.mega_moe import (
@@ -1930,8 +2042,23 @@ class Fp8MoEMethod(FusedMoEMethodBase):
AiterQuantType,
)
if _use_aiter and self.block_quant:
quant_type = AiterQuantType.PER_128X128
w13_weight = layer.w13_weight
w2_weight = layer.w2_weight
if self.block_quant:
quant_type = (
AiterQuantType.PER_1X32
if self.is_fp4_expert
else AiterQuantType.PER_128X128
)
if self.is_fp4_expert:
fp4_weight_dtype = _require_fp4_dtype()
w13_weight = w13_weight.view(fp4_weight_dtype)
w2_weight = w2_weight.view(fp4_weight_dtype)
if getattr(layer.w13_weight, "is_shuffled", False):
w13_weight.is_shuffled = True
w2_weight.is_shuffled = True
w13_scale = layer.w13_weight_scale_inv
w2_scale = layer.w2_weight_scale_inv
else:
@@ -1939,8 +2066,8 @@ class Fp8MoEMethod(FusedMoEMethodBase):
w13_scale = layer.w13_weight_scale1
w2_scale = layer.w2_weight_scale1
return AiterMoeQuantInfo(
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
w13_weight=w13_weight,
w2_weight=w2_weight,
quant_type=quant_type,
w13_scale=w13_scale,
w2_scale=w2_scale,
@@ -7,8 +7,11 @@ import torch
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
from sglang.srt.utils import is_hip
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
_is_hip = is_hip()
@dataclasses.dataclass
class KVAndScore:
@@ -22,16 +25,55 @@ class KVAndScore:
def score(self) -> torch.Tensor:
return self.kv_score[..., self._item_size :]
@property
def shape(self):
return self.kv_score.shape
def __post_init__(self):
self._item_size = self.kv_score.shape[-1] // 2
@staticmethod
def from_kv_score(*, kv: torch.Tensor, score: torch.Tensor) -> KVAndScore:
assert kv.shape == score.shape
return KVAndScore(torch.cat([kv, score], dim=-1))
def new_empty(self, new_shape) -> KVAndScore:
assert new_shape[-1] == self._item_size
new_shape = list(new_shape)
new_shape[-1] = 2 * self._item_size
return KVAndScore(self.kv_score.new_empty(new_shape, requires_grad=False))
def __getitem__(self, index) -> KVAndScore:
return KVAndScore(self.kv_score[index])
def __setitem__(self, index, value: KVAndScore):
self.kv_score[index] = value.kv_score
def clear(self):
self.kv.zero_()
self.score.fill_(float("-inf"))
def view(self, *args):
args = list(args)
if isinstance(args[-1], int) and args[-1] != -1:
args[-1] = 2 * self._item_size
return KVAndScore(self.kv_score.view(*args))
def clone(self) -> KVAndScore:
return KVAndScore(self.kv_score.clone())
@staticmethod
def cat(tensors: list[KVAndScore], dim: int) -> KVAndScore:
assert dim != -1, "Concatenation along last dim is not supported."
assert len(tensors) > 0, "At least one tensor is required for concatenation."
item_size = tensors[0]._item_size
for v in tensors:
assert (
v._item_size == item_size
), "All tensors must have the same item size."
return KVAndScore(torch.cat([v.kv_score for v in tensors], dim=dim))
class CompressStatePool:
def __init__(
@@ -45,8 +87,11 @@ class CompressStatePool:
enable_memory_saver: bool,
ratio: int,
online: bool = False,
swa_page_size: int = 0,
):
self.ring_size = ring_size
self.swa_page_size = swa_page_size
self.enable_memory_saver = enable_memory_saver
if online:
assert ring_size == 1, "online compress requires ring_size=1"
@@ -57,25 +102,47 @@ class CompressStatePool:
self._size = (self._size + ratio - 1) // ratio * ratio
last_dim = 2 * (1 + overlap) * head_dim
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
)
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
maybe_init_custom_mem_pool(device=device)
)
if _is_hip:
self.kv_score_buffer = KVAndScore(
torch.empty((self._size, last_dim), dtype=dtype, device=device)
)
if not online:
self.kv_score_buffer[-1].clear()
else:
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=enable_memory_saver
)
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
maybe_init_custom_mem_pool(device=device)
)
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
self.kv_score_buffer = KVAndScore(
torch.empty(
(self._size, last_dim),
dtype=dtype,
device=device,
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
self.kv_score_buffer = KVAndScore(
torch.empty(
(self._size, last_dim),
dtype=dtype,
device=device,
)
)
)
if not online:
self.kv_score_buffer[-1].clear()
if not online:
self.kv_score_buffer[-1].clear()
def translate_from_swa_loc_to_state_loc(
self, swa_loc: torch.Tensor
) -> torch.Tensor:
swa_pages = swa_loc // self.swa_page_size
state_loc = swa_pages * self.ring_size + (swa_loc % self.ring_size)
state_loc = torch.where(swa_loc < 0, -1, state_loc)
return state_loc
def get_state_by_state_loc(self, state_loc: torch.Tensor) -> KVAndScore:
return self.kv_score_buffer[state_loc]
def set_state_by_state_loc(self, state_loc: torch.Tensor, value: KVAndScore):
self.kv_score_buffer[state_loc] = value
self.kv_score_buffer[-1].clear()
@@ -18,11 +18,13 @@ from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
from sglang.srt.mem_cache.memory_pool import KVCache
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import ceil_div
from sglang.srt.utils import ceil_div, is_hip
logger = logging.getLogger(__name__)
ONLINE_C128 = envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
_is_hip = is_hip()
ONLINE_C128 = not _is_hip and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get()
def get_compress_state_ring_size(
@@ -144,6 +146,9 @@ class DeepSeekV4SingleKVPool(KVCache):
)
def get_key_buffer(self, layer_id: int):
if self.store_dtype != self.dtype:
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
return self.kv_buffer[layer_id]
def set_kv_buffer(self, *args, **kwargs) -> None:
@@ -466,7 +471,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
)
self.c4_indexer_kv_pool = DeepSeekV4IndexerPool(
self.c4_logical_size,
self.c4_logical_size if not _is_hip else c4_size,
c4_page_size,
dtype,
indexer_head_dim,
@@ -477,7 +482,10 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
self._init_compressed_layer_mapping()
self._init_paged_compress_states(enable_memory_saver)
if _is_hip:
self._init_paged_compress_states(False)
else:
self._init_paged_compress_states(enable_memory_saver)
self._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get()
self.cached_loc = None
@@ -585,6 +593,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
dtype=self.state_dtype,
enable_memory_saver=enable_memory_saver,
ratio=ratio,
swa_page_size=self.swa_page_size,
)
def _init_compressed_layer_mapping(self):
+5
View File
@@ -575,6 +575,11 @@ class DeepseekV2MoE(nn.Module):
use_grouped_topk=False,
scoring_func=config.scoring_func,
is_fp4_experts=getattr(quant_config, "is_fp4_experts", False),
apply_routed_scaling_factor_on_output=(
True
if _use_aiter
else self.experts.should_fuse_routed_scaling_factor_in_topk
),
)
self.topk = TopK(**topk_kwargs)
+53 -5
View File
@@ -58,6 +58,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessor
from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
from sglang.srt.layers.rotary_embedding import get_rope_wrapper
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
from sglang.srt.layers.utils.cp_utils import (
cp_all_gather_rerange_output,
@@ -76,6 +77,12 @@ from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_lo
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.dbrx import ReplicatedLinear
from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu
if not _is_hip:
from sglang.srt.layers.utils.cp_utils import (
prepare_context_parallel_metadata,
)
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
LazyValue,
@@ -94,6 +101,9 @@ if TYPE_CHECKING:
from sglang.srt.layers.attention.deepseek_v4_backend import (
DeepseekV4AttnBackend,
)
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
DeepseekV4HipRadixBackend,
)
from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -200,6 +210,16 @@ class MQALayer(nn.Module):
rope_base = config.compress_rope_theta if self.compress_ratio else rope_theta
self.rotary_emb = get_rope_wrapper(
head_size=self.rope_head_dim,
rotary_dim=self.rope_head_dim,
max_position=config.max_position_embeddings,
base=rope_base,
rope_scaling=rope_scaling,
is_neox_style=False,
device=get_global_server_args().device,
)
from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis
assert self.compress_ratio in {0, 4, 128}
@@ -243,6 +263,7 @@ class MQALayer(nn.Module):
head_dim=self.head_dim,
rotate=False,
prefix=add_prefix("compressor", prefix),
rotary_emb=getattr(self, "rotary_emb", None),
)
if self.compress_ratio == 4:
self.indexer = C4Indexer(
@@ -252,10 +273,11 @@ class MQALayer(nn.Module):
quant_config=quant_config,
prefix=add_prefix("indexer", prefix),
alt_streams=self.alt_streams_indexer,
rotary_emb=getattr(self, "rotary_emb", None),
)
self.attn_sink = nn.Parameter(torch.empty(self.n_heads, dtype=torch.float32))
self.fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get()
self.fuse_wqa_wkv = not _is_hip and envs.SGLANG_OPT_FUSE_WQA_WKV.get()
if self.fuse_wqa_wkv:
self.wqkv_a = ReplicatedLinear(
self.hidden_size,
@@ -409,7 +431,7 @@ class MQALayer(nn.Module):
x: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
attn_backend: DeepseekV4AttnBackend,
attn_backend,
q_out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
assert self.alt_streams is not None
@@ -469,7 +491,7 @@ class MQALayer(nn.Module):
x: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
attn_backend: DeepseekV4AttnBackend,
attn_backend,
q_out: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
if self.fuse_wqa_wkv:
@@ -530,7 +552,10 @@ class MQALayer(nn.Module):
attn_backend = forward_batch.attn_backend
if TYPE_CHECKING:
assert isinstance(attn_backend, DeepseekV4AttnBackend)
assert isinstance(
attn_backend,
(DeepseekV4AttnBackend, DeepseekV4HipRadixBackend),
)
enable_multi_stream = (
envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get()
@@ -717,6 +742,22 @@ class DeepseekV4DecoderLayer(nn.Module):
)
return y, post.squeeze(-1), comb, norm is not None
if _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_PRE.get():
from aiter.ops.mhc import mhc_pre
post, comb, y = mhc_pre(
residual=x,
fn=hc_fn,
hc_scale=hc_scale,
hc_base=hc_base,
rms_eps=self.rms_norm_eps,
hc_pre_eps=self.hc_eps,
hc_sinkhorn_eps=self.hc_eps,
hc_post_mult_value=2.0,
sinkhorn_repeat=self.hc_sinkhorn_iters,
)
return y, post.squeeze(-1), comb, False
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get():
import deep_gemm
@@ -765,6 +806,13 @@ class DeepseekV4DecoderLayer(nn.Module):
return mhc_post(x, residual, post, comb)
elif _is_hip and envs.SGLANG_OPT_USE_AITER_MHC_POST.get():
from aiter.ops.mhc import mhc_post
result = torch.empty_like(residual)
mhc_post(result, x, residual, post, comb)
return result
assert residual.shape == (x.shape[0], self.hc_mult, x.shape[-1])
assert post.shape == (x.shape[0], self.hc_mult)
assert comb.shape == (x.shape[0], self.hc_mult, self.hc_mult)
@@ -1284,7 +1332,7 @@ class DeepseekV4ForCausalLM(nn.Module):
cache_compressor_weight = {}
COMPRESSOR_PART = ".compressor.w"
fuse_wqa_wkv = envs.SGLANG_OPT_FUSE_WQA_WKV.get()
fuse_wqa_wkv = not _is_hip and envs.SGLANG_OPT_FUSE_WQA_WKV.get()
cache_wqkv_a_weight: dict[str, dict[str, torch.Tensor]] = {}
def auto_weight_loader(module):