From 38dc5839dd8d185b419be9e5bb2d22c2908db979 Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Thu, 8 Jan 2026 09:22:31 +0800 Subject: [PATCH] [1/n]deepseek_v2.py Refactor: attention backend handlers and forward method definition (#16306) --- .../srt/models/deepseek_common/__init__.py | 0 .../attention_backend_handler.py | 182 +++++++++++++ .../forward_methods.py | 32 +++ .../srt/models/deepseek_common/utils.py | 23 ++ python/sglang/srt/models/deepseek_v2.py | 246 ++---------------- 5 files changed, 255 insertions(+), 228 deletions(-) create mode 100644 python/sglang/srt/models/deepseek_common/__init__.py create mode 100644 python/sglang/srt/models/deepseek_common/attention_backend_handler.py create mode 100644 python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py create mode 100644 python/sglang/srt/models/deepseek_common/utils.py diff --git a/python/sglang/srt/models/deepseek_common/__init__.py b/python/sglang/srt/models/deepseek_common/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py new file mode 100644 index 000000000..9e99076e9 --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -0,0 +1,182 @@ +from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph +from sglang.srt.layers.attention.tbo_backend import TboAttnBackend +from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods import ( + AttnForwardMethod, +) +from sglang.srt.models.deepseek_common.utils import _is_hip +from sglang.srt.server_args import get_global_server_args +from sglang.srt.utils import use_intel_amx_backend + + +class AttentionBackendRegistry: + _handlers = {} + + @classmethod + def register(cls, backend_name, handler_func): + cls._handlers[backend_name] = handler_func + + @classmethod + def get_handler(cls, backend_name): + return cls._handlers.get(backend_name, cls._handlers.get("triton")) + + +def _dispatch_mla_subtype(attn, forward_batch): + if _is_hip: + if attn.rocm_fused_decode_mla and forward_batch.forward_mode.is_decode(): + return AttnForwardMethod.MLA_FUSED_ROPE + else: + return AttnForwardMethod.MLA + else: + if hasattr(attn, "fused_qkv_a_proj_with_mqa") and use_intel_amx_backend(attn): + return AttnForwardMethod.MLA_FUSED_ROPE_CPU + else: + return AttnForwardMethod.MLA + + +def handle_attention_ascend(attn, forward_batch): + if ( + forward_batch.forward_mode.is_extend() + and not forward_batch.forward_mode.is_target_verify() + and not forward_batch.forward_mode.is_draft_extend() + and not forward_batch.forward_mode.is_draft_extend_v2() + ): + if hasattr(attn, "indexer"): + return AttnForwardMethod.DSA_NPU + else: + return AttnForwardMethod.MHA_NPU + else: + if hasattr(attn, "indexer"): + return AttnForwardMethod.DSA_NPU + else: + return AttnForwardMethod.MLA_NPU + + +def _get_sum_extend_prefix_lens(forward_batch): + return ( + sum(forward_batch.extend_prefix_lens_cpu) + if forward_batch.extend_prefix_lens_cpu is not None + else 0 + ) + + +def _support_mha_one_shot(attn, forward_batch, backend_name): + attn_supported = backend_name in ["fa3", "flashinfer", "flashmla"] + sum_seq_lens = ( + sum(forward_batch.seq_lens_cpu) if forward_batch.seq_lens_cpu is not None else 0 + ) + return attn_supported and sum_seq_lens <= forward_batch.get_max_chunk_capacity() + + +def _handle_attention_backend(attn, forward_batch, backend_name): + if is_in_piecewise_cuda_graph(): + return AttnForwardMethod.MLA + + sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) + disable_ragged = ( + backend_name in ["flashinfer", "flashmla"] + ) and attn.flashinfer_mla_disable_ragged + + if ( + not disable_ragged + and forward_batch.forward_mode.is_extend_without_speculative() + and ( + ( + sum_extend_prefix_lens >= attn.chunked_prefix_cache_threshold + and not attn.disable_chunked_prefix_cache + ) + or sum_extend_prefix_lens == 0 + ) + ): + if _support_mha_one_shot(attn, forward_batch, backend_name): + return AttnForwardMethod.MHA_ONE_SHOT + return AttnForwardMethod.MHA_CHUNKED_KV + else: + return _dispatch_mla_subtype(attn, forward_batch) + + +def handle_attention_flashinfer(attn, forward_batch): + return _handle_attention_backend(attn, forward_batch, "flashinfer") + + +def handle_attention_fa3(attn, forward_batch): + # when deterministic inference is enabled, use MLA + if get_global_server_args().enable_deterministic_inference: + return _dispatch_mla_subtype(attn, forward_batch) + else: + return _handle_attention_backend(attn, forward_batch, "fa3") + + +def handle_attention_flashmla(attn, forward_batch): + return _handle_attention_backend(attn, forward_batch, "flashmla") + + +def handle_attention_cutlass_mla(attn, forward_batch): + return _handle_attention_backend(attn, forward_batch, "cutlass_mla") + + +def handle_attention_fa4(attn, forward_batch): + # TODO(cicirori): use FA4 MHA for DeepSeekV3 for now + return AttnForwardMethod.MHA_CHUNKED_KV + + +def handle_attention_trtllm_mla(attn, forward_batch): + if is_in_piecewise_cuda_graph(): + return AttnForwardMethod.MLA + + sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) + if forward_batch.forward_mode.is_extend_without_speculative() and ( + not attn.disable_chunked_prefix_cache or sum_extend_prefix_lens == 0 + ): + return AttnForwardMethod.MHA_CHUNKED_KV + else: + return _dispatch_mla_subtype(attn, forward_batch) + + +def handle_attention_aiter(attn, forward_batch): + if forward_batch.forward_mode.is_extend_without_speculative(): + return AttnForwardMethod.MHA + else: + return AttnForwardMethod.MLA + + +def handle_attention_nsa(attn, forward_batch): + """ + Dispatch logic is centralized in NativeSparseAttnBackend.set_nsa_prefill_impl and executed + in init_forward_metadata. Read the decision from backend.use_mha. + """ + + backend = forward_batch.attn_backend + if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend + backend = backend.primary + if hasattr(backend, "use_mha") and backend.use_mha: + return AttnForwardMethod.MHA_ONE_SHOT + return AttnForwardMethod.MLA + + +def handle_attention_triton(attn, forward_batch): + if is_in_piecewise_cuda_graph(): + return AttnForwardMethod.MLA + + # when deterministic inference is enabled, use MLA + if get_global_server_args().enable_deterministic_inference: + return _dispatch_mla_subtype(attn, forward_batch) + + if ( + forward_batch.forward_mode.is_extend_without_speculative() + and sum(forward_batch.extend_prefix_lens_cpu) == 0 + ): + return AttnForwardMethod.MHA + else: + return _dispatch_mla_subtype(attn, forward_batch) + + +AttentionBackendRegistry.register("ascend", handle_attention_ascend) +AttentionBackendRegistry.register("flashinfer", handle_attention_flashinfer) +AttentionBackendRegistry.register("fa3", handle_attention_fa3) +AttentionBackendRegistry.register("flashmla", handle_attention_flashmla) +AttentionBackendRegistry.register("cutlass_mla", handle_attention_cutlass_mla) +AttentionBackendRegistry.register("fa4", handle_attention_fa4) +AttentionBackendRegistry.register("trtllm_mla", handle_attention_trtllm_mla) +AttentionBackendRegistry.register("aiter", handle_attention_aiter) +AttentionBackendRegistry.register("nsa", handle_attention_nsa) +AttentionBackendRegistry.register("triton", handle_attention_triton) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py new file mode 100644 index 000000000..a75424a8d --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_methods.py @@ -0,0 +1,32 @@ +from enum import IntEnum, auto + + +class AttnForwardMethod(IntEnum): + # Use multi-head attention + MHA = auto() + + # Use absorbed multi-latent attention + MLA = auto() + + # Use multi-head attention, but with KV cache chunked. + # This method can avoid OOM when prefix lengths are long. + MHA_CHUNKED_KV = auto() + + # Use multi-head attention, execute the MHA for prefix and extended kv in one shot + # when the sequence lengths are below the threshold. + MHA_ONE_SHOT = auto() + + # Use MLA but with fused RoPE + MLA_FUSED_ROPE = auto() + + # Use MLA with fused RoPE kernel for CPU + MLA_FUSED_ROPE_CPU = auto() + + # Use multi-head attention for NPU + MHA_NPU = auto() + + # Use absorbed multi-latent attention for NPU + MLA_NPU = auto() + + # Use Deepseek V3.2 sparse multi-latent attention for NPU + DSA_NPU = auto() diff --git a/python/sglang/srt/models/deepseek_common/utils.py b/python/sglang/srt/models/deepseek_common/utils.py new file mode 100644 index 000000000..6c78f5683 --- /dev/null +++ b/python/sglang/srt/models/deepseek_common/utils.py @@ -0,0 +1,23 @@ +from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz +from sglang.srt.utils import ( + cpu_has_amx_support, + get_bool_env_var, + get_device_sm, + is_cpu, + is_cuda, + is_gfx95_supported, + is_hip, + is_npu, +) + +_is_hip = is_hip() +_is_cuda = is_cuda() +_is_npu = is_npu() +_is_fp8_fnuz = is_fp8_fnuz() +_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip +_is_cpu_amx_available = cpu_has_amx_support() +_is_cpu = is_cpu() +_device_sm = get_device_sm() +_is_gfx95_supported = is_gfx95_supported() + +_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index ed8cc7ada..81a2058f5 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -21,7 +21,6 @@ import concurrent.futures import logging import os from contextlib import nullcontext -from enum import IntEnum, auto from typing import Any, Dict, Iterable, List, Optional, Tuple, Union import torch @@ -109,7 +108,6 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8_kernel import ( fp8_dtype, - is_fp8_fnuz, per_tensor_quant_mla_fp8, per_token_group_quant_mla_deep_gemm_masked_fp8, ) @@ -138,6 +136,24 @@ from sglang.srt.model_loader.utils import ( should_deepgemm_weight_requant_ue8m0, ) 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 ( + AttnForwardMethod, +) +from sglang.srt.models.deepseek_common.utils import ( + _device_sm, + _is_cpu, + _is_cpu_amx_available, + _is_cuda, + _is_fp8_fnuz, + _is_gfx95_supported, + _is_hip, + _is_npu, + _use_aiter, + _use_aiter_gfx95, +) from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( @@ -145,33 +161,14 @@ from sglang.srt.utils import ( LazyValue, add_prefix, bind_or_assign, - cpu_has_amx_support, get_bool_env_var, - get_device_sm, - is_cpu, - is_cuda, - is_gfx95_supported, - is_hip, is_non_idle_and_non_empty, - is_npu, is_nvidia_cublas_cu12_version_ge_12_9, log_info_on_rank0, make_layers, use_intel_amx_backend, ) -_is_hip = is_hip() -_is_cuda = is_cuda() -_is_npu = is_npu() -_is_fp8_fnuz = is_fp8_fnuz() -_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip -_is_cpu_amx_available = cpu_has_amx_support() -_is_cpu = is_cpu() -_device_sm = get_device_sm() -_is_gfx95_supported = is_gfx95_supported() - -_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported - if _use_aiter_gfx95: from aiter.ops.triton.batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant import ( @@ -261,201 +258,6 @@ def add_forward_absorb_core_attention_backend(backend_name): logger.info(f"Added {backend_name} to FORWARD_ABSORB_CORE_ATTENTION_BACKENDS.") -class AttnForwardMethod(IntEnum): - # Use multi-head attention - MHA = auto() - - # Use absorbed multi-latent attention - MLA = auto() - - # Use multi-head attention, but with KV cache chunked. - # This method can avoid OOM when prefix lengths are long. - MHA_CHUNKED_KV = auto() - - # Use multi-head attention, execute the MHA for prefix and extended kv in one shot - # when the sequence lengths are below the threshold. - MHA_ONE_SHOT = auto() - - # Use MLA but with fused RoPE - MLA_FUSED_ROPE = auto() - - # Use MLA with fused RoPE kernel for CPU - MLA_FUSED_ROPE_CPU = auto() - - # Use multi-head attention for NPU - MHA_NPU = auto() - - # Use absorbed multi-latent attention for NPU - MLA_NPU = auto() - - # Use Deepseek V3.2 sparse multi-latent attention for NPU - DSA_NPU = auto() - - -def _dispatch_mla_subtype(attn, forward_batch): - if _is_hip: - if attn.rocm_fused_decode_mla and forward_batch.forward_mode.is_decode(): - return AttnForwardMethod.MLA_FUSED_ROPE - else: - return AttnForwardMethod.MLA - else: - if hasattr(attn, "fused_qkv_a_proj_with_mqa") and use_intel_amx_backend(attn): - return AttnForwardMethod.MLA_FUSED_ROPE_CPU - else: - return AttnForwardMethod.MLA - - -class AttentionBackendRegistry: - _handlers = {} - - @classmethod - def register(cls, backend_name, handler_func): - cls._handlers[backend_name] = handler_func - - @classmethod - def get_handler(cls, backend_name): - return cls._handlers.get(backend_name, cls._handlers.get("triton")) - - -def handle_attention_ascend(attn, forward_batch): - if ( - forward_batch.forward_mode.is_extend() - and not forward_batch.forward_mode.is_target_verify() - and not forward_batch.forward_mode.is_draft_extend() - and not forward_batch.forward_mode.is_draft_extend_v2() - ): - if hasattr(attn, "indexer"): - return AttnForwardMethod.DSA_NPU - else: - return AttnForwardMethod.MHA_NPU - else: - if hasattr(attn, "indexer"): - return AttnForwardMethod.DSA_NPU - else: - return AttnForwardMethod.MLA_NPU - - -def _get_sum_extend_prefix_lens(forward_batch): - return ( - sum(forward_batch.extend_prefix_lens_cpu) - if forward_batch.extend_prefix_lens_cpu is not None - else 0 - ) - - -def _support_mha_one_shot(attn: DeepseekV2AttentionMLA, forward_batch, backend_name): - attn_supported = backend_name in ["fa3", "flashinfer", "flashmla"] - sum_seq_lens = ( - sum(forward_batch.seq_lens_cpu) if forward_batch.seq_lens_cpu is not None else 0 - ) - return attn_supported and sum_seq_lens <= forward_batch.get_max_chunk_capacity() - - -def _handle_attention_backend( - attn: DeepseekV2AttentionMLA, forward_batch, backend_name -): - if is_in_piecewise_cuda_graph(): - return AttnForwardMethod.MLA - - sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) - disable_ragged = ( - backend_name in ["flashinfer", "flashmla"] - ) and attn.flashinfer_mla_disable_ragged - - if ( - not disable_ragged - and forward_batch.forward_mode.is_extend_without_speculative() - and ( - ( - sum_extend_prefix_lens >= attn.chunked_prefix_cache_threshold - and not attn.disable_chunked_prefix_cache - ) - or sum_extend_prefix_lens == 0 - ) - ): - if _support_mha_one_shot(attn, forward_batch, backend_name): - return AttnForwardMethod.MHA_ONE_SHOT - return AttnForwardMethod.MHA_CHUNKED_KV - else: - return _dispatch_mla_subtype(attn, forward_batch) - - -def handle_attention_flashinfer(attn, forward_batch): - return _handle_attention_backend(attn, forward_batch, "flashinfer") - - -def handle_attention_fa3(attn, forward_batch): - # when deterministic inference is enabled, use MLA - if get_global_server_args().enable_deterministic_inference: - return _dispatch_mla_subtype(attn, forward_batch) - else: - return _handle_attention_backend(attn, forward_batch, "fa3") - - -def handle_attention_flashmla(attn, forward_batch): - return _handle_attention_backend(attn, forward_batch, "flashmla") - - -def handle_attention_cutlass_mla(attn, forward_batch): - return _handle_attention_backend(attn, forward_batch, "cutlass_mla") - - -def handle_attention_fa4(attn, forward_batch): - # TODO(cicirori): use FA4 MHA for DeepSeekV3 for now - return AttnForwardMethod.MHA_CHUNKED_KV - - -def handle_attention_trtllm_mla(attn, forward_batch): - if is_in_piecewise_cuda_graph(): - return AttnForwardMethod.MLA - - sum_extend_prefix_lens = _get_sum_extend_prefix_lens(forward_batch) - if forward_batch.forward_mode.is_extend_without_speculative() and ( - not attn.disable_chunked_prefix_cache or sum_extend_prefix_lens == 0 - ): - return AttnForwardMethod.MHA_CHUNKED_KV - else: - return _dispatch_mla_subtype(attn, forward_batch) - - -def handle_attention_aiter(attn, forward_batch): - if forward_batch.forward_mode.is_extend_without_speculative(): - return AttnForwardMethod.MHA - else: - return AttnForwardMethod.MLA - - -def handle_attention_nsa(attn, forward_batch): - """ - Dispatch logic is centralized in NativeSparseAttnBackend.set_nsa_prefill_impl and executed - in init_forward_metadata. Read the decision from backend.use_mha. - """ - - backend = forward_batch.attn_backend - if isinstance(backend, TboAttnBackend): # if enable tbo, get primary backend - backend = backend.primary - if hasattr(backend, "use_mha") and backend.use_mha: - return AttnForwardMethod.MHA_ONE_SHOT - return AttnForwardMethod.MLA - - -def handle_attention_triton(attn, forward_batch): - if is_in_piecewise_cuda_graph(): - return AttnForwardMethod.MLA - - # when deterministic inference is enabled, use MLA - if get_global_server_args().enable_deterministic_inference: - return _dispatch_mla_subtype(attn, forward_batch) - - if ( - forward_batch.forward_mode.is_extend_without_speculative() - and sum(forward_batch.extend_prefix_lens_cpu) == 0 - ): - return AttnForwardMethod.MHA - else: - return _dispatch_mla_subtype(attn, forward_batch) - - class DeepseekV2MLP(nn.Module): def __init__( self, @@ -4015,18 +3817,6 @@ class DeepseekV2ForCausalLM(nn.Module): return list(weights_dict.items()) -AttentionBackendRegistry.register("ascend", handle_attention_ascend) -AttentionBackendRegistry.register("flashinfer", handle_attention_flashinfer) -AttentionBackendRegistry.register("fa3", handle_attention_fa3) -AttentionBackendRegistry.register("flashmla", handle_attention_flashmla) -AttentionBackendRegistry.register("cutlass_mla", handle_attention_cutlass_mla) -AttentionBackendRegistry.register("fa4", handle_attention_fa4) -AttentionBackendRegistry.register("trtllm_mla", handle_attention_trtllm_mla) -AttentionBackendRegistry.register("aiter", handle_attention_aiter) -AttentionBackendRegistry.register("nsa", handle_attention_nsa) -AttentionBackendRegistry.register("triton", handle_attention_triton) - - class DeepseekV3ForCausalLM(DeepseekV2ForCausalLM): pass