Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions lightllm/common/basemodel/attention/linear/gdn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
@@ -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
9 changes: 9 additions & 0 deletions lightllm/server/api_cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
1 change: 1 addition & 0 deletions lightllm/server/core/objs/start_args_type.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]})
1 change: 1 addition & 0 deletions test/benchmark/static_inference/model_infer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)],
Expand Down
Original file line number Diff line number Diff line change
@@ -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()
Loading