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
2 changes: 2 additions & 0 deletions docker/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 && \
Expand Down
34 changes: 34 additions & 0 deletions lightllm/common/basemodel/attention/linear/create_utils.py
Original file line number Diff line number Diff line change
@@ -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
51 changes: 43 additions & 8 deletions lightllm/common/basemodel/attention/linear/gdn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):

Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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):

Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down
7 changes: 4 additions & 3 deletions lightllm/models/qwen3next/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
54 changes: 49 additions & 5 deletions lightllm/utils/backend_validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand All @@ -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":
Expand All @@ -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]
Expand All @@ -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)

Expand Down
96 changes: 89 additions & 7 deletions unit_tests/common/basemodel/attention/linear/test_gdn.py
Original file line number Diff line number Diff line change
@@ -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])
Expand All @@ -13,19 +15,14 @@ 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(
mtp_step=0,
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,
Expand Down Expand Up @@ -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"]
Loading