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:
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
@@ -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())
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user