diff --git a/ENVs.md b/ENVs.md index 13e5f943d5..7b0f34647a 100644 --- a/ENVs.md +++ b/ENVs.md @@ -52,6 +52,7 @@ Each backend is gated by a single env var below. | ------------------------------ | ----------- | ---------- | ---------------------------------------------------------------------------------------------------------------------------------------- | | `FLA_DISABLE_BACKEND_DISPATCH` | unset (`0`) | `0` / `1` | Master switch. Set to `1` to bypass *all* backend dispatch and always use the default Triton implementation. Useful for debugging. | | `FLA_TILELANG` | unset (`1`) | `0` / `1` | Enable the TileLang backend when the `tilelang` package is installed. Set to `0` to force the Triton path (e.g. to work around #640). | +| `FLA_TLE_KDA` | unset (`1`) | `0` / `1` | Enable the Triton TLE forward backend for `chunk_kda` inference on NVIDIA Hopper GPUs. Requires Triton TLE. Set to `0` to skip it. | | `FLA_FLASH_KDA` | unset (`1`) | `0` / `1` | Enable the [FlashKDA](https://github.com/MoonshotAI/FlashKDA) CUTLASS forward for `chunk_kda` (inference only). Requires `flash_kda`. | | `FLA_INTRACARD_CP` | unset (`0`) | `0` / `1` | Opt in to the intra-card context-parallel backend for shared delta-rule ops (`chunk_gated_delta_rule_fwd_h`). Inference + varlen only. | | `FLA_INTRACARD_MAX_SPLITS` | `32` | int ≥ 1 | Max number of sub-sequences per original sequence used by the intra-card CP backend. Caps merge-chain depth to control precision loss. | diff --git a/benchmarks/ops/registry.py b/benchmarks/ops/registry.py index d5b8514d80..58a3a36e38 100644 --- a/benchmarks/ops/registry.py +++ b/benchmarks/ops/registry.py @@ -346,6 +346,32 @@ def generate_inputs( category='gate_beta', )) +register_op(OpConfig( + name='chunk_kda_inference', + import_path='fla.ops.kda', + inputs={ + 'q': TensorSpec(shape_BTHD, requires_grad=False), + 'k': TensorSpec(shape_BTHD, requires_grad=False), + 'v': TensorSpec(shape_BTHD, requires_grad=False), + 'g': TensorSpec(shape_BTHD, requires_grad=False), + 'beta': TensorSpec(shape_BTH, requires_grad=False), + 'A_log': TensorSpec(shape_H, requires_grad=False, dtype='float32'), + 'dt_bias': TensorSpec(shape_HD, requires_grad=False, dtype='float32'), + }, + func_name='chunk_kda', + extra_kwargs={ + 'use_qk_l2norm_in_kernel': True, + 'use_gate_in_kernel': True, + 'use_beta_sigmoid_in_kernel': True, + 'safe_gate': True, + 'lower_bound': -5.0, + 'state_v_first': True, + }, + skip_backward=True, + category='gate_beta', + dim_constraints={'D': [128]}, +)) + # --- +head gate (g=[B,T,H] with logsigmoid) --- register_op(OpConfig( diff --git a/benchmarks/ops/run.py b/benchmarks/ops/run.py index 0894281108..4381c097f1 100644 --- a/benchmarks/ops/run.py +++ b/benchmarks/ops/run.py @@ -281,6 +281,8 @@ def benchmark_op( config = get_op(op_name) op_fn = _import_op(config) + if config.skip_backward: + op_fn = torch.inference_mode()(op_fn) # `--backend` selects an op backend by toggling its dispatch env var (see OpConfig.backend_env), # matching how FLA backends are enabled at runtime. 'triton' (or unset) leaves the default path. @@ -338,12 +340,13 @@ def benchmark_op( inputs = generate_inputs(config, B, T, H, D, dtype=dtype, device=device_name, **extra_shape_kw) out = op_fn(**inputs, **call_kwargs) out_tensor = out[0] if config.output_is_tuple else out - do = torch.randn_like(out_tensor) + do = None if config.skip_backward else torch.randn_like(out_tensor) def _fwdbwd_fn(inputs=inputs, do=do): result = op_fn(**inputs, **call_kwargs) - t = result[0] if config.output_is_tuple else result - t.backward(do) + if do is not None: + t = result[0] if config.output_is_tuple else result + t.backward(do) _warmup_autotune(_fwdbwd_fn, device=device_name) except Exception as e: diff --git a/fla/ops/kda/backends/__init__.py b/fla/ops/kda/backends/__init__.py index cdbdaebc44..c114c6c991 100644 --- a/fla/ops/kda/backends/__init__.py +++ b/fla/ops/kda/backends/__init__.py @@ -10,12 +10,14 @@ from fla.ops.backends import BackendRegistry, dispatch from fla.ops.kda.backends.flash_kda import FlashKDABackend from fla.ops.kda.backends.tilelang import KDATileLangBackend +from fla.ops.kda.backends.tle import KDATLEBackend from fla.ops.kda.backends.triton_ascend import TritonAscendKDABackend kda_registry = BackendRegistry("kda") kda_registry.register(TritonAscendKDABackend()) kda_registry.register(FlashKDABackend()) kda_registry.register(KDATileLangBackend()) +kda_registry.register(KDATLEBackend()) __all__ = ['dispatch', 'kda_registry'] diff --git a/fla/ops/kda/backends/tle/__init__.py b/fla/ops/kda/backends/tle/__init__.py new file mode 100644 index 0000000000..b55fd9042f --- /dev/null +++ b/fla/ops/kda/backends/tle/__init__.py @@ -0,0 +1,247 @@ +# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# For a list of all contributors, visit: +# https://github.com/fla-org/flash-linear-attention/graphs/contributors +# Copyright 2026, The FlagOS Contributors. + +"""TLE backend for KDA chunk inference (BT=16, TMA + warp-specialized).""" + +from __future__ import annotations + +import warnings +from typing import TYPE_CHECKING + +import torch + +from fla.ops.backends import BaseBackend +from fla.utils import IS_NVIDIA_HOPPER + +if TYPE_CHECKING: + from fla.ops.cp import FLACPContext + + +def _has_tle() -> bool: + try: + import triton + from packaging.version import Version + ver = Version(triton.__version__.split("+")[0]) + if ver < Version("3.6.0"): + return False + import triton.experimental.tle.language # noqa: F401 + return True + except Exception: + return False + + +_TLE_AVAILABLE = _has_tle() + + +def _tle_input_error( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + initial_state: torch.Tensor | None, + state_v_first: bool, + cu_seqlens: torch.Tensor | None, + safe_gate: bool, + lower_bound: float | None, + A_log: torch.Tensor | None, + dt_bias: torch.Tensor | None, +) -> str | None: + inputs = {"q": q, "k": k, "v": v, "g": g, "beta": beta} + if any(x.dtype != torch.bfloat16 for x in inputs.values()): + actual = ", ".join(f"{name}={x.dtype}" for name, x in inputs.items()) + return f"TLE KDA requires bfloat16 inputs, got {actual}" + if any(not x.is_cuda for x in inputs.values()): + return "TLE KDA requires CUDA inputs" + if any(x.device != q.device for x in inputs.values()): + return "TLE KDA requires all inputs on the same device" + if q.ndim != 4 or k.ndim != 4 or v.ndim != 4 or g.ndim != 4 or beta.ndim != 3: + return "TLE KDA expects q/k/v/g with rank 4 and beta with rank 3" + + B, T, H, D = q.shape + if T == 0: + return "TLE KDA requires a non-empty sequence" + if D != 128: + return f"TLE KDA requires K=128, got {D}" + if k.shape != q.shape or v.shape != q.shape or g.shape != q.shape: + return "TLE KDA requires q, k, v, and g to share shape [B, T, H, 128]" + if beta.shape != (B, T, H): + return f"TLE KDA requires beta shape {(B, T, H)}, got {tuple(beta.shape)}" + + if not state_v_first: + return "TLE KDA requires state_v_first=True" + if not safe_gate: + return "TLE KDA requires safe_gate=True" + if lower_bound is None or not -5 <= lower_bound < 0: + return f"TLE KDA requires -5 <= lower_bound < 0, got {lower_bound}" + + if A_log is None or A_log.dtype != torch.float32 or A_log.shape != (H,): + actual = None if A_log is None else (tuple(A_log.shape), A_log.dtype) + return f"TLE KDA requires float32 A_log with shape {(H,)}, got {actual}" + if A_log.device != q.device: + return "TLE KDA requires A_log on the input device" + if dt_bias is not None: + if dt_bias.dtype != torch.float32 or dt_bias.shape not in ((H * D,), (H, D)): + actual = (tuple(dt_bias.shape), dt_bias.dtype) + return f"TLE KDA requires float32 dt_bias with shape {(H * D,)} or {(H, D)}, got {actual}" + if dt_bias.device != q.device: + return "TLE KDA requires dt_bias on the input device" + + N = B + if cu_seqlens is not None: + if B != 1: + return "TLE KDA requires B=1 when cu_seqlens is provided" + if ( + cu_seqlens.device != q.device + or cu_seqlens.dtype not in (torch.int32, torch.int64) + or cu_seqlens.ndim != 1 + ): + return "TLE KDA requires a 1D int32 or int64 cu_seqlens tensor on the input device" + if cu_seqlens.numel() < 2: + return "TLE KDA requires cu_seqlens to contain at least two elements" + N = cu_seqlens.numel() - 1 + + if initial_state is not None: + expected = (N, H, D, D) + if initial_state.dtype != torch.float32 or initial_state.shape != expected: + return f"TLE KDA requires float32 initial_state with shape {expected}" + if initial_state.device != q.device: + return "TLE KDA requires initial_state on the input device" + return None + + +class KDATLEBackend(BaseBackend): + """TLE-accelerated KDA chunk forward (inference only). + + Uses TMA + warp-specialized fused kernels with BT=16. + Requires Triton >= 3.6.0 with TLE extension. + """ + + backend_type = "tle" + package_name = None + env_var = "FLA_TLE_KDA" + default_enable = True + priority = 2 + + @classmethod + def is_available(cls) -> bool: + return _TLE_AVAILABLE and IS_NVIDIA_HOPPER + + def chunk_kda_verifier( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float | None = None, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + use_gate_in_kernel: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + allow_neg_eigval: bool = False, + safe_gate: bool = False, + lower_bound: float | None = None, + disable_recompute: bool = False, + return_intermediate_states: bool = False, + state_v_first: bool = False, + cu_seqlens: torch.Tensor | None = None, + cu_seqlens_cpu: torch.LongTensor | None = None, + cp_context: FLACPContext | None = None, + **kwargs, + ) -> tuple[bool, str | None]: + if torch.is_grad_enabled(): + return False, "TLE KDA only supports inference mode" + if not use_gate_in_kernel: + return False, "TLE KDA requires use_gate_in_kernel=True" + if not use_qk_l2norm_in_kernel: + return False, "TLE KDA requires use_qk_l2norm_in_kernel=True" + if not use_beta_sigmoid_in_kernel: + return False, "TLE KDA requires use_beta_sigmoid_in_kernel=True" + if allow_neg_eigval: + return False, "TLE KDA does not support allow_neg_eigval=True" + if cp_context is not None: + return False, "TLE KDA does not support context parallel" + if return_intermediate_states: + return False, "TLE KDA does not support return_intermediate_states" + if "transpose_state_layout" in kwargs: + if state_v_first: + return False, "Cannot pass both state_v_first and transpose_state_layout" + state_v_first = kwargs["transpose_state_layout"] + chunk_size = kwargs.get("chunk_size", 64) + if chunk_size not in (32, 64): + return False, f"chunk_size must be either 32 or 64, got {chunk_size}" + reason = _tle_input_error( + q=q, + k=k, + v=v, + g=g, + beta=beta, + initial_state=initial_state, + state_v_first=state_v_first, + cu_seqlens=cu_seqlens, + safe_gate=safe_gate, + lower_bound=lower_bound, + A_log=kwargs.get("A_log"), + dt_bias=kwargs.get("dt_bias"), + ) + return reason is None, reason + + def chunk_kda( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float | None = None, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + use_gate_in_kernel: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + allow_neg_eigval: bool = False, + safe_gate: bool = False, + lower_bound: float | None = None, + disable_recompute: bool = False, + return_intermediate_states: bool = False, + state_v_first: bool = False, + cu_seqlens: torch.Tensor | None = None, + cu_seqlens_cpu: torch.LongTensor | None = None, + cp_context: FLACPContext | None = None, + **kwargs, + ): + from fla.ops.kda.backends.tle.chunk_kda import chunk_kda_fwd_infer + + if "transpose_state_layout" in kwargs: + if state_v_first: + raise ValueError("Cannot pass both `state_v_first` and the deprecated `transpose_state_layout`.") + warnings.warn( + "`transpose_state_layout` is deprecated and renamed to `state_v_first`.", + DeprecationWarning, + stacklevel=2, + ) + state_v_first = kwargs.pop("transpose_state_layout") + + return chunk_kda_fwd_infer( + q=q, + k=k, + v=v, + g=g, + beta=beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + state_v_first=state_v_first, + cu_seqlens=cu_seqlens, + safe_gate=safe_gate, + lower_bound=lower_bound, + A_log=kwargs.get("A_log"), + dt_bias=kwargs.get("dt_bias"), + ) diff --git a/fla/ops/kda/backends/tle/chunk_kda.py b/fla/ops/kda/backends/tle/chunk_kda.py new file mode 100644 index 0000000000..7c7924c72a --- /dev/null +++ b/fla/ops/kda/backends/tle/chunk_kda.py @@ -0,0 +1,787 @@ +# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li +# +# This source code is licensed under the MIT license found in the +# LICENSE file in the root directory of this source tree. +# For a list of all contributors, visit: +# https://github.com/fla-org/flash-linear-attention/graphs/contributors +# Copyright 2026, The FlagOS Contributors. + +"""TLE forward path for KDA chunk inference.""" + +from __future__ import annotations + +import torch +import triton +import triton.experimental.tle.language as tle +import triton.language as tl +from triton.runtime._allocation import NullAllocator, _allocator +from triton.tools.tensor_descriptor import TensorDescriptor + +from fla.ops.kda.backends.tle import _tle_input_error +from fla.ops.utils.index import prepare_chunk_indices, prepare_chunk_offsets +from fla.utils import autotune_cache_kwargs, input_guard +from fla.utils._device import _default_alloc_fn + +__all__ = ["chunk_kda_fwd_infer"] + + +RCP_LN2 = 1.4426950216 + + +@triton.jit +def exp2(x): + return tl.math.exp2(x.to(tl.float32)) + + +@triton.heuristics( + { + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + "HAS_DT_BIAS": lambda args: args["dt_bias"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({}, num_warps=num_warps, num_stages=num_stages, maxnreg=maxnreg) + for num_warps in [2, 4, 8] + for num_stages in [2, 4, 8] + for maxnreg in [None, 32, 64, 72] + ], + key=["H", "K", "BT", "IS_VARLEN", "HAS_DT_BIAS"], + **autotune_cache_kwargs, +) +@triton.jit(do_not_specialize=["T"]) +def _kda_fwd_intra_kernel( + q, + k, + g, + beta, + ws, + Aqk, + Akk, + g_last, + A_log, + dt_bias, + lower_bound, + scale, + g_scale, + l2norm_eps, + cu_seqlens, + chunk_indices, + T, + NT_TOTAL, + H: tl.constexpr, + K: tl.constexpr, + BT: tl.constexpr, + IS_VARLEN: tl.constexpr, + HAS_DT_BIAS: tl.constexpr, +): + i_chunk_pid = tl.program_id(0).to(tl.int64) + i_bh = tl.program_id(1).to(tl.int64) + i_h = i_bh % H + + if IS_VARLEN: + i_chunk_global = i_chunk_pid + i_n = tl.load(chunk_indices + i_chunk_pid * 2).to(tl.int64) + i_chunk = tl.load(chunk_indices + i_chunk_pid * 2 + 1).to(tl.int64) + bos = tl.load(cu_seqlens + i_n).to(tl.int64) + eos = tl.load(cu_seqlens + i_n + 1).to(tl.int64) + T = (eos - bos).to(tl.int32) + else: + i_n = i_bh // H + i_chunk = i_chunk_pid + bos = i_n.to(tl.int64) * T + i_chunk_global = i_n * tl.cdiv(T, BT) + i_chunk + + if i_chunk * BT >= T: + return + + q += (bos * H + i_h) * K + k += (bos * H + i_h) * K + g += (bos * H + i_h) * K + g_last += (i_chunk_global * H + i_h).to(tl.int64) * K + if IS_VARLEN: + a_chunk = i_h * NT_TOTAL + i_chunk_global + else: + a_chunk = (i_n * H + i_h) * NT_TOTAL + i_chunk + Aqk += a_chunk.to(tl.int64) * BT * BT + Akk += a_chunk.to(tl.int64) * BT * BT + ws += (bos * H + i_h) * 3 * K + beta += bos * H + i_h + + o_i = tl.arange(0, BT) + o_i64 = o_i.to(tl.int64) + o_k = tl.arange(0, K) + o_k64 = o_k.to(tl.int64) + token_start = i_chunk * BT + m_c = token_start + o_i < T + + q_buf = tle.gpu.alloc([BT, K], dtype=q.dtype.element_ty, scope=tle.gpu.smem) + k_buf = tle.gpu.alloc([BT, K], dtype=k.dtype.element_ty, scope=tle.gpu.smem) + gc_buf = tle.gpu.alloc([BT, K], dtype=tl.float32, scope=tle.gpu.smem) + + rows = tl.broadcast_to(o_i[:, None], (BT, K)) + cols = tl.broadcast_to(o_k[None, :], (BT, K)) + q_sp = tle.gpu.local_ptr(q_buf, (rows, cols)) + k_sp = tle.gpu.local_ptr(k_buf, (rows, cols)) + gc_sp = tle.gpu.local_ptr(gc_buf, (rows, cols)) + + offsets_qkg = (token_start + o_i64)[:, None] * H * K + o_k64[None, :] + b_q = tl.load(q + offsets_qkg, mask=m_c[:, None], other=0.0) + b_k = tl.load(k + offsets_qkg, mask=m_c[:, None], other=0.0) + tl.store(q_sp, b_q) + tl.store(k_sp, b_k) + + b_qf = b_q.to(tl.float32) + b_kf = b_k.to(tl.float32) + + b_q_rstd = 1.0 / tl.sqrt(tl.sum(b_qf * b_qf, 1) + l2norm_eps) + b_k_rstd = 1.0 / tl.sqrt(tl.sum(b_kf * b_kf, 1) + l2norm_eps) + + b_g = tl.load(g + offsets_qkg, mask=m_c[:, None], other=0.0).to(tl.float32) + b_A = exp2(tl.load(A_log + i_h).to(tl.float32) * g_scale) + if HAS_DT_BIAS: + b_bias = tl.load(dt_bias + i_h * K + o_k64).to(tl.float32) + b_g += b_bias[None, :] + b_g = (lower_bound * g_scale) * tl.sigmoid(b_A * b_g) + tl.store(gc_sp, b_g) + one_row = tl.broadcast_to(tl.arange(0, 1)[:, None], (1, K)) + col_row = tl.broadcast_to(tl.arange(0, K)[None, :], (1, K)) + b_acc = tl.zeros([1, K], dtype=tl.float32) + for r in tl.static_range(BT): + rp = tle.gpu.local_ptr(gc_buf, (tl.broadcast_to(one_row + r, (1, K)), col_row)) + b_acc = b_acc + tl.load(rp) + tl.store(rp, b_acc) + + b_gq = tl.where(m_c[:, None], exp2(tl.load(gc_sp)), 0.0) + b_gk = tl.where(m_c[:, None], exp2(-tl.load(gc_sp)), 0.0) + + b_kgt = tl.trans(b_kf * b_gk).to(b_k.dtype) + b_Aqk = tl.dot( + (b_qf * b_gq).to(b_q.dtype), + b_kgt, + out_dtype=tl.float32, + ) + b_Akk = tl.dot( + (b_kf * b_gq).to(b_k.dtype), + b_kgt, + out_dtype=tl.float32, + ) + + b_Aqk = b_Aqk * b_q_rstd[:, None] * b_k_rstd[None, :] + b_Akk = b_Akk * b_k_rstd[:, None] * b_k_rstd[None, :] + + b_beta = tl.sigmoid(tl.load(beta + (token_start + o_i64) * H, mask=m_c, other=0.0).to(tl.float32)) + + m_Aqk = o_i[:, None] >= o_i[None, :] + m_Akk = o_i[:, None] > o_i[None, :] + m_I = o_i[:, None] == o_i[None, :] + + b_Aqk = tl.where(m_Aqk, b_Aqk * scale, 0.0) + b_Akk = tl.where(m_Akk, b_Akk * b_beta[:, None], 0.0) + + offsets_A = o_i64[:, None] * BT + o_i64[None, :] + tl.store(Aqk + offsets_A, b_Aqk.to(Aqk.dtype.element_ty)) + + b_L = b_Akk.to(tl.float16) + b_Ai = m_I.to(tl.float16) - b_L + b_L2 = tl.dot(b_L, b_L, out_dtype=tl.float16) + b_Ai = b_Ai + tl.dot(b_Ai, b_L2, out_dtype=tl.float16) + b_L4 = tl.dot(b_L2, b_L2, out_dtype=tl.float16) + b_Ai = b_Ai + tl.dot(b_Ai, b_L4, out_dtype=tl.float16) + b_L8 = tl.dot(b_L4, b_L4, out_dtype=tl.float16) + b_Ai = b_Ai + tl.dot(b_Ai, b_L8, out_dtype=tl.float16) + + tl.store(Akk + offsets_A, b_Ai.to(Akk.dtype.element_ty)) + + b_k3 = tl.load(k_sp).to(tl.float32) * b_k_rstd[:, None] + b_gk3 = tl.load(gc_sp) + b_kb = b_k3 * b_beta[:, None] * exp2(b_gk3) + offsets_ws = (token_start + o_i64)[:, None] * H * 3 * K + o_k64[None, :] + tl.store(ws + offsets_ws, b_kb.to(ws.dtype.element_ty), mask=m_c[:, None]) + + b_qg_val = tl.load(q_sp).to(tl.float32) * b_q_rstd[:, None] * exp2(b_gk3) + tl.store(ws + offsets_ws + K, b_qg_val.to(ws.dtype.element_ty), mask=m_c[:, None]) + + last_local = (tl.minimum(BT, T - token_start) - 1).to(tl.int32) + gn_rows = tl.broadcast_to(last_local + tl.zeros([1, K], dtype=tl.int32), (1, K)) + gn_cols = tl.broadcast_to(tl.arange(0, K)[None, :], (1, K)) + b_gn = tl.load(tle.gpu.local_ptr(gc_buf, (gn_rows, gn_cols))) + tl.store(g_last + o_k64, b_gn.reshape([K]).to(g_last.dtype.element_ty)) + b_kg_val = b_k3 * tl.where(m_c[:, None], exp2(b_gn - b_gk3), 0) + tl.store(ws + offsets_ws + 2 * K, b_kg_val.to(ws.dtype.element_ty), mask=m_c[:, None]) + + +def _kda_fwd_intra( + q, + k, + g, + beta, + scale, + cu_seqlens=None, + chunk_indices=None, + chunk_size=16, + lower_bound=None, + A_log=None, + dt_bias=None, +): + B, T_len, H, K = q.shape + BT = chunk_size + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, BT) + NT = triton.cdiv(T_len, BT) if cu_seqlens is None else len(chunk_indices) + grid = (NT, B * H) + + g_last = torch.empty(B * NT, H, K, device=q.device, dtype=torch.float32) + ws = torch.empty(B, NT * BT, H, 3 * K, device=q.device, dtype=q.dtype) + Aqk = torch.empty(B, H, NT, BT, BT, device=q.device, dtype=q.dtype) + Akk = torch.empty(B, H, NT, BT, BT, device=q.device, dtype=q.dtype) + + _kda_fwd_intra_kernel[grid]( + q=q, + k=k, + g=g, + beta=beta, + ws=ws, + Aqk=Aqk, + Akk=Akk, + g_last=g_last, + A_log=A_log, + dt_bias=dt_bias, + lower_bound=lower_bound, + scale=scale, + g_scale=RCP_LN2, + l2norm_eps=1e-6, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T_len, + NT_TOTAL=NT, + H=H, + K=K, + BT=BT, + ) + return ws, Aqk, Akk, g_last + + +@triton.jit +def _kda_state_output_load_producer( + writer, + ws_desc, + v_ptr, + beta_ptr, + gk_desc, + Aqk_desc, + Akk_desc, + K: tl.constexpr, + T, + H: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BV: tl.constexpr, + NT, + i_v, + USE_HOST_DESCRIPTORS: tl.constexpr, + state_row_start, + chunk_start, + i_h, +): + for i_chunk in tl.range(NT): + slot = writer.acquire(i_chunk) + + ws_row = i_chunk * BT + ws_col = 0 + A_row = i_chunk * BT + gk_row = i_chunk + gk_col = 0 + if USE_HOST_DESCRIPTORS: + ws_row += state_row_start + ws_col = i_h * 3 * K + A_row = state_row_start * H + i_h * NT * BT + i_chunk * BT + gk_row += chunk_start + gk_col = i_h * K + + ws_row = ws_row.to(tl.int32) + ws_col = ws_col.to(tl.int32) + A_row = A_row.to(tl.int32) + gk_row = gk_row.to(tl.int32) + gk_col = gk_col.to(tl.int32) + + tle.gpu.copy(ws_desc, slot.w, [BT, K], [ws_row, ws_col]) + tle.gpu.copy(ws_desc, slot.qg, [BT, K], [ws_row, ws_col + K]) + tle.gpu.copy(ws_desc, slot.kg, [BT, K], [ws_row, ws_col + 2 * K]) + tle.gpu.copy(Aqk_desc, slot.Aqk, [BT, BT], [A_row, 0]) + tle.gpu.copy(Akk_desc, slot.Akk, [BT, BT], [A_row, 0]) + tle.gpu.copy(gk_desc, slot.gk, [1, K], [gk_row, gk_col]) + + o_t = tl.arange(0, BT) + o_t64 = o_t.to(tl.int64) + o_v = tl.arange(0, BV) + o_v64 = o_v.to(tl.int64) + token = i_chunk.to(tl.int64) * BT + o_t64 + value = i_v.to(tl.int64) * BV + o_v64 + mask_t = token < T + b_v_raw = tl.load( + v_ptr + token[:, None] * H * V + value[None, :], + mask=mask_t[:, None] & (value[None, :] < V), + other=0.0, + ) + b_beta = tl.sigmoid(tl.load(beta_ptr + token * H, mask=mask_t, other=0.0).to(tl.float32)) + b_v = (b_v_raw.to(tl.float32) * b_beta[:, None]).to(b_v_raw.dtype) + tl.store(tle.gpu.local_ptr(slot.v), b_v) + + writer.commit(i_chunk) + + +@triton.jit +def _kda_state_output_mma_consumer( + load_reader, + store_writer, + h0, + ht, + scale, + i_v, + i_nh, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BV: tl.constexpr, + NT, + USE_INITIAL_STATE: tl.constexpr, + STORE_FINAL_STATE: tl.constexpr, +): + state_dtype: tl.constexpr = tl.bfloat16 + if USE_INITIAL_STATE: + o_v = i_v.to(tl.int64) * BV + tl.arange(0, BV).to(tl.int64) + o_k = tl.arange(0, K).to(tl.int64) + offsets_h = i_nh.to(tl.int64) * K * V + o_v[:, None] * K + o_k[None, :] + b_h = tl.trans(tl.load(h0 + offsets_h, mask=o_v[:, None] < V, other=0.0)).to(tl.float32) + else: + b_h = tl.zeros([K, BV], dtype=tl.float32) + + for i_chunk in tl.range(NT): + wait = load_reader.wait(i_chunk) + slot = wait.slot + + b_w = tl.load(tle.gpu.local_ptr(slot.w)) + b_v_raw = tl.load(tle.gpu.local_ptr(slot.v)) + b_qg = tl.load(tle.gpu.local_ptr(slot.qg)) + b_kg = tl.load(tle.gpu.local_ptr(slot.kg)) + b_Aqk = tl.load(tle.gpu.local_ptr(slot.Aqk)) + b_Akk = tl.load(tle.gpu.local_ptr(slot.Akk)) + b_gk = tl.load(tle.gpu.local_ptr(slot.gk)).reshape([K]) + + b_h_bf = b_h.to(state_dtype) + + b_kh = tl.dot(b_w, b_h_bf).to(tl.float32) + b_diff = b_v_raw.to(tl.float32) - b_kh + b_v = tl.dot(b_Akk, b_diff.to(state_dtype)).to(tl.float32) + + b_qh = tl.dot(b_qg, b_h_bf).to(tl.float32) + b_o = scale * b_qh + b_v_cast = b_v.to(state_dtype) + b_o += tl.dot(b_Aqk, b_v_cast).to(tl.float32) + + out_slot = store_writer.acquire(i_chunk) + tl.store(tle.gpu.local_ptr(out_slot.output), b_o) + store_writer.commit(i_chunk) + + load_reader.release(i_chunk) + + b_h = b_h * exp2(b_gk)[:, None] + tl.dot(tl.trans(b_kg), b_v_cast).to(tl.float32) + + if STORE_FINAL_STATE: + o_v = i_v.to(tl.int64) * BV + tl.arange(0, BV).to(tl.int64) + o_k = tl.arange(0, K).to(tl.int64) + offsets_ht = i_nh.to(tl.int64) * K * V + o_v[:, None] * K + o_k[None, :] + tl.store(ht + offsets_ht, tl.trans(b_h).to(ht.dtype.element_ty), mask=o_v[:, None] < V) + + +@triton.jit +def _kda_state_output_store_consumer( + store_reader, + store_target, + BT: tl.constexpr, + BV: tl.constexpr, + NT, + i_v, + USE_HOST_DESCRIPTORS: tl.constexpr, + output_row_start, + output_col_start, +): + for i_chunk in tl.range(NT): + store_wait = store_reader.wait(i_chunk) + slot = store_wait.slot + output_row = i_chunk * BT + output_col = i_v * BV + if USE_HOST_DESCRIPTORS: + output_row += output_row_start + output_col += output_col_start + output_row = output_row.to(tl.int32) + output_col = output_col.to(tl.int32) + tle.gpu.copy(slot.output, store_target, [BT, BV], [output_row, output_col]) + store_reader.release(i_chunk) + + +PIPE_STAGES = tl.constexpr(4) + + +@triton.heuristics( + { + "USE_INITIAL_STATE": lambda args: args["h0"].numel() > 1, + "STORE_FINAL_STATE": lambda args: args["ht"].numel() > 1, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + } +) +@triton.jit(do_not_specialize=["T"]) +def _kda_fwd_state_output_kernel( + v, + beta, + gk, + Aqk, + Akk, + o, + ws, + h0, + ht, + ws_host_desc, + gk_host_desc, + Aqk_host_desc, + Akk_host_desc, + output_host_desc, + cu_seqlens, + chunk_offsets, + scale, + T, + NT_TOTAL, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BV: tl.constexpr, + USE_HOST_DESCRIPTORS: tl.constexpr, + USE_INITIAL_STATE: tl.constexpr, + STORE_FINAL_STATE: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_v = tl.program_id(0).to(tl.int64) + i_nh = tl.program_id(1).to(tl.int64) + i_n = i_nh // H + i_h = i_nh % H + + if IS_VARLEN: + bos = tl.load(cu_seqlens + i_n).to(tl.int64) + eos = tl.load(cu_seqlens + i_n + 1).to(tl.int64) + T = (eos - bos).to(tl.int32) + NT = tl.cdiv(T, BT) + chunk_start = tl.load(chunk_offsets + i_n).to(tl.int32) + state_row_start = bos + else: + bos = i_n.to(tl.int64) * T + NT = tl.cdiv(T, BT) + chunk_start = i_n * NT + state_row_start = i_n.to(tl.int64) * NT * BT + + v += (bos * H + i_h) * V + beta += bos * H + i_h + gk += (chunk_start * H + i_h).to(tl.int64) * K + o += (bos * H + i_h) * V + ws_base = ws + (bos * H + i_h) * 3 * K + + if USE_HOST_DESCRIPTORS: + ws_desc = ws_host_desc + gk_desc = gk_host_desc + Aqk_desc = Aqk_host_desc + Akk_desc = Akk_host_desc + else: + if IS_VARLEN: + a_chunk = (i_h * NT_TOTAL + chunk_start) * BT * BT + else: + a_chunk = (i_n * H + i_h) * NT_TOTAL * BT * BT + Aqk += a_chunk.to(tl.int64) + Akk += a_chunk.to(tl.int64) + ws_desc = tl.make_tensor_descriptor(ws_base, shape=[T, 3 * K], strides=[H * 3 * K, 1], block_shape=[BT, K]) + gk_desc = tl.make_tensor_descriptor(gk, shape=[NT, K], strides=[H * K, 1], block_shape=[1, K]) + Aqk_desc = tl.make_tensor_descriptor(Aqk, shape=[NT * BT, BT], strides=[BT, 1], block_shape=[BT, BT]) + Akk_desc = tl.make_tensor_descriptor(Akk, shape=[NT * BT, BT], strides=[BT, 1], block_shape=[BT, BT]) + + w_smem = tle.gpu.alloc([PIPE_STAGES, BT, K], dtype=tl.bfloat16, scope=tle.gpu.smem) + v_smem = tle.gpu.alloc([PIPE_STAGES, BT, BV], dtype=tl.bfloat16, scope=tle.gpu.smem) + qg_smem = tle.gpu.alloc([PIPE_STAGES, BT, K], dtype=tl.bfloat16, scope=tle.gpu.smem) + kg_smem = tle.gpu.alloc([PIPE_STAGES, BT, K], dtype=tl.bfloat16, scope=tle.gpu.smem) + Aqk_smem = tle.gpu.alloc([PIPE_STAGES, BT, BT], dtype=tl.bfloat16, scope=tle.gpu.smem) + Akk_smem = tle.gpu.alloc([PIPE_STAGES, BT, BT], dtype=tl.bfloat16, scope=tle.gpu.smem) + gk_smem = tle.gpu.alloc([PIPE_STAGES, 1, K], dtype=tl.float32, scope=tle.gpu.smem) + out_smem = tle.gpu.alloc([PIPE_STAGES, BT, BV], dtype=tl.bfloat16, scope=tle.gpu.smem) + if USE_HOST_DESCRIPTORS: + output_store_target = output_host_desc + else: + output_store_target = tl.make_tensor_descriptor( + o, + shape=[T, V], + strides=[H * V, 1], + block_shape=[BT, BV], + ) + + load_pipe = tle.pipe( + capacity=PIPE_STAGES, + scope="cta", + name="kda_load", + w=w_smem, + v=v_smem, + qg=qg_smem, + kg=kg_smem, + Aqk=Aqk_smem, + Akk=Akk_smem, + gk=gk_smem, + ) + store_pipe = tle.pipe( + capacity=PIPE_STAGES, + scope="cta", + name="kda_store", + output=out_smem, + ) + tle.gpu.warp_specialize( + [ + ( + _kda_state_output_load_producer, + ( + load_pipe.writer(), + ws_desc, + v, + beta, + gk_desc, + Aqk_desc, + Akk_desc, + K, + T, + H, + V, + BT, + BV, + NT, + i_v, + USE_HOST_DESCRIPTORS, + state_row_start, + chunk_start, + i_h, + ), + ), + ( + _kda_state_output_mma_consumer, + ( + load_pipe.reader(), + store_pipe.writer(), + h0, + ht, + scale, + i_v, + i_nh, + H, + K, + V, + BV, + NT, + USE_INITIAL_STATE, + STORE_FINAL_STATE, + ), + ), + ( + _kda_state_output_store_consumer, + ( + store_pipe.reader(), + output_store_target, + BT, + BV, + NT, + i_v, + USE_HOST_DESCRIPTORS, + bos.to(tl.int64), + i_h.to(tl.int64) * V, + ), + ), + ], + [4, 1], + [240, 32], + ) + + +def _kda_fwd_state_output( + v: torch.Tensor, + beta: torch.Tensor, + Akk: torch.Tensor, + gk: torch.Tensor, + Aqk: torch.Tensor, + ws: torch.Tensor, + scale: float, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + cu_seqlens: torch.Tensor | None = None, + chunk_size: int = 16, +) -> tuple[torch.Tensor, torch.Tensor | None]: + B, _, H, packed_K = ws.shape + K = packed_K // 3 + T_actual = v.shape[1] + V = v.shape[-1] + BT = chunk_size + + if cu_seqlens is None: + N = B + chunk_offsets = None + else: + N = len(cu_seqlens) - 1 + chunk_offsets = prepare_chunk_offsets(cu_seqlens, BT) + + final_state = ws.new_empty(N, H, V, K, dtype=torch.float32) if output_final_state else None + + o = torch.empty(B, T_actual, H, V, device=ws.device, dtype=v.dtype) + + h0_arg = initial_state if initial_state is not None else ws.new_empty(1, dtype=torch.float32) + ht_arg = final_state if final_state is not None else ws.new_empty(1, dtype=torch.float32) + + use_host_descriptors = cu_seqlens is None and T_actual % BT == 0 + if use_host_descriptors: + NT_total = ws.shape[1] // BT + descriptor_rows = B * H * NT_total * BT + ws_desc_arg = TensorDescriptor( + ws, + shape=[descriptor_rows, H * 3 * K], + strides=[H * 3 * K, 1], + block_shape=[BT, K], + ) + gk_desc_arg = TensorDescriptor( + gk, + shape=[gk.shape[0], H * K], + strides=[H * K, 1], + block_shape=[1, K], + ) + Aqk_desc_arg = TensorDescriptor( + Aqk, + shape=[descriptor_rows, BT], + strides=[BT, 1], + block_shape=[BT, BT], + ) + Akk_desc_arg = TensorDescriptor( + Akk, + shape=[descriptor_rows, BT], + strides=[BT, 1], + block_shape=[BT, BT], + ) + output_desc_arg = TensorDescriptor( + o, + shape=[B * T_actual, H * V], + strides=[H * V, 1], + block_shape=[BT, 128], + ) + else: + ws_desc_arg = ws + gk_desc_arg = gk + Aqk_desc_arg = Aqk + Akk_desc_arg = Akk + output_desc_arg = o + + _kda_fwd_state_output_kernel[(1, N * H)]( + v=v, + beta=beta, + gk=gk, + Aqk=Aqk, + Akk=Akk, + o=o, + ws=ws, + h0=h0_arg, + ht=ht_arg, + ws_host_desc=ws_desc_arg, + gk_host_desc=gk_desc_arg, + Aqk_host_desc=Aqk_desc_arg, + Akk_host_desc=Akk_desc_arg, + output_host_desc=output_desc_arg, + cu_seqlens=cu_seqlens, + chunk_offsets=chunk_offsets, + scale=scale, + T=T_actual, + NT_TOTAL=ws.shape[1] // BT, + H=H, + K=K, + V=V, + BT=BT, + BV=128, + USE_HOST_DESCRIPTORS=use_host_descriptors, + num_warps=4, + ) + + return o, final_state + + +@input_guard +def chunk_kda_fwd_infer( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float | None = None, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + state_v_first: bool = False, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = 16, + safe_gate: bool = False, + lower_bound: float | None = None, + A_log: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor | None]: + if torch.is_grad_enabled(): + raise RuntimeError("TLE KDA only supports inference mode") + reason = _tle_input_error( + q=q, + k=k, + v=v, + g=g, + beta=beta, + initial_state=initial_state, + state_v_first=state_v_first, + cu_seqlens=cu_seqlens, + safe_gate=safe_gate, + lower_bound=lower_bound, + A_log=A_log, + dt_bias=dt_bias, + ) + if reason is not None: + raise ValueError(reason) + + if isinstance(_allocator.get(), NullAllocator): + triton.set_allocator(_default_alloc_fn) + + if scale is None: + scale = q.shape[-1] ** -0.5 + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) + + ws, Aqk, Akk, g_last = _kda_fwd_intra( + q=q, + k=k, + g=g, + beta=beta, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + lower_bound=lower_bound, + A_log=A_log, + dt_bias=dt_bias, + ) + + return _kda_fwd_state_output( + v=v, + beta=beta, + Akk=Akk, + gk=g_last, + Aqk=Aqk, + ws=ws, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_size=chunk_size, + ) diff --git a/tests/ops/test_kda.py b/tests/ops/test_kda.py index 00de429537..adafbda7a8 100644 --- a/tests/ops/test_kda.py +++ b/tests/ops/test_kda.py @@ -12,6 +12,7 @@ import torch.nn.functional as F from fla.ops.kda import chunk_kda, fused_recurrent_kda +from fla.ops.kda.backends.tle import KDATLEBackend from fla.ops.kda.fused_recurrent import fused_recurrent_kda_fwd from fla.ops.kda.gate import fused_kda_gate, naive_kda_gate, naive_kda_lowerbound_gate from fla.ops.kda.naive import naive_chunk_kda, naive_recurrent_kda @@ -1212,6 +1213,7 @@ def _flash_kda_make_gate_params(H, D): def _flash_kda_run(monkeypatch, **kwargs): + monkeypatch.setenv("FLA_TLE_KDA", "0") monkeypatch.setenv("FLA_FLASH_KDA", "1") with torch.inference_mode(): return chunk_kda(**kwargs, **_FLASH_KDA_REQUIRED_KWARGS) @@ -1446,3 +1448,223 @@ def test_triton_ascend_backend_routing(): finally: for name in _TRITON_ASCEND_KDA_OPS: delattr(backend, name) + + +# TLE backend (inference-only) + +_SKIP_TLE_KDA = pytest.mark.skipif( + not KDATLEBackend.is_available(), + reason="TLE KDA backend requires an NVIDIA Hopper GPU and Triton TLE", +) + + +def _tle_kda_make_gate_params(H, D): + A_log = torch.log(torch.empty(H, dtype=torch.float32, device=device).uniform_(1, 16)) + dt_bias = torch.randn(H, D, dtype=torch.float32, device=device) + return A_log, dt_bias + + +def _tle_kda_run(monkeypatch, **kwargs): + monkeypatch.setenv("FLA_TLE_KDA", "1") + monkeypatch.setenv("FLA_FLASH_KDA", "0") + monkeypatch.setenv("FLA_TILELANG", "0") + dispatched = [] + impl = KDATLEBackend.chunk_kda + + def spy(self, *args, **kw): + dispatched.append(True) + return impl(self, *args, **kw) + + monkeypatch.setattr(KDATLEBackend, "chunk_kda", spy) + required_kwargs = dict(_FLASH_KDA_REQUIRED_KWARGS) + if "transpose_state_layout" in kwargs: + required_kwargs.pop("state_v_first") + with torch.inference_mode(): + out = chunk_kda(**kwargs, **required_kwargs) + assert dispatched, "TLE KDA backend was not dispatched; the test would only compare the Triton fallback" + return out + + +@_SKIP_TLE_KDA +@pytest.mark.parametrize( + ("B", "T", "H", "D"), + [ + pytest.param(*test, id="B{}-T{}-H{}-D{}".format(*test)) + for test in [ + (1, 257, 4, 128), + (2, 1024, 8, 128), + ] + ], +) +def test_tle_kda_chunk(B, T, H, D, monkeypatch): + torch.manual_seed(42) + dtype = torch.bfloat16 + q = torch.rand(B, T, H, D, dtype=dtype, device=device) + k = torch.rand(B, T, H, D, dtype=dtype, device=device) + v = torch.rand(B, T, H, D, dtype=dtype, device=device) + g = torch.randn(B, T, H, D, dtype=dtype, device=device) + beta = torch.randn(B, T, H, dtype=dtype, device=device) + A_log, dt_bias = _tle_kda_make_gate_params(H, D) + h0 = torch.randn(B, H, D, D, dtype=torch.float32, device=device) + scale = D ** -0.5 + + ref_o, ref_ht = _flash_kda_gold( + q, k, v, g, beta, A_log, dt_bias, scale, h0.clone()) + tri_o, tri_ht = _tle_kda_run( + monkeypatch, + q=q, k=k, v=v, g=g, beta=beta, + A_log=A_log, dt_bias=dt_bias, + scale=scale, + initial_state=h0.clone(), + output_final_state=True, + ) + assert_close("o", ref_o, tri_o, _FLASH_KDA_RTOL) + assert_close("ht", ref_ht, tri_ht.to(ref_ht.dtype), _FLASH_KDA_RTOL) + + +@_SKIP_TLE_KDA +@pytest.mark.parametrize("bias_kind", ["none", "flat"]) +def test_tle_kda_input_contract(bias_kind, monkeypatch): + torch.manual_seed(42) + B, T, H, D = 1, 65, 2, 128 + + def make_noncontiguous(*shape, dtype): + return torch.randn(*shape, 2, dtype=dtype, device=device)[..., 0] + + q = make_noncontiguous(B, T, H, D, dtype=torch.bfloat16) + k = make_noncontiguous(B, T, H, D, dtype=torch.bfloat16) + v = make_noncontiguous(B, T, H, D, dtype=torch.bfloat16) + g = make_noncontiguous(B, T, H, D, dtype=torch.bfloat16) + beta = make_noncontiguous(B, T, H, dtype=torch.bfloat16) + A_log = torch.log(make_noncontiguous(H, dtype=torch.float32).abs() + 1) + dt_bias = None if bias_kind == "none" else make_noncontiguous(H * D, dtype=torch.float32) + h0 = make_noncontiguous(B, H, D, D, dtype=torch.float32) + ref_bias = torch.zeros(H * D, dtype=torch.float32, device=device) if dt_bias is None else dt_bias + + ref_o, ref_ht = _flash_kda_gold( + q, k, v, g, beta, A_log, ref_bias, D ** -0.5, h0.clone()) + tri_o, tri_ht = _tle_kda_run( + monkeypatch, + q=q, k=k, v=v, g=g, beta=beta, + A_log=A_log, dt_bias=dt_bias, + scale=D ** -0.5, + initial_state=h0.clone(), + output_final_state=True, + ) + assert_close("o", ref_o, tri_o, _FLASH_KDA_RTOL) + assert_close("ht", ref_ht, tri_ht.to(ref_ht.dtype), _FLASH_KDA_RTOL) + + +@_SKIP_TLE_KDA +def test_tle_kda_transpose_state_layout(monkeypatch): + torch.manual_seed(42) + B, T, H, D = 1, 65, 2, 128 + q = torch.randn(B, T, H, D, dtype=torch.bfloat16, device=device) + k = torch.randn_like(q) + v = torch.randn_like(q) + g = torch.randn_like(q) + beta = torch.randn(B, T, H, dtype=torch.bfloat16, device=device) + A_log, dt_bias = _tle_kda_make_gate_params(H, D) + h0 = torch.randn(B, H, D, D, dtype=torch.float32, device=device) + + with pytest.warns(DeprecationWarning, match="transpose_state_layout"): + out, ht = _tle_kda_run( + monkeypatch, + q=q, k=k, v=v, g=g, beta=beta, + A_log=A_log, dt_bias=dt_bias, + initial_state=h0, + output_final_state=True, + transpose_state_layout=True, + ) + assert out.shape == v.shape + assert ht.shape == h0.shape + + +@_SKIP_TLE_KDA +@pytest.mark.parametrize( + ("H", "D", "cu_seqlens"), + [ + pytest.param(H, D, cu, id=f"H{H}-D{D}-cu{cu}") + for (H, D, cu) in [ + (4, 128, [0, 17, 129, 257]), + (8, 128, [0, 101, 303, 1205]), + ] + ], +) +def test_tle_kda_chunk_varlen(H, D, cu_seqlens, monkeypatch): + torch.manual_seed(42) + dtype = torch.bfloat16 + cu_seqlens_t = torch.LongTensor(cu_seqlens).to(device) + T = cu_seqlens[-1] + N = len(cu_seqlens) - 1 + + q = torch.randn(1, T, H, D, dtype=dtype, device=device) + k = torch.randn(1, T, H, D, dtype=dtype, device=device) + v = torch.randn(1, T, H, D, dtype=dtype, device=device) + g = torch.randn(1, T, H, D, dtype=dtype, device=device) + beta = torch.randn(1, T, H, dtype=dtype, device=device) + A_log, dt_bias = _tle_kda_make_gate_params(H, D) + h0 = torch.randn(N, H, D, D, dtype=torch.float32, device=device) + scale = D ** -0.5 + + ref_o, ref_ht = _flash_kda_gold( + q, k, v, g, beta, A_log, dt_bias, scale, h0.clone(), + cu_seqlens=cu_seqlens_t, + ) + tri_o, tri_ht = _tle_kda_run( + monkeypatch, + q=q, k=k, v=v, g=g, beta=beta, + A_log=A_log, dt_bias=dt_bias, + scale=scale, + initial_state=h0.clone(), + output_final_state=True, + cu_seqlens=cu_seqlens_t, + ) + assert_close("o", ref_o, tri_o, _FLASH_KDA_RTOL) + assert_close("ht", ref_ht, tri_ht.to(ref_ht.dtype), _FLASH_KDA_RTOL) + + +@_SKIP_TLE_KDA +def test_tle_kda_fallback(monkeypatch): + torch.manual_seed(42) + B, T, H, D = 1, 65, 2, 64 + dtype = torch.bfloat16 + q = torch.rand(B, T, H, D, dtype=dtype, device=device) + k = torch.rand(B, T, H, D, dtype=dtype, device=device) + v = torch.rand(B, T, H, D, dtype=dtype, device=device) + g = torch.randn(B, T, H, D, dtype=dtype, device=device) + beta = torch.randn(B, T, H, dtype=dtype, device=device) + A_log, dt_bias = _tle_kda_make_gate_params(H, D) + h0 = torch.randn(B, H, D, D, dtype=torch.float32, device=device) + kwargs = { + "q": q, + "k": k, + "v": v, + "g": g, + "beta": beta, + "A_log": A_log, + "dt_bias": dt_bias, + "initial_state": h0, + "output_final_state": True, + **_FLASH_KDA_REQUIRED_KWARGS, + } + + monkeypatch.setenv("FLA_TLE_KDA", "1") + monkeypatch.setenv("FLA_FLASH_KDA", "0") + monkeypatch.setenv("FLA_TILELANG", "0") + dispatched = [] + impl = KDATLEBackend.chunk_kda + + def spy(self, *args, **kw): + dispatched.append(True) + return impl(self, *args, **kw) + + monkeypatch.setattr(KDATLEBackend, "chunk_kda", spy) + with torch.inference_mode(): + out, ht = chunk_kda(**kwargs) + monkeypatch.setenv("FLA_TLE_KDA", "0") + ref_out, ref_ht = chunk_kda(**kwargs) + + assert not dispatched, "TLE KDA must reject unsupported head dimensions" + assert_close("o", ref_out, out, _FLASH_KDA_RTOL) + assert_close("ht", ref_ht, ht, _FLASH_KDA_RTOL)