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,
|
make_cpp_args,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
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:
|
if TYPE_CHECKING:
|
||||||
from tvm_ffi.module import Module
|
from tvm_ffi.module import Module
|
||||||
@@ -644,6 +651,23 @@ def fused_rope(
|
|||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
inverse: bool = False,
|
inverse: bool = False,
|
||||||
) -> None:
|
) -> 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()
|
freqs_real = torch.view_as_real(freqs_cis).flatten(-2).contiguous()
|
||||||
module = _jit_fused_rope_module()
|
module = _jit_fused_rope_module()
|
||||||
module.forward(q, k, freqs_real, positions, inverse)
|
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)
|
z = x.new_empty(x.size(0), y.size(0), dtype=torch.float32)
|
||||||
deep_gemm.bf16_gemm_nt(x, y, z)
|
deep_gemm.bf16_gemm_nt(x, y, z)
|
||||||
return z
|
return z
|
||||||
|
elif _use_aiter:
|
||||||
|
return tgemm.mm(x, y, otype=torch.float32)
|
||||||
else:
|
else:
|
||||||
return torch.nn.functional.linear(x.float(), y.float())
|
return torch.nn.functional.linear(x.float(), y.float())
|
||||||
|
|||||||
@@ -571,6 +571,13 @@ class Envs:
|
|||||||
|
|
||||||
# ====================================================================
|
# ====================================================================
|
||||||
# DeepSeek V4
|
# 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.
|
# Set False when using FP4-to-FP8 converted DeepSeek V4 checkpoint.
|
||||||
|
|||||||
@@ -105,11 +105,24 @@ def create_nsa_backend(runner):
|
|||||||
|
|
||||||
@register_attention_backend("dsv4")
|
@register_attention_backend("dsv4")
|
||||||
def create_dsv4_backend(runner):
|
def create_dsv4_backend(runner):
|
||||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
from sglang.srt.utils import is_hip
|
||||||
DeepseekV4AttnBackend,
|
|
||||||
)
|
|
||||||
|
|
||||||
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")
|
@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.layernorm import RMSNorm
|
||||||
from sglang.srt.layers.linear import ReplicatedLinear
|
from sglang.srt.layers.linear import ReplicatedLinear
|
||||||
from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_output
|
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.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
|
from sglang.srt.models.deepseek_v2 import _is_hip
|
||||||
from sglang.srt.utils import add_prefix
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
|
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
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
|
|
||||||
@@ -293,6 +297,7 @@ class Compressor(nn.Module):
|
|||||||
head_dim: int,
|
head_dim: int,
|
||||||
rotate: bool = False,
|
rotate: bool = False,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
rotary_emb: Optional[RotaryEmbedding] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
@@ -304,7 +309,7 @@ class Compressor(nn.Module):
|
|||||||
self.ratio = compress_ratio
|
self.ratio = compress_ratio
|
||||||
self.overlap = self.ratio == 4
|
self.overlap = self.ratio == 4
|
||||||
self.rotate = rotate
|
self.rotate = rotate
|
||||||
coff = 1 + self.overlap
|
self.coff = coff = 1 + self.overlap
|
||||||
|
|
||||||
self.ape = nn.Parameter(
|
self.ape = nn.Parameter(
|
||||||
torch.empty(self.ratio, coff * self.head_dim, dtype=torch.float32)
|
torch.empty(self.ratio, coff * self.head_dim, dtype=torch.float32)
|
||||||
@@ -321,6 +326,7 @@ class Compressor(nn.Module):
|
|||||||
self.norm = RMSNorm(
|
self.norm = RMSNorm(
|
||||||
self.head_dim, eps=config.rms_norm_eps, weight_dtype=torch.float32
|
self.head_dim, eps=config.rms_norm_eps, weight_dtype=torch.float32
|
||||||
)
|
)
|
||||||
|
self.rotary_emb = rotary_emb
|
||||||
self.freqs_cis = freqs_cis
|
self.freqs_cis = freqs_cis
|
||||||
|
|
||||||
self.ape_converted = False
|
self.ape_converted = False
|
||||||
@@ -350,6 +356,8 @@ class Compressor(nn.Module):
|
|||||||
# NOTE: used by v2 compressor backend
|
# NOTE: used by v2 compressor backend
|
||||||
def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch):
|
def compute_kv_score(self, x: torch.Tensor, forward_batch: ForwardBatch):
|
||||||
kv_score = linear_bf16_fp32(x, self.wkv_gate.weight)
|
kv_score = linear_bf16_fp32(x, self.wkv_gate.weight)
|
||||||
|
|
||||||
|
# CUDA path: delegate to backend
|
||||||
if nsa_use_prefill_cp(forward_batch):
|
if nsa_use_prefill_cp(forward_batch):
|
||||||
kv_score = cp_all_gather_rerange_output(
|
kv_score = cp_all_gather_rerange_output(
|
||||||
kv_score,
|
kv_score,
|
||||||
@@ -383,3 +391,9 @@ class Compressor(nn.Module):
|
|||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
is_paged=True,
|
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
|
from sglang.srt.utils import add_prefix, is_hip
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
|
|
||||||
from sglang.srt.layers.attention.dsv4.compressor import (
|
from sglang.srt.layers.attention.dsv4.compressor import (
|
||||||
CompressorBackendMixin,
|
CompressorBackendMixin,
|
||||||
)
|
)
|
||||||
@@ -100,7 +99,7 @@ def topk_transform_512_pytorch_vectorized(
|
|||||||
out_raw_indices: Optional[torch.Tensor] = None,
|
out_raw_indices: Optional[torch.Tensor] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
||||||
TOPK = 512
|
TOPK = out_page_indices.shape[1]
|
||||||
batch_size = scores.shape[0]
|
batch_size = scores.shape[0]
|
||||||
max_seq_len = scores.shape[1]
|
max_seq_len = scores.shape[1]
|
||||||
device = scores.device
|
device = scores.device
|
||||||
@@ -332,11 +331,6 @@ class C4IndexerBackendMixin:
|
|||||||
indexer_metadata = metadata.indexer_metadata
|
indexer_metadata = metadata.indexer_metadata
|
||||||
core_metadata = metadata.core_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)
|
assert isinstance(indexer_metadata, PagedIndexerMetadata)
|
||||||
|
|
||||||
if enable_multi_stream:
|
if enable_multi_stream:
|
||||||
@@ -374,7 +368,7 @@ class C4IndexerBackendMixin:
|
|||||||
assert len(weights.shape) == 3
|
assert len(weights.shape) == 3
|
||||||
weights = weights.squeeze(2)
|
weights = weights.squeeze(2)
|
||||||
if envs.SGLANG_OPT_USE_TILELANG_INDEXER.get():
|
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,
|
tilelang_fp8_paged_mqa_logits as fn,
|
||||||
)
|
)
|
||||||
elif envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get():
|
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
|
from deep_gemm import fp8_paged_mqa_logits as fn
|
||||||
|
|
||||||
_c4sl = indexer_metadata.c4_seq_lens
|
_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)
|
_c4sl = _c4sl.unsqueeze(-1)
|
||||||
logits = fn(
|
logits = fn(
|
||||||
q_fp8,
|
q_fp8,
|
||||||
@@ -479,6 +474,7 @@ class C4Indexer(nn.Module):
|
|||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
alt_streams: Optional[List[torch.cuda.Stream]] = None,
|
alt_streams: Optional[List[torch.cuda.Stream]] = None,
|
||||||
|
rotary_emb=None,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
@@ -486,6 +482,7 @@ class C4Indexer(nn.Module):
|
|||||||
self.n_heads = config.index_n_heads
|
self.n_heads = config.index_n_heads
|
||||||
self.head_dim = config.index_head_dim
|
self.head_dim = config.index_head_dim
|
||||||
self.rope_head_dim = config.qk_rope_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.q_lora_rank = config.q_lora_rank
|
||||||
self.softmax_scale = self.head_dim**-0.5
|
self.softmax_scale = self.head_dim**-0.5
|
||||||
self.n_local_heads = self.n_heads
|
self.n_local_heads = self.n_heads
|
||||||
@@ -514,7 +511,9 @@ class C4Indexer(nn.Module):
|
|||||||
head_dim=self.head_dim,
|
head_dim=self.head_dim,
|
||||||
rotate=True,
|
rotate=True,
|
||||||
prefix=add_prefix("compressor", prefix),
|
prefix=add_prefix("compressor", prefix),
|
||||||
|
rotary_emb=rotary_emb,
|
||||||
)
|
)
|
||||||
|
self.rotary_emb = rotary_emb
|
||||||
self.freqs_cis = freqs_cis
|
self.freqs_cis = freqs_cis
|
||||||
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
|
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
|
||||||
self.alt_streams = alt_streams
|
self.alt_streams = alt_streams
|
||||||
@@ -545,8 +544,6 @@ class C4Indexer(nn.Module):
|
|||||||
enable_multi_stream: bool = False,
|
enable_multi_stream: bool = False,
|
||||||
q_lora_ready: Optional[torch.cuda.Event] = None,
|
q_lora_ready: Optional[torch.cuda.Event] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
if TYPE_CHECKING:
|
|
||||||
assert isinstance(forward_batch.attn_backend, DeepseekV4AttnBackend)
|
|
||||||
return forward_batch.attn_backend.forward_c4_indexer(
|
return forward_batch.attn_backend.forward_c4_indexer(
|
||||||
x=x,
|
x=x,
|
||||||
q_lora=q_lora,
|
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 (
|
assert (
|
||||||
page_size % 16 == 0
|
page_size % 16 == 0
|
||||||
), f"HIP preshuffle requires page_size to be a multiple of 16, got {page_size}"
|
), 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:
|
else:
|
||||||
assert page_size == 64
|
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
|
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,
|
scoring_func=scoring_func,
|
||||||
)
|
)
|
||||||
elif custom_routing_function is None:
|
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":
|
if scoring_func == "sqrtsoftplus":
|
||||||
_biased_topk = (
|
_biased_topk = (
|
||||||
biased_topk_jit_kernel_impl
|
biased_topk_jit_kernel_impl
|
||||||
|
|||||||
@@ -84,6 +84,7 @@ from sglang.srt.utils import (
|
|||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_musa,
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
@@ -111,9 +112,21 @@ _is_cpu = is_cpu()
|
|||||||
_is_fp8_fnuz = is_fp8_fnuz()
|
_is_fp8_fnuz = is_fp8_fnuz()
|
||||||
_use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip
|
_use_hip_int4 = get_bool_env_var("SGLANG_INT4_WEIGHT") and _is_hip
|
||||||
_use_aiter = envs.SGLANG_USE_AITER.get() 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:
|
if _use_aiter or _use_hip_int4:
|
||||||
from aiter.ops.shuffle import shuffle_weight
|
from aiter.ops.shuffle import shuffle_weight
|
||||||
|
from aiter.utility.fp4_utils import e8m0_shuffle
|
||||||
|
|
||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
from sglang.srt.layers.quantization.fp8_utils import (
|
from sglang.srt.layers.quantization.fp8_utils import (
|
||||||
@@ -998,12 +1011,13 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
# WEIGHT_SCALES
|
# WEIGHT_SCALES
|
||||||
if self.is_fp4_expert:
|
if self.is_fp4_expert:
|
||||||
fp4_block_k = 32
|
fp4_block_k = 32
|
||||||
|
fp4_scale_dtype = torch.float8_e8m0fnu if _use_aiter else torch.float32
|
||||||
w13_weight_scale = torch.nn.Parameter(
|
w13_weight_scale = torch.nn.Parameter(
|
||||||
torch.ones(
|
torch.ones(
|
||||||
num_experts,
|
num_experts,
|
||||||
2 * intermediate_size_per_partition,
|
2 * intermediate_size_per_partition,
|
||||||
hidden_size // fp4_block_k,
|
hidden_size // fp4_block_k,
|
||||||
dtype=torch.float32,
|
dtype=fp4_scale_dtype,
|
||||||
),
|
),
|
||||||
requires_grad=False,
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
@@ -1012,7 +1026,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
num_experts,
|
num_experts,
|
||||||
hidden_size,
|
hidden_size,
|
||||||
intermediate_size_per_partition // fp4_block_k,
|
intermediate_size_per_partition // fp4_block_k,
|
||||||
dtype=torch.float32,
|
dtype=fp4_scale_dtype,
|
||||||
),
|
),
|
||||||
requires_grad=False,
|
requires_grad=False,
|
||||||
)
|
)
|
||||||
@@ -1123,6 +1137,105 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.w2_input_scale = None
|
layer.w2_input_scale = None
|
||||||
|
|
||||||
def process_weights_after_loading_block_quant(self, layer: Module) -> 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 ROCm, normalize the weights and scales to e4m3fnuz
|
||||||
if _is_fp8_fnuz:
|
if _is_fp8_fnuz:
|
||||||
# activation_scheme: dynamic
|
# activation_scheme: dynamic
|
||||||
@@ -1148,8 +1261,6 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
)
|
)
|
||||||
layer.w2_input_scale = None
|
layer.w2_input_scale = None
|
||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
# add this section for MI300
|
|
||||||
# Pre-shuffle weights
|
|
||||||
layer.w13_weight.data = shuffle_weight(
|
layer.w13_weight.data = shuffle_weight(
|
||||||
layer.w13_weight.contiguous(), (16, 16)
|
layer.w13_weight.contiguous(), (16, 16)
|
||||||
)
|
)
|
||||||
@@ -1158,12 +1269,12 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
)
|
)
|
||||||
elif _use_aiter:
|
elif _use_aiter:
|
||||||
# Pre-shuffle weights
|
# Pre-shuffle weights
|
||||||
layer.w13_weight.data = shuffle_weight(
|
t = shuffle_weight(layer.w13_weight, (16, 16))
|
||||||
layer.w13_weight.contiguous(), (16, 16)
|
layer.w13_weight.copy_(t)
|
||||||
)
|
del t
|
||||||
layer.w2_weight.data = shuffle_weight(
|
t = shuffle_weight(layer.w2_weight, (16, 16))
|
||||||
layer.w2_weight.contiguous(), (16, 16)
|
layer.w2_weight.copy_(t)
|
||||||
)
|
del t
|
||||||
elif _is_cpu:
|
elif _is_cpu:
|
||||||
assert (
|
assert (
|
||||||
_is_cpu_amx_available
|
_is_cpu_amx_available
|
||||||
@@ -1190,8 +1301,9 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.w2_weight.data = layer.w2_weight.data.view(torch.int8)
|
layer.w2_weight.data = layer.w2_weight.data.view(torch.int8)
|
||||||
return
|
return
|
||||||
|
|
||||||
layer.w13_weight.data = layer.w13_weight.data.view(torch.int8)
|
fp4_weight_dtype = _require_fp4_dtype() if _use_aiter else torch.int8
|
||||||
layer.w2_weight.data = layer.w2_weight.data.view(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():
|
if get_moe_a2a_backend().is_megamoe():
|
||||||
from sglang.srt.layers.moe.mega_moe import (
|
from sglang.srt.layers.moe.mega_moe import (
|
||||||
@@ -1930,8 +2042,23 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
AiterQuantType,
|
AiterQuantType,
|
||||||
)
|
)
|
||||||
|
|
||||||
if _use_aiter and self.block_quant:
|
w13_weight = layer.w13_weight
|
||||||
quant_type = AiterQuantType.PER_128X128
|
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
|
w13_scale = layer.w13_weight_scale_inv
|
||||||
w2_scale = layer.w2_weight_scale_inv
|
w2_scale = layer.w2_weight_scale_inv
|
||||||
else:
|
else:
|
||||||
@@ -1939,8 +2066,8 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
|||||||
w13_scale = layer.w13_weight_scale1
|
w13_scale = layer.w13_weight_scale1
|
||||||
w2_scale = layer.w2_weight_scale1
|
w2_scale = layer.w2_weight_scale1
|
||||||
return AiterMoeQuantInfo(
|
return AiterMoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=w13_weight,
|
||||||
w2_weight=layer.w2_weight,
|
w2_weight=w2_weight,
|
||||||
quant_type=quant_type,
|
quant_type=quant_type,
|
||||||
w13_scale=w13_scale,
|
w13_scale=w13_scale,
|
||||||
w2_scale=w2_scale,
|
w2_scale=w2_scale,
|
||||||
|
|||||||
@@ -7,8 +7,11 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
|
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.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
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
|
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class KVAndScore:
|
class KVAndScore:
|
||||||
@@ -22,16 +25,55 @@ class KVAndScore:
|
|||||||
def score(self) -> torch.Tensor:
|
def score(self) -> torch.Tensor:
|
||||||
return self.kv_score[..., self._item_size :]
|
return self.kv_score[..., self._item_size :]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def shape(self):
|
||||||
|
return self.kv_score.shape
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
self._item_size = self.kv_score.shape[-1] // 2
|
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:
|
def __getitem__(self, index) -> KVAndScore:
|
||||||
return KVAndScore(self.kv_score[index])
|
return KVAndScore(self.kv_score[index])
|
||||||
|
|
||||||
|
def __setitem__(self, index, value: KVAndScore):
|
||||||
|
self.kv_score[index] = value.kv_score
|
||||||
|
|
||||||
def clear(self):
|
def clear(self):
|
||||||
self.kv.zero_()
|
self.kv.zero_()
|
||||||
self.score.fill_(float("-inf"))
|
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:
|
class CompressStatePool:
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -45,8 +87,11 @@ class CompressStatePool:
|
|||||||
enable_memory_saver: bool,
|
enable_memory_saver: bool,
|
||||||
ratio: int,
|
ratio: int,
|
||||||
online: bool = False,
|
online: bool = False,
|
||||||
|
swa_page_size: int = 0,
|
||||||
):
|
):
|
||||||
self.ring_size = ring_size
|
self.ring_size = ring_size
|
||||||
|
self.swa_page_size = swa_page_size
|
||||||
|
self.enable_memory_saver = enable_memory_saver
|
||||||
|
|
||||||
if online:
|
if online:
|
||||||
assert ring_size == 1, "online compress requires ring_size=1"
|
assert ring_size == 1, "online compress requires ring_size=1"
|
||||||
@@ -57,25 +102,47 @@ class CompressStatePool:
|
|||||||
self._size = (self._size + ratio - 1) // ratio * ratio
|
self._size = (self._size + ratio - 1) // ratio * ratio
|
||||||
last_dim = 2 * (1 + overlap) * head_dim
|
last_dim = 2 * (1 + overlap) * head_dim
|
||||||
|
|
||||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
if _is_hip:
|
||||||
enable=enable_memory_saver
|
self.kv_score_buffer = KVAndScore(
|
||||||
)
|
torch.empty((self._size, last_dim), dtype=dtype, device=device)
|
||||||
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
|
)
|
||||||
maybe_init_custom_mem_pool(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 self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
|
||||||
with (
|
with (
|
||||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||||
if self.custom_mem_pool
|
if self.custom_mem_pool
|
||||||
else nullcontext()
|
else nullcontext()
|
||||||
):
|
):
|
||||||
self.kv_score_buffer = KVAndScore(
|
self.kv_score_buffer = KVAndScore(
|
||||||
torch.empty(
|
torch.empty(
|
||||||
(self._size, last_dim),
|
(self._size, last_dim),
|
||||||
dtype=dtype,
|
dtype=dtype,
|
||||||
device=device,
|
device=device,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
)
|
if not online:
|
||||||
if not online:
|
self.kv_score_buffer[-1].clear()
|
||||||
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.deepseek_v4_compress_state import CompressStatePool
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache
|
from sglang.srt.mem_cache.memory_pool import KVCache
|
||||||
from sglang.srt.server_args import get_global_server_args
|
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__)
|
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(
|
def get_compress_state_ring_size(
|
||||||
@@ -144,6 +146,9 @@ class DeepSeekV4SingleKVPool(KVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def get_key_buffer(self, layer_id: int):
|
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]
|
return self.kv_buffer[layer_id]
|
||||||
|
|
||||||
def set_kv_buffer(self, *args, **kwargs) -> None:
|
def set_kv_buffer(self, *args, **kwargs) -> None:
|
||||||
@@ -466,7 +471,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.c4_indexer_kv_pool = DeepSeekV4IndexerPool(
|
self.c4_indexer_kv_pool = DeepSeekV4IndexerPool(
|
||||||
self.c4_logical_size,
|
self.c4_logical_size if not _is_hip else c4_size,
|
||||||
c4_page_size,
|
c4_page_size,
|
||||||
dtype,
|
dtype,
|
||||||
indexer_head_dim,
|
indexer_head_dim,
|
||||||
@@ -477,7 +482,10 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
|
|
||||||
self._init_compressed_layer_mapping()
|
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._should_cache_swa = envs.SGLANG_OPT_CACHE_SWA_TRANSLATION.get()
|
||||||
self.cached_loc = None
|
self.cached_loc = None
|
||||||
@@ -585,6 +593,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
dtype=self.state_dtype,
|
dtype=self.state_dtype,
|
||||||
enable_memory_saver=enable_memory_saver,
|
enable_memory_saver=enable_memory_saver,
|
||||||
ratio=ratio,
|
ratio=ratio,
|
||||||
|
swa_page_size=self.swa_page_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_compressed_layer_mapping(self):
|
def _init_compressed_layer_mapping(self):
|
||||||
|
|||||||
@@ -575,6 +575,11 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
use_grouped_topk=False,
|
use_grouped_topk=False,
|
||||||
scoring_func=config.scoring_func,
|
scoring_func=config.scoring_func,
|
||||||
is_fp4_experts=getattr(quant_config, "is_fp4_experts", False),
|
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)
|
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 import get_moe_a2a_backend
|
||||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
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.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 import PPMissingLayer, get_layer_id
|
||||||
from sglang.srt.layers.utils.cp_utils import (
|
from sglang.srt.layers.utils.cp_utils import (
|
||||||
cp_all_gather_rerange_output,
|
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.model_loader.weight_utils import default_weight_loader
|
||||||
from sglang.srt.models.dbrx import ReplicatedLinear
|
from sglang.srt.models.dbrx import ReplicatedLinear
|
||||||
from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu
|
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.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
LazyValue,
|
LazyValue,
|
||||||
@@ -94,6 +101,9 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
||||||
DeepseekV4AttnBackend,
|
DeepseekV4AttnBackend,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||||
|
DeepseekV4HipRadixBackend,
|
||||||
|
)
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
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
|
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
|
from sglang.srt.layers.deepseek_v4_rope import precompute_freqs_cis
|
||||||
|
|
||||||
assert self.compress_ratio in {0, 4, 128}
|
assert self.compress_ratio in {0, 4, 128}
|
||||||
@@ -243,6 +263,7 @@ class MQALayer(nn.Module):
|
|||||||
head_dim=self.head_dim,
|
head_dim=self.head_dim,
|
||||||
rotate=False,
|
rotate=False,
|
||||||
prefix=add_prefix("compressor", prefix),
|
prefix=add_prefix("compressor", prefix),
|
||||||
|
rotary_emb=getattr(self, "rotary_emb", None),
|
||||||
)
|
)
|
||||||
if self.compress_ratio == 4:
|
if self.compress_ratio == 4:
|
||||||
self.indexer = C4Indexer(
|
self.indexer = C4Indexer(
|
||||||
@@ -252,10 +273,11 @@ class MQALayer(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("indexer", prefix),
|
prefix=add_prefix("indexer", prefix),
|
||||||
alt_streams=self.alt_streams_indexer,
|
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.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:
|
if self.fuse_wqa_wkv:
|
||||||
self.wqkv_a = ReplicatedLinear(
|
self.wqkv_a = ReplicatedLinear(
|
||||||
self.hidden_size,
|
self.hidden_size,
|
||||||
@@ -409,7 +431,7 @@ class MQALayer(nn.Module):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
attn_backend: DeepseekV4AttnBackend,
|
attn_backend,
|
||||||
q_out: Optional[torch.Tensor] = None,
|
q_out: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
assert self.alt_streams is not None
|
assert self.alt_streams is not None
|
||||||
@@ -469,7 +491,7 @@ class MQALayer(nn.Module):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
positions: torch.Tensor,
|
positions: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
attn_backend: DeepseekV4AttnBackend,
|
attn_backend,
|
||||||
q_out: Optional[torch.Tensor] = None,
|
q_out: Optional[torch.Tensor] = None,
|
||||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||||
if self.fuse_wqa_wkv:
|
if self.fuse_wqa_wkv:
|
||||||
@@ -530,7 +552,10 @@ class MQALayer(nn.Module):
|
|||||||
|
|
||||||
attn_backend = forward_batch.attn_backend
|
attn_backend = forward_batch.attn_backend
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
assert isinstance(attn_backend, DeepseekV4AttnBackend)
|
assert isinstance(
|
||||||
|
attn_backend,
|
||||||
|
(DeepseekV4AttnBackend, DeepseekV4HipRadixBackend),
|
||||||
|
)
|
||||||
|
|
||||||
enable_multi_stream = (
|
enable_multi_stream = (
|
||||||
envs.SGLANG_OPT_USE_MULTI_STREAM_OVERLAP.get()
|
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
|
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():
|
if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get():
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
|
|
||||||
@@ -765,6 +806,13 @@ class DeepseekV4DecoderLayer(nn.Module):
|
|||||||
|
|
||||||
return mhc_post(x, residual, post, comb)
|
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 residual.shape == (x.shape[0], self.hc_mult, x.shape[-1])
|
||||||
assert post.shape == (x.shape[0], self.hc_mult)
|
assert post.shape == (x.shape[0], self.hc_mult)
|
||||||
assert comb.shape == (x.shape[0], self.hc_mult, 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 = {}
|
cache_compressor_weight = {}
|
||||||
COMPRESSOR_PART = ".compressor.w"
|
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]] = {}
|
cache_wqkv_a_weight: dict[str, dict[str, torch.Tensor]] = {}
|
||||||
|
|
||||||
def auto_weight_loader(module):
|
def auto_weight_loader(module):
|
||||||
|
|||||||
Reference in New Issue
Block a user