[1/n]deepseek_v2.py Refactor: attention backend handlers and forward method definition (#16306)
This commit is contained in:
@@ -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)
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -21,7 +21,6 @@ import concurrent.futures
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
from enum import IntEnum, auto
|
|
||||||
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
|
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
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 import Fp8Config
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import (
|
from sglang.srt.layers.quantization.fp8_kernel import (
|
||||||
fp8_dtype,
|
fp8_dtype,
|
||||||
is_fp8_fnuz,
|
|
||||||
per_tensor_quant_mla_fp8,
|
per_tensor_quant_mla_fp8,
|
||||||
per_token_group_quant_mla_deep_gemm_masked_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,
|
should_deepgemm_weight_requant_ue8m0,
|
||||||
)
|
)
|
||||||
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.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.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -145,33 +161,14 @@ from sglang.srt.utils import (
|
|||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
bind_or_assign,
|
bind_or_assign,
|
||||||
cpu_has_amx_support,
|
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
get_device_sm,
|
|
||||||
is_cpu,
|
|
||||||
is_cuda,
|
|
||||||
is_gfx95_supported,
|
|
||||||
is_hip,
|
|
||||||
is_non_idle_and_non_empty,
|
is_non_idle_and_non_empty,
|
||||||
is_npu,
|
|
||||||
is_nvidia_cublas_cu12_version_ge_12_9,
|
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
_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:
|
if _use_aiter_gfx95:
|
||||||
|
|
||||||
from aiter.ops.triton.batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant import (
|
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.")
|
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):
|
class DeepseekV2MLP(nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -4015,18 +3817,6 @@ class DeepseekV2ForCausalLM(nn.Module):
|
|||||||
return list(weights_dict.items())
|
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):
|
class DeepseekV3ForCausalLM(DeepseekV2ForCausalLM):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user