diff --git a/docker/Dockerfile b/docker/Dockerfile index 8fc267dd8..824737daf 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -9,6 +9,7 @@ ARG FLASH_MLA_REF=47c35a7 ARG DEEPGEMM_REF=891d57b4db1071624b5c8fa0d1e51cb317fa709f ARG DEEPEP_REF=60d44037a702f651a6e18bd4aea65ed8409051c2 ARG DEEPEP_NCCL_VERSION=2.30.4 +ARG FLASHQLA_VERSION=0.1.2 ARG TARGETPLATFORM ARG ENABLE_DEEPEP=1 ARG ENABLE_NIXL=1 @@ -57,6 +58,7 @@ RUN pip install --no-cache-dir \ RUN pip install -r /lightllm/requirements.txt --no-cache-dir \ -i https://pypi.org/simple \ --extra-index-url https://download.pytorch.org/whl/cu130 +RUN PIP_NO_INDEX=0 pip install --no-cache-dir "flash-qla==${FLASHQLA_VERSION}" -i https://pypi.org/simple RUN export CPATH=/usr/local/cuda/targets/x86_64-linux/include/cccl:/usr/local/cuda/targets/x86_64-linux/include${CPATH:+:${CPATH}} && \ git clone https://github.com/deepseek-ai/FlashMLA.git /root/FlashMLA && \ cd /root/FlashMLA && \ diff --git a/lightllm/common/basemodel/attention/linear/create_utils.py b/lightllm/common/basemodel/attention/linear/create_utils.py new file mode 100644 index 000000000..004aa04b2 --- /dev/null +++ b/lightllm/common/basemodel/attention/linear/create_utils.py @@ -0,0 +1,34 @@ +import os + +from lightllm.common.basemodel.attention.linear.gdn import ( + FlaLinearAttBackend, + FlashQlaLinearAttBackend, + LinearAttBackend, +) +from lightllm.utils.backend_validator import validate +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + +linear_att_backend_classes = { + "flashqla": FlashQlaLinearAttBackend, + "fla": FlaLinearAttBackend, +} + + +def get_linear_att_backend_class(model, priority_list=("flashqla", "fla")): + if os.environ.get("FLA_FLASH_QLA", "1") == "0": + priority_list = ("fla",) + + backend_args = LinearAttBackend.get_gdn_prefill_validation_args(model) + for backend_name in priority_list: + backend_class = linear_att_backend_classes[backend_name] + if backend_name == "fla": + logger.info("Linear attention backend: fla.") + return backend_class + if validate(backend_name, *backend_args): + logger.info(f"Linear attention backend: {backend_name} (validated).") + return backend_class + + logger.warning("No linear attention backend validation succeeded, falling back to FLA.") + return FlaLinearAttBackend diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index ae2816840..13cc51f46 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -5,12 +5,14 @@ from lightllm.utils.envs_utils import get_env_start_args from lightllm.common.basemodel.triton_kernel.linear_att.causal_conv1d import causal_conv1d_fn from lightllm.common.basemodel.triton_kernel.linear_att.fused_gdn_gating import fused_gdn_gating -from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import chunk_gated_delta_rule from lightllm.common.basemodel.triton_kernel.linear_att.gdn_decode_pack import conv_pack_gdn_decode_inputs from lightllm.common.basemodel.triton_kernel.linear_att.mtp_fused_recurrent import ( mtp_fused_recurrent_gated_delta_rule, ) -from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import fused_recurrent_gated_delta_rule +from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import ( + chunk_gated_delta_rule as fla_chunk_gated_delta_rule, + fused_recurrent_gated_delta_rule, +) if TYPE_CHECKING: from lightllm.common.basemodel.basemodel import TpPartBaseModel @@ -23,6 +25,29 @@ class LinearAttBackend(BaseAttBackend): def __init__(self, model: "TpPartBaseModel"): super().__init__(model=model) self._init_linear_layer_metadata(network_config=model.config, tp_world_size=model.tp_world_size_) + self._chunk_gated_delta_rule = self._get_chunk_gated_delta_rule() + + @staticmethod + def _get_ssm_state_dtype(): + ssm_dtype_dict = {"bfloat16": torch.bfloat16, "float32": torch.float32} + return ssm_dtype_dict.get(get_env_start_args().linear_att_ssm_data_type, torch.bfloat16) + + @classmethod + def get_gdn_prefill_validation_args(cls, model: "TpPartBaseModel"): + network_config = model.config + tp_world_size = model.tp_world_size_ + return ( + network_config["linear_num_key_heads"] // tp_world_size, + network_config["linear_num_value_heads"] // tp_world_size, + network_config["linear_key_head_dim"], + network_config["linear_value_head_dim"], + model.data_type, + cls._get_ssm_state_dtype(), + ) + + @staticmethod + def _get_chunk_gated_delta_rule(): + raise NotImplementedError def _init_linear_layer_metadata(self, network_config, tp_world_size): @@ -50,10 +75,7 @@ def _init_linear_layer_metadata(self, network_config, tp_world_size): self.num_v_heads_per_k_head = self.num_v_heads // self.num_k_heads # SSM state dtype optimization - ssm_dtype_dict = {"bfloat16": torch.bfloat16, "float32": torch.float32} - start_args = get_env_start_args() - self.ssm_state_dtype = ssm_dtype_dict.get(start_args.linear_att_ssm_data_type, torch.bfloat16) - + self.ssm_state_dtype = self._get_ssm_state_dtype() return def _split_qkvzba(self, mixed_qkvzba): @@ -97,6 +119,20 @@ def create_att_decode_state(self, infer_state: "InferStateInfo") -> "LinearAttDe return LinearAttDecodeAttState(backend=self, infer_state=infer_state) +class FlashQlaLinearAttBackend(LinearAttBackend): + @staticmethod + def _get_chunk_gated_delta_rule(): + from flash_qla import chunk_gated_delta_rule + + return chunk_gated_delta_rule + + +class FlaLinearAttBackend(LinearAttBackend): + @staticmethod + def _get_chunk_gated_delta_rule(): + return fla_chunk_gated_delta_rule + + @dataclasses.dataclass class LinearAttPrefillAttState(BasePrefillAttState): @@ -171,7 +207,7 @@ def _gdn_prefill_kernel( query, key, value = backend._rearrange_mixed_qkv(mixed_qkv) initial_state = ssm_states[self.b_ssm_buffer_idx] # g and beta have shape (total_tokens, num_heads), need to unsqueeze to get (1, total_tokens, num_heads) - core_attn_out, last_recurrent_state = chunk_gated_delta_rule( + core_attn_out, last_recurrent_state = backend._chunk_gated_delta_rule( q=query, k=key, v=value, @@ -180,7 +216,6 @@ def _gdn_prefill_kernel( initial_state=initial_state, output_final_state=True, cu_seqlens=infer_state.b1_cu_q_seq_len, - head_first=False, use_qk_l2norm_in_kernel=True, ) # The chunk kernel accumulates the recurrent state in float32 even when diff --git a/lightllm/models/qwen3next/model.py b/lightllm/models/qwen3next/model.py index e31b83ffe..62d06c5e9 100644 --- a/lightllm/models/qwen3next/model.py +++ b/lightllm/models/qwen3next/model.py @@ -112,8 +112,9 @@ def _init_att_backend1(self): self.prefill_att_backend1 = None self.decode_att_backend1 = None return - from lightllm.common.basemodel.attention.linear.gdn import LinearAttBackend + from lightllm.common.basemodel.attention.linear.create_utils import get_linear_att_backend_class - self.prefill_att_backend1 = LinearAttBackend(model=self) - self.decode_att_backend1 = LinearAttBackend(model=self) + linear_att_backend_class = get_linear_att_backend_class(self) + self.prefill_att_backend1 = linear_att_backend_class(model=self) + self.decode_att_backend1 = linear_att_backend_class(model=self) return diff --git a/lightllm/utils/backend_validator.py b/lightllm/utils/backend_validator.py index ab5c0a88a..ae8525399 100644 --- a/lightllm/utils/backend_validator.py +++ b/lightllm/utils/backend_validator.py @@ -92,6 +92,48 @@ def _validate_flashinfer(): return True, None +def _validate_flashqla(num_k_heads, num_v_heads, head_k_dim, head_v_dim, qkv_dtype, state_dtype): + """Validate FlashQLA against LightLLM's vendored FLA kernel.""" + from flash_qla import chunk_gated_delta_rule as flashqla_chunk_gated_delta_rule + + from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import ( + chunk_gated_delta_rule as fla_chunk_gated_delta_rule, + ) + + batch, seq = 1, 64 + torch.manual_seed(0) + q = torch.randn(batch, seq, num_k_heads, head_k_dim, dtype=qkv_dtype, device="cuda") + k = torch.randn_like(q) + v = torch.randn(batch, seq, num_v_heads, head_v_dim, dtype=qkv_dtype, device="cuda") + g = -torch.rand(batch, seq, num_v_heads, dtype=torch.float32, device="cuda") + beta = torch.rand(batch, seq, num_v_heads, dtype=torch.float32, device="cuda") + initial_state = torch.randn(batch, num_v_heads, head_k_dim, head_v_dim, dtype=state_dtype, device="cuda") + cu_seqlens = torch.tensor([0, seq], dtype=torch.int32, device="cuda") + kwargs = { + "q": q, + "k": k, + "v": v, + "g": g, + "beta": beta, + "initial_state": initial_state, + "output_final_state": True, + "cu_seqlens": cu_seqlens, + "use_qk_l2norm_in_kernel": True, + } + + expected_out, expected_state = fla_chunk_gated_delta_rule(**kwargs) + out, final_state = flashqla_chunk_gated_delta_rule(**kwargs) + torch.cuda.synchronize() + + for name, actual, expected in ( + ("output", out, expected_out), + ("final state", final_state, expected_state), + ): + if not torch.allclose(actual, expected, rtol=1e-2, atol=1e-2): + return False, f"{name} mismatch: max diff {(actual - expected).abs().max().item():.6f}" + return True, None + + def _validate_triton(): """Validate Triton with softmax ground truth.""" import triton @@ -230,7 +272,7 @@ def _validate_flashmla_sparse(): return True, None -def _run_in_subprocess(backend_name, pipe): +def _run_in_subprocess(backend_name, backend_args, pipe): """Run validation in subprocess with suppressed output.""" import sys @@ -248,6 +290,8 @@ def _run_in_subprocess(backend_name, pipe): success, err = _validate_sdpa() elif backend_name == "flashinfer": success, err = _validate_flashinfer() + elif backend_name == "flashqla": + success, err = _validate_flashqla(*backend_args) elif backend_name == "triton": success, err = _validate_triton() elif backend_name == "flashmla_sparse": @@ -262,9 +306,9 @@ def _run_in_subprocess(backend_name, pipe): @lru_cache(maxsize=None) -def validate(backend_name: str) -> bool: +def validate(backend_name: str, *backend_args) -> bool: if get_global_rank() == 0: - validate_ok = _validate(backend_name) + validate_ok = _validate(backend_name, *backend_args) torch.distributed.broadcast_object_list([validate_ok], src=0) else: validate_ok = [None] @@ -273,13 +317,13 @@ def validate(backend_name: str) -> bool: return validate_ok -def _validate(backend_name: str) -> bool: +def _validate(backend_name: str, *backend_args) -> bool: """Validate backend in subprocess with ground truth check.""" try: ctx = mp.get_context("spawn") parent, child = ctx.Pipe(duplex=False) logger.info(f"Validating {backend_name} backend start, please wait ...") - proc = ctx.Process(target=_run_in_subprocess, args=(backend_name, child)) + proc = ctx.Process(target=_run_in_subprocess, args=(backend_name, backend_args, child)) proc.start() proc.join(timeout=_VALIDATION_TIMEOUT) diff --git a/unit_tests/common/basemodel/attention/linear/test_gdn.py b/unit_tests/common/basemodel/attention/linear/test_gdn.py index 241e0a5d2..56e2c45ce 100644 --- a/unit_tests/common/basemodel/attention/linear/test_gdn.py +++ b/unit_tests/common/basemodel/attention/linear/test_gdn.py @@ -1,9 +1,11 @@ -from types import SimpleNamespace +import sys +from types import ModuleType, SimpleNamespace import pytest import torch import lightllm.common.basemodel.attention.linear.gdn as gdn +import lightllm.common.basemodel.attention.linear.create_utils as linear_create_utils @pytest.mark.parametrize("cache_dtype", [torch.bfloat16, torch.float32]) @@ -13,12 +15,6 @@ def test_prefill_casts_final_state_to_cache_dtype(monkeypatch, cache_dtype): monkeypatch.setattr(gdn, "fused_gdn_gating", lambda _log, a, b, _bias: (a, b)) monkeypatch.setattr(gdn, "causal_conv1d_fn", lambda mixed, *args, **kwargs: mixed) - monkeypatch.setattr( - gdn, - "chunk_gated_delta_rule", - lambda *args, **kwargs: (None, final_state), - ) - qkv = torch.zeros((1, 3), dtype=cache_dtype) q = torch.zeros((1, 1, 1, 1), dtype=cache_dtype) backend = SimpleNamespace( @@ -26,6 +22,7 @@ def test_prefill_casts_final_state_to_cache_dtype(monkeypatch, cache_dtype): activation="silu", ssm_state_dtype=cache_dtype, _rearrange_mixed_qkv=lambda mixed: (q, q, q), + _chunk_gated_delta_rule=lambda *args, **kwargs: (None, final_state), ) state = gdn.LinearAttPrefillAttState( backend=backend, @@ -54,3 +51,88 @@ def test_prefill_casts_final_state_to_cache_dtype(monkeypatch, cache_dtype): assert ssm_states.dtype == cache_dtype assert torch.equal(ssm_states, final_state.to(cache_dtype)) + + +@pytest.fixture +def linear_model(monkeypatch): + monkeypatch.setattr( + gdn, + "get_env_start_args", + lambda: SimpleNamespace(linear_att_ssm_data_type="float32"), + ) + return SimpleNamespace( + config={ + "linear_num_key_heads": 2, + "linear_num_value_heads": 4, + "linear_key_head_dim": 128, + "linear_value_head_dim": 128, + }, + tp_world_size_=2, + data_type=torch.bfloat16, + ) + + +def test_gdn_prefill_backend_uses_validated_flashqla(monkeypatch, linear_model): + backend_args = (1, 2, 128, 128, torch.bfloat16, torch.float32) + validate_calls = [] + monkeypatch.setenv("FLA_FLASH_QLA", "1") + flashqla = ModuleType("flash_qla") + flashqla.chunk_gated_delta_rule = lambda **kwargs: ("flashqla", kwargs) + monkeypatch.setitem(sys.modules, "flash_qla", flashqla) + monkeypatch.setattr( + linear_create_utils, + "validate", + lambda name, *args: validate_calls.append((name, args)) or True, + ) + + backend_class = linear_create_utils.get_linear_att_backend_class(linear_model) + + assert backend_class is gdn.FlashQlaLinearAttBackend + assert backend_class._get_chunk_gated_delta_rule()(q="q")[0] == "flashqla" + assert validate_calls == [("flashqla", backend_args)] + + +def test_gdn_prefill_backend_falls_back_when_flashqla_validation_fails(monkeypatch, linear_model): + monkeypatch.setenv("FLA_FLASH_QLA", "1") + monkeypatch.setattr(linear_create_utils, "validate", lambda *args: False) + + backend_class = linear_create_utils.get_linear_att_backend_class(linear_model) + + assert backend_class is gdn.FlaLinearAttBackend + + +def test_flashqla_backend_respects_disable_env(monkeypatch, linear_model): + monkeypatch.setenv("FLA_FLASH_QLA", "0") + monkeypatch.setattr( + linear_create_utils, + "validate", + lambda *args: pytest.fail("disabled FlashQLA must not be validated"), + ) + + backend_class = linear_create_utils.get_linear_att_backend_class(linear_model) + + assert backend_class is gdn.FlaLinearAttBackend + + +def test_gdn_prefill_backend_tries_candidates_in_order(monkeypatch, linear_model): + validate_calls = [] + monkeypatch.setenv("FLA_FLASH_QLA", "1") + flashqla2_backend = type("FlashQla2LinearAttBackend", (), {}) + flashqla3_backend = type("FlashQla3LinearAttBackend", (), {}) + monkeypatch.setattr( + linear_create_utils, + "linear_att_backend_classes", + {"flashqla2": flashqla2_backend, "flashqla3": flashqla3_backend}, + ) + monkeypatch.setattr( + linear_create_utils, + "validate", + lambda name, *args: validate_calls.append(name) or name == "flashqla3", + ) + + backend_class = linear_create_utils.get_linear_att_backend_class( + linear_model, priority_list=("flashqla2", "flashqla3") + ) + + assert backend_class is flashqla3_backend + assert validate_calls == ["flashqla2", "flashqla3"]