Skip to content
Merged
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
24 changes: 21 additions & 3 deletions aiter/ops/triton/_triton_kernels/gated_delta_rule/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,15 +18,33 @@
"""

from .decode.fused_recurrent import _fused_recurrent_gated_delta_rule_fwd_kernel
from .prefill.chunk import chunk_gated_delta_rule_fwd
from .prefill.chunk_delta_h import chunk_gated_delta_rule_fwd_h
from .prefill.chunk_o import chunk_fwd_o
from .prefill.chunk import (
chunk_gated_delta_rule_fwd,
chunk_gated_delta_rule_fwd_opt,
chunk_gated_delta_rule_fwd_opt_vk,
)
from .prefill.chunk_delta_h import (
chunk_gated_delta_rule_fwd_h,
chunk_gated_delta_rule_fwd_h_opt,
chunk_gated_delta_rule_fwd_h_opt_vk,
)
from .prefill.chunk_o import chunk_fwd_o, chunk_fwd_o_opt, chunk_fwd_o_opt_vk
from .prefill.fused_cumsum_kkt import fused_chunk_local_cumsum_scaled_dot_kkt_fwd
from .prefill.fused_solve_tril_recompute import fused_solve_tril_recompute_w_u
from . import gated_delta_rule_utils

__all__ = [
"_fused_recurrent_gated_delta_rule_fwd_kernel",
"chunk_gated_delta_rule_fwd",
"chunk_gated_delta_rule_fwd_opt",
"chunk_gated_delta_rule_fwd_opt_vk",
"chunk_gated_delta_rule_fwd_h",
"chunk_gated_delta_rule_fwd_h_opt",
"chunk_gated_delta_rule_fwd_h_opt_vk",
"chunk_fwd_o",
"chunk_fwd_o_opt",
"chunk_fwd_o_opt_vk",
"fused_chunk_local_cumsum_scaled_dot_kkt_fwd",
"fused_solve_tril_recompute_w_u",
"gated_delta_rule_utils",
]
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,33 @@
{"cache_results": FLA_CACHE_RESULTS} if SUPPORTS_AUTOTUNE_CACHE else {}
)

FLA_USE_AUTOTUNE = False


def maybe_autotune(configs, default_config=None, **kwargs):
"""
Conditional autotune decorator.

When FLA_USE_AUTOTUNE is True, behaves identically to @triton.autotune.
When FLA_USE_AUTOTUNE is False (default), uses only the single default_config
(first config in the list if not specified), skipping all benchmark overhead.

Usage::

@maybe_autotune(
configs=[triton.Config(...), triton.Config(...), ...],
default_config=triton.Config({"BV": 64}, num_warps=4, num_stages=2),
key=[...],
)
@triton.jit
def my_kernel(...):
...
"""
if FLA_USE_AUTOTUNE:
return triton.autotune(configs=configs, **kwargs)
cfg = default_config if default_config is not None else configs[0]
return triton.autotune(configs=[cfg], **kwargs)


@lru_cache(maxsize=1)
def check_environments():
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,16 +8,36 @@
This module provides optimized Triton kernels for prefill/training operations.
"""

from .chunk import chunk_gated_delta_rule_fwd
from .chunk_delta_h import chunk_gated_delta_rule_fwd_h
from .chunk_o import chunk_fwd_o
from .fused_cumsum_kkt import fused_cumsum_kkt
from .chunk import (
chunk_gated_delta_rule_fwd,
chunk_gated_delta_rule_fwd_opt,
chunk_gated_delta_rule_fwd_opt_vk,
)
from .chunk_delta_h import (
chunk_gated_delta_rule_fwd_h,
chunk_gated_delta_rule_fwd_h_opt,
chunk_gated_delta_rule_fwd_h_opt_vk,
)
from .chunk_o import chunk_fwd_o, chunk_fwd_o_opt, chunk_fwd_o_opt_vk
from .fused_cumsum_kkt import (
fused_cumsum_kkt,
fused_chunk_local_cumsum_scaled_dot_kkt_fwd,
)
from .fused_solve_tril_recompute import fused_solve_tril_recompute_w_u
from .fused_gdn_gating_prefill import fused_gdn_gating_and_sigmoid

__all__ = [
"chunk_gated_delta_rule_fwd",
"chunk_gated_delta_rule_fwd_opt",
"chunk_gated_delta_rule_fwd_opt_vk",
"chunk_gated_delta_rule_fwd_h",
"chunk_gated_delta_rule_fwd_h_opt",
"chunk_gated_delta_rule_fwd_h_opt_vk",
"chunk_fwd_o",
"chunk_fwd_o_opt",
"chunk_fwd_o_opt_vk",
"fused_cumsum_kkt",
"fused_chunk_local_cumsum_scaled_dot_kkt_fwd",
"fused_solve_tril_recompute_w_u",
"fused_gdn_gating_and_sigmoid",
]
164 changes: 162 additions & 2 deletions aiter/ops/triton/_triton_kernels/gated_delta_rule/prefill/chunk.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,14 @@

import torch

