[2/n] deepseek_v2.py Refactor: Migrate MHA forward method in deepseek_v2.py (#16817)
This commit is contained in:
@@ -56,7 +56,6 @@ from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.amx_utils import PackWeightMethod
|
||||
from sglang.srt.layers.attention.nsa.dequant_k_cache import dequantize_k_cache_paged
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
from sglang.srt.layers.attention.nsa.utils import (
|
||||
can_cp_split,
|
||||
@@ -67,8 +66,6 @@ from sglang.srt.layers.attention.nsa.utils import (
|
||||
nsa_use_prefill_cp,
|
||||
prepare_input_dp_with_cp_dsa,
|
||||
)
|
||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||
from sglang.srt.layers.attention.utils import concat_and_cast_mha_k_triton
|
||||
from sglang.srt.layers.communicator import (
|
||||
LayerCommunicator,
|
||||
LayerScatterModes,
|
||||
@@ -142,8 +139,9 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.deepseek_common.attention_backend_handler import (
|
||||
AttentionBackendRegistry,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import (
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods import (
|
||||
AttnForwardMethod,
|
||||
DeepseekMHAForwardMixin,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.utils import (
|
||||
_device_sm,
|
||||
@@ -195,14 +193,7 @@ if _use_aiter_gfx95:
|
||||
)
|
||||
|
||||
if _is_cuda:
|
||||
from sgl_kernel import (
|
||||
awq_dequantize,
|
||||
bmm_fp8,
|
||||
concat_mla_k,
|
||||
dsv3_fused_a_gemm,
|
||||
dsv3_router_gemm,
|
||||
merge_state_v2,
|
||||
)
|
||||
from sgl_kernel import awq_dequantize, bmm_fp8, dsv3_fused_a_gemm, dsv3_router_gemm
|
||||
elif _is_cpu and _is_cpu_amx_available:
|
||||
pass
|
||||
elif _is_hip:
|
||||
@@ -255,12 +246,6 @@ FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [
|
||||
]
|
||||
|
||||
|
||||
def add_forward_absorb_core_attention_backend(backend_name):
|
||||
if backend_name not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS:
|
||||
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS.append(backend_name)
|
||||
logger.info(f"Added {backend_name} to FORWARD_ABSORB_CORE_ATTENTION_BACKENDS.")
|
||||
|
||||
|
||||
class DeepseekV2MLP(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
@@ -1074,7 +1059,7 @@ def _get_llama_4_scaling(
|
||||
return scaling[..., None, None]
|
||||
|
||||
|
||||
class DeepseekV2AttentionMLA(nn.Module):
|
||||
class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -1264,9 +1249,6 @@ class DeepseekV2AttentionMLA(nn.Module):
|
||||
self.flashinfer_mla_disable_ragged = (
|
||||
get_global_server_args().flashinfer_mla_disable_ragged
|
||||
)
|
||||
self.disable_chunked_prefix_cache = (
|
||||
get_global_server_args().disable_chunked_prefix_cache
|
||||
)
|
||||
|
||||
self.current_attention_backend = (
|
||||
None # Attention backend used by current forward batch
|
||||
@@ -1275,11 +1257,6 @@ class DeepseekV2AttentionMLA(nn.Module):
|
||||
"SGLANG_ROCM_FUSED_DECODE_MLA", "false"
|
||||
)
|
||||
|
||||
# TODO: Design a finer way to determine the threshold
|
||||
self.chunked_prefix_cache_threshold = (
|
||||
envs.SGLANG_CHUNKED_PREFIX_CACHE_THRESHOLD.get()
|
||||
)
|
||||
|
||||
# If we have self.fused_qkv_a_proj_with_mqa and we're running on CPU, we will choose the torch.ops.sgl_kernel.qkv_proj_with_rope_fused_weight kernel
|
||||
# which requires self.w_kc and self.w_vc to be packed.
|
||||
# If not, we will use torch.bmm and weight shouldn't be packed in this case
|
||||
@@ -1334,6 +1311,8 @@ class DeepseekV2AttentionMLA(nn.Module):
|
||||
self.fused_qkv_a_proj_with_mqa.quant_method.quant_config.weight_block_size
|
||||
)
|
||||
|
||||
self.init_mha_forward()
|
||||
|
||||
def dispatch_attn_forward_method(
|
||||
self, forward_batch: ForwardBatch
|
||||
) -> AttnForwardMethod:
|
||||
@@ -1503,161 +1482,6 @@ class DeepseekV2AttentionMLA(nn.Module):
|
||||
qkv_latent = self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
|
||||
return qkv_latent
|
||||
|
||||
def forward_normal_prepare(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
zero_allocator: BumpAllocator,
|
||||
):
|
||||
if self.q_lora_rank is not None:
|
||||
q, latent_cache = (
|
||||
get_attn_tp_context()
|
||||
.fetch_qkv_latent()
|
||||
.split(
|
||||
[self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim],
|
||||
dim=-1,
|
||||
)
|
||||
)
|
||||
|
||||
# NSA Indexer: cache quantized keys, auto-skip topk for sequences <= nsa_index_topk
|
||||
|
||||
if self.use_nsa:
|
||||
# NSA requires unquantized q_lora for the indexer. When q_b_proj is FP8
|
||||
# on gfx95, we can still use fused RMSNorm+FP8 quant, but MUST request
|
||||
# the unquantized output for q_lora; otherwise q_lora becomes the (fp8,scale)
|
||||
# tuple.
|
||||
if (
|
||||
_use_aiter_gfx95
|
||||
and self.q_b_proj.weight.dtype == torch.float8_e4m3fn
|
||||
):
|
||||
q_quanted, q_lora, _, _ = fused_rms_fp8_group_quant(
|
||||
q,
|
||||
self.q_a_layernorm.weight,
|
||||
self.q_a_layernorm.variance_epsilon,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
group_size=128,
|
||||
dtype_quant=torch.float8_e4m3fn,
|
||||
res1=None,
|
||||
output_unquantized_inp1=True,
|
||||
)
|
||||
q = self.q_b_proj(q_quanted)[0].view(
|
||||
-1, self.num_local_heads, self.qk_head_dim
|
||||
)
|
||||
else:
|
||||
q_lora = self.q_a_layernorm(q)
|
||||
q = self.q_b_proj(q_lora)[0].view(
|
||||
-1, self.num_local_heads, self.qk_head_dim
|
||||
)
|
||||
_ = self.indexer(
|
||||
x=hidden_states,
|
||||
q_lora=q_lora,
|
||||
positions=positions,
|
||||
forward_batch=forward_batch,
|
||||
layer_id=self.layer_id,
|
||||
return_indices=False,
|
||||
)
|
||||
elif _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.uint8:
|
||||
# MXFP4: fused RMSNorm + quant
|
||||
q, _, _, _ = fused_rms_mxfp4_quant(
|
||||
q,
|
||||
self.q_a_layernorm.weight,
|
||||
self.q_a_layernorm.variance_epsilon,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim)
|
||||
elif _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.float8_e4m3fn:
|
||||
|
||||
q, _, _, _ = fused_rms_fp8_group_quant(
|
||||
q,
|
||||
self.q_a_layernorm.weight,
|
||||
self.q_a_layernorm.variance_epsilon,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
group_size=128,
|
||||
dtype_quant=torch.float8_e4m3fn,
|
||||
res1=None,
|
||||
output_unquantized_inp1=False,
|
||||
)
|
||||
q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim)
|
||||
else:
|
||||
q = self.q_a_layernorm(q)
|
||||
q = self.q_b_proj(q)[0].view(-1, self.num_local_heads, self.qk_head_dim)
|
||||
|
||||
else:
|
||||
q = self.q_proj(hidden_states)[0].view(
|
||||
-1, self.num_local_heads, self.qk_head_dim
|
||||
)
|
||||
latent_cache = self.kv_a_proj_with_mqa(hidden_states)[0]
|
||||
|
||||
_, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
|
||||
kv_a, _ = latent_cache.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
|
||||
latent_cache = latent_cache.unsqueeze(1)
|
||||
|
||||
if _use_aiter_gfx95 and self.kv_b_proj.weight.dtype == torch.float8_e4m3fn:
|
||||
|
||||
kv_a_quanted, kv_a, _, _ = fused_rms_fp8_group_quant(
|
||||
kv_a,
|
||||
self.kv_a_layernorm.weight,
|
||||
self.kv_a_layernorm.variance_epsilon,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
group_size=128,
|
||||
dtype_quant=torch.float8_e4m3fn,
|
||||
res1=None,
|
||||
output_unquantized_inp1=True, # return unqaunt kv_a
|
||||
)
|
||||
|
||||
else:
|
||||
kv_a = self.kv_a_layernorm(kv_a)
|
||||
|
||||
# kv_a = self.kv_a_layernorm(kv_a)
|
||||
|
||||
k_pe = latent_cache[:, :, self.kv_lora_rank :]
|
||||
if self.rotary_emb is not None:
|
||||
q_pe, k_pe = self.rotary_emb(positions, q_pe, k_pe)
|
||||
q[..., self.qk_nope_head_dim :] = q_pe
|
||||
|
||||
self._set_mla_kv_buffer(latent_cache, kv_a, k_pe, forward_batch)
|
||||
if (
|
||||
forward_batch.mha_one_shot
|
||||
and sum(forward_batch.extend_prefix_lens_cpu) != 0
|
||||
):
|
||||
if self.use_nsa and self.kv_cache_dtype == "fp8_e4m3":
|
||||
# FP8 path: dequantize NSA-specific FP8 format to BF16
|
||||
kv_a, k_pe = self._get_mla_kv_buffer_from_fp8(forward_batch)
|
||||
else:
|
||||
# BF16/FP16 path: directly fetch from cache
|
||||
kv_a, k_pe = self._get_mla_kv_buffer(
|
||||
forward_batch.fetch_mha_one_shot_kv_indices(),
|
||||
q.dtype,
|
||||
forward_batch,
|
||||
)
|
||||
if _use_aiter_gfx95 and self.kv_b_proj.weight.dtype == torch.float8_e4m3fn:
|
||||
kv = self.kv_b_proj(
|
||||
kv_a_quanted,
|
||||
)[0]
|
||||
else:
|
||||
kv = self.kv_b_proj(kv_a)[0]
|
||||
kv = kv.view(-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim)
|
||||
k_nope = kv[..., : self.qk_nope_head_dim]
|
||||
v = kv[..., self.qk_nope_head_dim :]
|
||||
|
||||
k = self._concat_and_cast_mha_k(k_nope, k_pe, forward_batch)
|
||||
return q, k, v, forward_batch
|
||||
|
||||
def forward_normal_core(self, q, k, v, forward_batch):
|
||||
attn_output = self.attn_mha(q, k, v, forward_batch, save_kv_cache=False)
|
||||
attn_output = attn_output.reshape(-1, self.num_local_heads * self.v_head_dim)
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
|
||||
def _fuse_rope_for_trtllm_mla(self, forward_batch: ForwardBatch) -> bool:
|
||||
"""
|
||||
Check if we should skip rope and do fused rope+quantize for TRTLLM MLA decode in fp8_e4m3 path.
|
||||
@@ -2377,231 +2201,6 @@ class DeepseekV2AttentionMLA(nn.Module):
|
||||
|
||||
return output
|
||||
|
||||
def _chunked_prefix_attn_mha(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
accum_output: torch.Tensor,
|
||||
accum_lse: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
) -> torch.Tensor:
|
||||
|
||||
assert forward_batch.num_prefix_chunks is not None
|
||||
for i in range(forward_batch.num_prefix_chunks):
|
||||
forward_batch.set_prefix_chunk_idx(i)
|
||||
|
||||
kv_indices = forward_batch.prefix_chunk_kv_indices[i]
|
||||
# Fetch latent cache from memory pool with precomputed chunked kv indices
|
||||
kv_a_normed, k_pe = self._get_mla_kv_buffer(
|
||||
kv_indices, q.dtype, forward_batch
|
||||
)
|
||||
kv = self.kv_b_proj(kv_a_normed)[0]
|
||||
kv = kv.view(
|
||||
-1, self.num_local_heads, self.qk_nope_head_dim + self.v_head_dim
|
||||
)
|
||||
v = kv[..., self.qk_nope_head_dim :]
|
||||
k_nope = kv[..., : self.qk_nope_head_dim]
|
||||
|
||||
k = torch.empty(
|
||||
(
|
||||
k_nope.shape[0],
|
||||
self.num_local_heads,
|
||||
self.qk_nope_head_dim + self.qk_rope_head_dim,
|
||||
),
|
||||
dtype=v.dtype,
|
||||
device=v.device,
|
||||
)
|
||||
k[..., : self.qk_nope_head_dim] = k_nope
|
||||
k[..., self.qk_nope_head_dim :] = k_pe
|
||||
|
||||
output, lse = self.attn_mha(q, k, v, forward_batch, save_kv_cache=False)
|
||||
tmp_output = torch.empty_like(accum_output)
|
||||
tmp_lse = torch.empty_like(accum_lse)
|
||||
merge_state_v2(output, lse, accum_output, accum_lse, tmp_output, tmp_lse)
|
||||
accum_output, accum_lse = tmp_output, tmp_lse
|
||||
del kv, k, v, output, lse, tmp_output, tmp_lse
|
||||
|
||||
return accum_output
|
||||
|
||||
def forward_normal_chunked_kv_prepare(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
zero_allocator: BumpAllocator,
|
||||
):
|
||||
# In normal mha, the k and v tensors will become overly large when the prefix length is long.
|
||||
# To avoid this, we split the kv cache into chunks and process them one after another.
|
||||
# Since mha is compute friendly, the for loop induced here will not introduce significant overhead.
|
||||
# The top comments in https://github.com/vllm-project/vllm/blob/main/vllm/v1/attention/backends/mla/common.py
|
||||
# will be helpful for understanding the purpose of this function.
|
||||
|
||||
# First do normal mha forward to get output for extended part
|
||||
return self.forward_normal_prepare(
|
||||
positions, hidden_states, forward_batch, zero_allocator
|
||||
)
|
||||
|
||||
def forward_normal_chunked_kv_core(self, q, k, v, forward_batch):
|
||||
has_extend_prefix = forward_batch.extend_prefix_lens_cpu is not None and any(
|
||||
forward_batch.extend_prefix_lens_cpu
|
||||
)
|
||||
# Only initialize the info once
|
||||
if has_extend_prefix and forward_batch.num_prefix_chunks is None:
|
||||
forward_batch.prepare_chunked_prefix_cache_info(q.device)
|
||||
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"):
|
||||
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch)
|
||||
|
||||
forward_batch.mha_return_lse = has_extend_prefix
|
||||
# Do mha for extended part without prefix
|
||||
forward_batch.set_attn_attend_prefix_cache(False)
|
||||
attn_output = self.attn_mha(q, k, v, forward_batch, save_kv_cache=False)
|
||||
|
||||
# Do mha attention with chunked prefix cache if there are any sequence with prefix
|
||||
if has_extend_prefix:
|
||||
attn_output, lse = attn_output
|
||||
forward_batch.set_attn_attend_prefix_cache(True)
|
||||
attn_output = self._chunked_prefix_attn_mha(
|
||||
q=q,
|
||||
accum_output=attn_output,
|
||||
accum_lse=lse,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
|
||||
attn_output = attn_output.reshape(-1, self.num_local_heads * self.v_head_dim)
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
|
||||
def forward_normal_one_shot_prepare(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
hidden_states: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
zero_allocator: BumpAllocator,
|
||||
):
|
||||
forward_batch.mha_one_shot = True
|
||||
return self.forward_normal_prepare(
|
||||
positions, hidden_states, forward_batch, zero_allocator
|
||||
)
|
||||
|
||||
def forward_normal_one_shot_core(self, q, k, v, forward_batch):
|
||||
has_extend_prefix = any(forward_batch.extend_prefix_lens_cpu)
|
||||
# Only initialize the info once
|
||||
if has_extend_prefix and forward_batch.num_prefix_chunks is None:
|
||||
forward_batch.num_prefix_chunks = 0
|
||||
if hasattr(forward_batch.attn_backend, "init_mha_chunk_metadata"):
|
||||
forward_batch.attn_backend.init_mha_chunk_metadata(forward_batch)
|
||||
forward_batch.mha_return_lse = False
|
||||
# Do mha for extended part without prefix
|
||||
forward_batch.set_attn_attend_prefix_cache(False)
|
||||
return self.forward_normal_core(q, k, v, forward_batch)
|
||||
|
||||
def _set_mla_kv_buffer(
|
||||
self,
|
||||
latent_cache: torch.Tensor,
|
||||
kv_a: torch.Tensor,
|
||||
k_pe: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
if _is_cuda or _use_aiter_gfx95:
|
||||
# Save latent cache
|
||||
forward_batch.token_to_kv_pool.set_mla_kv_buffer(
|
||||
self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe
|
||||
)
|
||||
elif _is_npu:
|
||||
# To reduce a time-costing split operation
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
self.attn_mha, forward_batch.out_cache_loc, kv_a.unsqueeze(1), k_pe
|
||||
)
|
||||
else:
|
||||
latent_cache[:, :, : self.kv_lora_rank] = kv_a.unsqueeze(1)
|
||||
latent_cache[:, :, self.kv_lora_rank :] = k_pe
|
||||
|
||||
# Save latent cache
|
||||
forward_batch.token_to_kv_pool.set_kv_buffer(
|
||||
self.attn_mha, forward_batch.out_cache_loc, latent_cache, None
|
||||
)
|
||||
|
||||
def _get_mla_kv_buffer(
|
||||
self,
|
||||
kv_indices: torch.Tensor,
|
||||
dst_dtype: torch.dtype,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
if _is_cuda or _use_aiter_gfx95:
|
||||
kv_a, k_pe = forward_batch.token_to_kv_pool.get_mla_kv_buffer(
|
||||
self.attn_mha, kv_indices, dst_dtype
|
||||
)
|
||||
kv_a = kv_a.squeeze(1)
|
||||
else:
|
||||
latent_cache_buf = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mha.layer_id
|
||||
)
|
||||
latent_cache = latent_cache_buf[kv_indices].contiguous().to(dst_dtype)
|
||||
|
||||
kv_a, k_pe = latent_cache.split(
|
||||
[self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
|
||||
)
|
||||
kv_a = kv_a.squeeze(1).contiguous()
|
||||
return kv_a, k_pe
|
||||
|
||||
def _get_mla_kv_buffer_from_fp8(
|
||||
self,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
"""
|
||||
Dequantize FP8 KV cache to BF16 for MLA attention (NSA-specific format).
|
||||
|
||||
Returns: (kv_a, k_pe) both in BF16
|
||||
"""
|
||||
backend = forward_batch.attn_backend
|
||||
if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend
|
||||
backend = backend.primary
|
||||
kv_indices = backend.forward_metadata.page_table_1_flattened
|
||||
assert (
|
||||
kv_indices is not None
|
||||
), "page_table_1_flattened should have been generated for FP8 MHA path"
|
||||
|
||||
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_key_buffer(
|
||||
self.attn_mha.layer_id
|
||||
)
|
||||
|
||||
kv_latent_bf16 = dequantize_k_cache_paged(kv_cache_fp8, kv_indices)
|
||||
|
||||
kv_a = kv_latent_bf16[:, :, : self.kv_lora_rank].squeeze(1).contiguous()
|
||||
k_pe = kv_latent_bf16[:, :, self.kv_lora_rank :]
|
||||
|
||||
return kv_a, k_pe
|
||||
|
||||
def _concat_and_cast_mha_k(self, k_nope, k_pe, forward_batch):
|
||||
# Temporary for DeepSeek V3/R1 only, but can generalize if needed
|
||||
k_shape = (k_nope.shape[0], self.num_local_heads, self.qk_head_dim)
|
||||
if (
|
||||
_is_cuda
|
||||
and (self.num_local_heads == 128)
|
||||
and (self.qk_nope_head_dim == 128)
|
||||
and (self.qk_rope_head_dim == 64)
|
||||
):
|
||||
k = k_nope.new_empty(*k_shape)
|
||||
concat_mla_k(k=k, k_nope=k_nope, k_rope=k_pe)
|
||||
elif _is_cuda:
|
||||
# fa3 mha support fp8 inputs
|
||||
if (
|
||||
self.current_attention_backend == "fa3"
|
||||
and self.kv_cache_dtype != "auto"
|
||||
):
|
||||
attn_dtype = forward_batch.token_to_kv_pool.dtype
|
||||
else:
|
||||
attn_dtype = k_nope.dtype
|
||||
k = k_nope.new_empty(*k_shape, dtype=attn_dtype)
|
||||
concat_and_cast_mha_k_triton(k, k_nope, k_pe)
|
||||
elif _is_hip and self.current_attention_backend == "aiter":
|
||||
k = k_nope.new_empty(*k_shape)
|
||||
concat_and_cast_mha_k_triton(k, k_nope, k_pe)
|
||||
else:
|
||||
k = k_nope.new_empty(*k_shape)
|
||||
k[..., : self.qk_nope_head_dim] = k_nope
|
||||
k[..., self.qk_nope_head_dim :] = k_pe
|
||||
return k
|
||||
|
||||
@staticmethod
|
||||
def _get_q_b_proj_quant_config(quant_config):
|
||||
if envs.SGLANG_NVFP4_CKPT_FP8_GEMM_IN_ATTN.get():
|
||||
|
||||
Reference in New Issue
Block a user