Refactor: Extract DeepSeek common utilities into shared module (#16969)
This commit is contained in:
@@ -11,9 +11,15 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
import logging
|
||||||
|
import math
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import get_moe_runner_backend
|
from sglang.srt.layers.moe.fused_moe_triton.layer import get_moe_runner_backend
|
||||||
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -24,6 +30,7 @@ from sglang.srt.utils import (
|
|||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
|
is_nvidia_cublas_cu12_version_ge_12_9,
|
||||||
)
|
)
|
||||||
|
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
@@ -38,6 +45,22 @@ _is_gfx95_supported = is_gfx95_supported()
|
|||||||
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
|
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
|
||||||
|
|
||||||
|
|
||||||
|
_is_cublas_ge_129 = is_nvidia_cublas_cu12_version_ge_12_9()
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
NVFP4_CKPT_FP8_ATTN_QUANT_MODULES = ["q_b_proj"]
|
||||||
|
|
||||||
|
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [
|
||||||
|
"fa3",
|
||||||
|
"nsa",
|
||||||
|
"flashinfer",
|
||||||
|
"cutlass_mla",
|
||||||
|
"trtllm_mla",
|
||||||
|
"ascend",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def awq_dequantize_func():
|
def awq_dequantize_func():
|
||||||
"""
|
"""
|
||||||
Get the AWQ dequantize function for the current device
|
Get the AWQ dequantize function for the current device
|
||||||
@@ -66,10 +89,27 @@ def awq_dequantize_func():
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def enable_nextn_moe_bf16_cast_to_fp8(quant_config):
|
def enable_nextn_moe_bf16_cast_to_fp8(
|
||||||
|
quant_config: Optional[QuantizationConfig],
|
||||||
|
) -> bool:
|
||||||
return (
|
return (
|
||||||
envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE.get()
|
envs.SGLANG_NVFP4_CKPT_FP8_NEXTN_MOE.get()
|
||||||
and quant_config is not None
|
and quant_config is not None
|
||||||
and quant_config.get_name() == "modelopt_fp4"
|
and quant_config.get_name() == "modelopt_fp4"
|
||||||
and get_moe_runner_backend().is_deep_gemm()
|
and get_moe_runner_backend().is_deep_gemm()
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def yarn_get_mscale(scale: float = 1, mscale: float = 1) -> float:
|
||||||
|
if scale <= 1:
|
||||||
|
return 1.0
|
||||||
|
return 0.1 * mscale * math.log(scale) + 1.0
|
||||||
|
|
||||||
|
|
||||||
|
def _get_llama_4_scaling(
|
||||||
|
original_max_position_embeddings: int, scaling_beta: float, positions: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
scaling = 1 + scaling_beta * torch.log(
|
||||||
|
1 + torch.floor(positions / original_max_position_embeddings)
|
||||||
|
)
|
||||||
|
return scaling[..., None, None]
|
||||||
|
|||||||
@@ -128,15 +128,19 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
|
|||||||
DeepseekV2WeightLoaderMixin,
|
DeepseekV2WeightLoaderMixin,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.deepseek_common.utils import (
|
from sglang.srt.models.deepseek_common.utils import (
|
||||||
|
FORWARD_ABSORB_CORE_ATTENTION_BACKENDS,
|
||||||
_device_sm,
|
_device_sm,
|
||||||
|
_get_llama_4_scaling,
|
||||||
_is_cpu,
|
_is_cpu,
|
||||||
_is_cpu_amx_available,
|
_is_cpu_amx_available,
|
||||||
|
_is_cublas_ge_129,
|
||||||
_is_cuda,
|
_is_cuda,
|
||||||
_is_gfx95_supported,
|
_is_gfx95_supported,
|
||||||
_is_hip,
|
_is_hip,
|
||||||
_is_npu,
|
_is_npu,
|
||||||
_use_aiter,
|
_use_aiter,
|
||||||
_use_aiter_gfx95,
|
_use_aiter_gfx95,
|
||||||
|
yarn_get_mscale,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -146,7 +150,6 @@ from sglang.srt.utils import (
|
|||||||
add_prefix,
|
add_prefix,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
is_non_idle_and_non_empty,
|
is_non_idle_and_non_empty,
|
||||||
is_nvidia_cublas_cu12_version_ge_12_9,
|
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
make_layers,
|
make_layers,
|
||||||
use_intel_amx_backend,
|
use_intel_amx_backend,
|
||||||
@@ -193,8 +196,6 @@ elif _is_npu:
|
|||||||
else:
|
else:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
_is_cublas_ge_129 = is_nvidia_cublas_cu12_version_ge_12_9()
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -1049,24 +1050,6 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
state.hidden_states_mlp_output = final_hidden_states
|
state.hidden_states_mlp_output = final_hidden_states
|
||||||
|
|
||||||
|
|
||||||
def yarn_get_mscale(scale: float = 1, mscale: float = 1) -> float:
|
|
||||||
import math
|
|
||||||
|
|
||||||
if scale <= 1:
|
|
||||||
return 1.0
|
|
||||||
return 0.1 * mscale * math.log(scale) + 1.0
|
|
||||||
|
|
||||||
|
|
||||||
def _get_llama_4_scaling(
|
|
||||||
original_max_position_embeddings: int, scaling_beta: float, positions: torch.Tensor
|
|
||||||
) -> torch.Tensor:
|
|
||||||
scaling = 1 + scaling_beta * torch.log(
|
|
||||||
1 + torch.floor(positions / original_max_position_embeddings)
|
|
||||||
)
|
|
||||||
# Broadcast over num_heads and head_dim
|
|
||||||
return scaling[..., None, None]
|
|
||||||
|
|
||||||
|
|
||||||
class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin):
|
class DeepseekV2AttentionMLA(nn.Module, DeepseekMHAForwardMixin):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|||||||
Reference in New Issue
Block a user