from .chunk_delta_h import chunk_gated_delta_rule_fwd_h
from .chunk_o import chunk_fwd_o
from .chunk_delta_h import (
chunk_gated_delta_rule_fwd_h,
chunk_gated_delta_rule_fwd_h_opt,
chunk_gated_delta_rule_fwd_h_opt_vk,
)
from .chunk_o import chunk_fwd_o, chunk_fwd_o_opt, chunk_fwd_o_opt_vk
from .fused_cumsum_kkt import fused_chunk_local_cumsum_scaled_dot_kkt_fwd
from .fused_solve_tril_recompute import fused_solve_tril_recompute_w_u
from ..utils import (
chunk_local_cumsum,
chunk_scaled_dot_kkt_fwd,
Expand Down Expand Up @@ -106,3 +112,157 @@ def chunk_gated_delta_rule_fwd(
)

return g, o, A, final_state


def chunk_gated_delta_rule_fwd_opt(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
Comment thread
yiijin marked this conversation as resolved.
cu_seqlens: torch.LongTensor | None = None,
):
"""
Optimized chunk gated delta rule forward computation (Forward only).

This function implements an optimized chunk-based parallel computation for
the gated delta rule, using fused kernels and transposed intermediate layouts
to reduce global memory round-trips.

Note: This implementation only supports forward pass. Backward pass is not available.

Args:
q: Query tensor of shape [B, T, Hg, K]
k: Key tensor of shape [B, T, Hg, K]
v: Value tensor of shape [B, T, H, V]
g: Gate tensor (in log space, pre-cumsum) of shape [B, T, H]
beta: Beta parameter tensor of shape [B, T, H]
scale: Scaling factor for queries
initial_state: Initial hidden state of shape [N, H, K, V]
output_final_state: Whether to output the final state
cu_seqlens: Cumulative sequence lengths for variable-length inputs (optional) [N+1]

Returns:
tuple: (g_cumsum, o, final_state) where:
- g_cumsum: Cumulative gate values [B, T, H]
- o: Output tensor [B, T, H, V]
- final_state: Final hidden state [N, H, K, V] if output_final_state=True, else None
"""
# Step 1: Compute fused local cumulative sum of gates and KKT
g_cumsum, A_raw = fused_chunk_local_cumsum_scaled_dot_kkt_fwd(
k=k,
beta=beta,
g=g,
cu_seqlens=cu_seqlens,
)

# Step 2: Compute fused triangular solve and recompute w, u
# w, u are already in [B, H, T, K/V] head-major contiguous layout
w, u = fused_solve_tril_recompute_w_u(
A_raw=A_raw,
k=k,
v=v,
beta=beta,
g_cumsum=g_cumsum,
cu_seqlens=cu_seqlens,
)

# Step 3: Compute hidden states
h, v_new, final_state = chunk_gated_delta_rule_fwd_h_opt(
k=k,
w=w,
u=u,
g=g_cumsum,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
)

# Step 4: Compute output (directly in [B, T, H, V] layout)
o = chunk_fwd_o_opt(
q=q,
k=k,
v=v_new,
h=h,
g=g_cumsum,
scale=scale,
cu_seqlens=cu_seqlens,
)

return g_cumsum, o, final_state


def chunk_gated_delta_rule_fwd_opt_vk(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float,
initial_state: torch.Tensor,
output_final_state: bool,
cu_seqlens: torch.LongTensor | None = None,
Comment thread
yiijin marked this conversation as resolved.
):
"""
Optimized chunk gated delta rule forward with h layout [V, K].

Uses the same fused K12/K34 kernels as opt, but K5/K6 use transposed
h layout [V, K] instead of [K, V].

Args:
q: [B, T, Hg, K]
k: [B, T, Hg, K]
v: [B, T, H, V]
g: [B, T, H] — raw gate (pre-cumsum)
beta: [B, T, H]
scale: float
initial_state: [N, H, V, K] — note transposed h layout
output_final_state: bool
cu_seqlens: [N+1] optional

Returns:
tuple: (g_cumsum, o, final_state) where:
- g_cumsum: [B, T, H]
- o: [B, T, H, V]
- final_state: [N, H, V, K] if output_final_state=True, else None
"""
g_cumsum, A_raw = fused_chunk_local_cumsum_scaled_dot_kkt_fwd(
k=k,
beta=beta,
g=g,
cu_seqlens=cu_seqlens,
)

w, u = fused_solve_tril_recompute_w_u(
A_raw=A_raw,
k=k,
v=v,
beta=beta,
g_cumsum=g_cumsum,
cu_seqlens=cu_seqlens,
)

h, v_new, final_state = chunk_gated_delta_rule_fwd_h_opt_vk(
k=k,
w=w,
u=u,
g=g_cumsum,
initial_state=initial_state,
output_final_state=output_final_state,
cu_seqlens=cu_seqlens,
)

o = chunk_fwd_o_opt_vk(
q=q,
k=k,
v=v_new,
h=h,
g=g_cumsum,
scale=scale,
cu_seqlens=cu_seqlens,
)

return g_cumsum, o, final_state
Loading
Loading