diff --git a/lightllm/common/basemodel/attention/linear/gdn.py b/lightllm/common/basemodel/attention/linear/gdn.py index c4c3e15cb7..f313d2e9b9 100644 --- a/lightllm/common/basemodel/attention/linear/gdn.py +++ b/lightllm/common/basemodel/attention/linear/gdn.py @@ -5,8 +5,8 @@ from lightllm.utils.envs_utils import get_env_start_args, get_llm_data_type 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.gdn_prefill_backend import get_gdn_prefill_chunk_fn from lightllm.common.basemodel.triton_kernel.linear_att.mtp_fused_recurrent import ( mtp_fused_recurrent_gated_delta_rule, ) @@ -58,6 +58,7 @@ def _init_linear_layer_metadata(self, network_config, tp_world_size): # GDN kernel output dtype is self.data_type # Conversion needed only if SSM state uses different dtype self.needs_ssm_dtype_conversion = get_llm_data_type() != self.ssm_state_dtype + self._gdn_prefill_chunk = get_gdn_prefill_chunk_fn() return def _split_qkvzba(self, mixed_qkvzba): @@ -175,7 +176,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._gdn_prefill_chunk( q=query, k=key, v=value, diff --git a/lightllm/common/basemodel/triton_kernel/linear_att/gdn_prefill_backend.py b/lightllm/common/basemodel/triton_kernel/linear_att/gdn_prefill_backend.py new file mode 100644 index 0000000000..7642286717 --- /dev/null +++ b/lightllm/common/basemodel/triton_kernel/linear_att/gdn_prefill_backend.py @@ -0,0 +1,70 @@ +import functools + +import torch + +from lightllm.common.basemodel.triton_kernel.linear_att.fla.ops import ( + chunk_gated_delta_rule as _fla_chunk_gated_delta_rule, +) +from lightllm.utils.envs_utils import get_env_start_args +from lightllm.utils.log_utils import init_logger + +logger = init_logger(__name__) + + +@functools.lru_cache(maxsize=1) +def get_gdn_prefill_chunk_fn(): + backend = getattr(get_env_start_args(), "gdn_prefill_backend", "fla") + if backend != "flashqla": + return _fla_chunk_gated_delta_rule + + capability = torch.cuda.get_device_capability() + if capability[0] < 9: + logger.warning( + f"gdn_prefill_backend=flashqla requires Hopper (SM90+), got SM{capability[0]}{capability[1]}; " + "falling back to the FLA triton kernel." + ) + return _fla_chunk_gated_delta_rule + + try: + from flash_qla import chunk_gated_delta_rule as flashqla_chunk_gated_delta_rule + except Exception as exc: + logger.warning( + f"gdn_prefill_backend=flashqla but importing flash_qla failed ({exc!r}); " + "falling back to the FLA triton kernel. " + "Install FlashQLA (https://github.com/QwenLM/FlashQLA)." + ) + return _fla_chunk_gated_delta_rule + + flashqla_initialized = False + use_flashqla = True + + def flashqla_chunk(q, k, v, **kwargs): + nonlocal flashqla_initialized, use_flashqla + + if not use_flashqla: + return _fla_chunk_gated_delta_rule(q=q, k=k, v=v, **kwargs) + + flashqla_q = q.contiguous() + flashqla_k = k.contiguous() + flashqla_v = v.contiguous() + try: + result = flashqla_chunk_gated_delta_rule( + q=flashqla_q, + k=flashqla_k, + v=flashqla_v, + **kwargs, + ) + except Exception as exc: + if flashqla_initialized: + raise + use_flashqla = False + logger.warning( + f"FlashQLA failed during its first invocation ({exc!r}); " "falling back to the FLA triton kernel." + ) + return _fla_chunk_gated_delta_rule(q=q, k=k, v=v, **kwargs) + + flashqla_initialized = True + return result + + logger.info("GDN chunked-prefill backend: FlashQLA (TileLang, Hopper).") + return flashqla_chunk diff --git a/lightllm/server/api_cli.py b/lightllm/server/api_cli.py index e85d083075..8f86c155e0 100644 --- a/lightllm/server/api_cli.py +++ b/lightllm/server/api_cli.py @@ -820,6 +820,15 @@ def add_cli_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: default="float32", help="the data type of linear att smm data type", ) + parser.add_argument( + "--gdn_prefill_backend", + type=str, + choices=["fla", "flashqla"], + default="fla", + help="""GDN chunked-prefill kernel backend for hybrid linear-attention models. + 'fla' uses the vendored flash-linear-attention Triton kernel. 'flashqla' uses the + TileLang FlashQLA kernel on SM90+ and falls back to 'fla' if FlashQLA is unavailable.""", + ) parser.add_argument( "--disable_linear_att_small_page_cpu_cache", action="store_true", diff --git a/lightllm/server/core/objs/start_args_type.py b/lightllm/server/core/objs/start_args_type.py index 9e92b02e1b..1ab26869c7 100644 --- a/lightllm/server/core/objs/start_args_type.py +++ b/lightllm/server/core/objs/start_args_type.py @@ -226,3 +226,4 @@ class StartArgs: disable_linear_att_small_page_cpu_cache: bool = field(default=False) linear_att_cache_size: Optional[int] = field(default=None) linear_att_ssm_data_type: Optional[str] = field(default="float32", metadata={"choices": ["bfloat16", "float32"]}) + gdn_prefill_backend: str = field(default="fla", metadata={"choices": ["fla", "flashqla"]}) diff --git a/test/benchmark/static_inference/model_infer.py b/test/benchmark/static_inference/model_infer.py index f2c900af09..5d55230b74 100644 --- a/test/benchmark/static_inference/model_infer.py +++ b/test/benchmark/static_inference/model_infer.py @@ -218,6 +218,7 @@ def decode( b_req_idx=b_req_idx, b_seq_len=b_seq_len, b_mtp_index=b_mtp_index, + b_position_delta=torch.zeros(batch_size, dtype=torch.int32, device="cpu"), mem_indexes_cpu=mem_indexes, is_prefill=False, multimodal_params=[{"images": [], "audios": []} for _ in range(batch_size)], diff --git a/unit_tests/common/basemodel/triton_kernel/test_gdn_prefill_backend.py b/unit_tests/common/basemodel/triton_kernel/test_gdn_prefill_backend.py new file mode 100644 index 0000000000..21f205791b --- /dev/null +++ b/unit_tests/common/basemodel/triton_kernel/test_gdn_prefill_backend.py @@ -0,0 +1,61 @@ +import sys +import types +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import torch + +from lightllm.common.basemodel.triton_kernel.linear_att import gdn_prefill_backend + + +@pytest.fixture(autouse=True) +def clear_backend_cache(): + gdn_prefill_backend.get_gdn_prefill_chunk_fn.cache_clear() + yield + gdn_prefill_backend.get_gdn_prefill_chunk_fn.cache_clear() + + +def configure_flashqla(monkeypatch, flashqla_chunk): + flashqla = types.ModuleType("flash_qla") + flashqla.chunk_gated_delta_rule = flashqla_chunk + monkeypatch.setitem(sys.modules, "flash_qla", flashqla) + monkeypatch.setattr( + gdn_prefill_backend, + "get_env_start_args", + lambda: SimpleNamespace(gdn_prefill_backend="flashqla"), + ) + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: (9, 0)) + + +def test_first_flashqla_failure_falls_back_permanently(monkeypatch): + flashqla_chunk = Mock(side_effect=RuntimeError("nvcc not found")) + fla_result = object() + fla_chunk = Mock(return_value=fla_result) + configure_flashqla(monkeypatch, flashqla_chunk) + monkeypatch.setattr(gdn_prefill_backend, "_fla_chunk_gated_delta_rule", fla_chunk) + + chunk = gdn_prefill_backend.get_gdn_prefill_chunk_fn() + q, k, v = torch.randn(1), torch.randn(1), torch.randn(1) + + assert chunk(q=q, k=k, v=v, head_first=False) is fla_result + assert chunk(q=q, k=k, v=v, head_first=False) is fla_result + assert flashqla_chunk.call_count == 1 + assert fla_chunk.call_count == 2 + fla_chunk.assert_called_with(q=q, k=k, v=v, head_first=False) + + +def test_flashqla_failure_after_success_is_raised(monkeypatch): + flashqla_result = object() + flashqla_chunk = Mock(side_effect=[flashqla_result, RuntimeError("kernel failed")]) + fla_chunk = Mock() + configure_flashqla(monkeypatch, flashqla_chunk) + monkeypatch.setattr(gdn_prefill_backend, "_fla_chunk_gated_delta_rule", fla_chunk) + + chunk = gdn_prefill_backend.get_gdn_prefill_chunk_fn() + q, k, v = torch.randn(1), torch.randn(1), torch.randn(1) + + assert chunk(q=q, k=k, v=v) is flashqla_result + with pytest.raises(RuntimeError, match="kernel failed"): + chunk(q=q, k=k, v=v) + fla_chunk.assert_not_called()