"""
Tilelang implementation of the Mamba3 backward kernels with MIMO support
and variable-length sequence (varlen) support.

This is the varlen counterpart of ``mamba3_mimo_bwd.py``.  The two files
share the same mathematical kernels; the differences mirror those in the
forward file (see ``mamba3_mimo_fwd_varlen.py``):

* **NS parameter** — number of packed sequences; activates the extra
  ``i_ns`` grid dimension.
* **CU_SEQLENS tensor** — int32 prefix-sum array of shape ``[NS+1]``.
* **max_nchunks** — ``(S // chunk_size) + NS``.
* **DMIMO_O / DMIMO_Z / DMIMO_V / DD shapes** — ``[B, H, NS, R, P]`` /
  ``[B, H, NS]``, vs. ``[B, H, R, P]`` / ``[B, H]`` in the non-varlen
  version.  The NS dimension is summed-reduced in the wrapper before
  returning to the caller.

The backward pass is split into two TileLang kernels (same as non-varlen):
    1. ``mamba_mimo_bwd_fwd_varlen`` — bwd-fwd pass: re-runs the
       forward recurrence, saves STATES and QK_DOT for the bwd-bwd pass,
       and computes DMIMO_O / DMIMO_Z.
    2. ``mamba_mimo_bwd_bwd_varlen`` — bwd-bwd pass: uses the cached
       STATES and QK_DOT to compute all remaining gradients (DQ, DK, DV,
       DMIMO_V, DD, DANGLES, DDA, DFACTOR, DGAMMA_DIAG, …).

Public API:
    mamba_mimo_bwd_combined_varlen(..., cu_seqlens=None) — combined backward;
        falls back to the non-varlen ``mamba_mimo_bwd_combined`` when
        cu_seqlens is None.

Copyright (c) 2026, Dao AI Lab, Goombalab
"""

import torch
import tilelang
import tilelang.language as T
from triton.testing import do_bench

import argparse
from typing import Optional, Tuple

from mamba_ssm.ops.triton.mamba3.mamba3_mimo_utils import bwd_dadt_fused_triton_varlen, bwd_dtrap_ddt_triton_varlen
from mamba_ssm.ops.triton.mamba3.grouped_head_reduction import (
    reduce_grouped_qk_grads_and_bias_triton,
)
from mamba_ssm.ops.tilelang.mamba3.mamba3_mimo_bwd import mamba_mimo_bwd_combined

# def get_configs():
#     iter_params = dict(num_stages=[0, 1, 2, 3], threads=[128, 256, 512])
#     # iter_params = dict(num_stages=[2], threads=[128])
#     return [dict(zip(iter_params, values)) for values in itertools.product(*iter_params.values())]

# @autotune(
#     configs=get_configs(),
#     warmup=3,
#     rep=20,
# )
@tilelang.jit(
    out_idx=[],
    pass_configs={
        tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
        tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
        tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True,
    })
def mamba_mimo_bwd_fwd(
    B,
    H,
    G,
    N,
    P,
    R,
    hasZ,
    hasD,
    reduceO,
    fuse_pregate_headwise_rms_norm=False,
    isVarlen: bool = True,
    chunk_size: int = 16,
    rotary_dim_divisor: int = 4,
    dtype: str = 'float16',
    outproj_norm_eps: float = 1e-5,
    threads: int = 128,
    num_stages: int = 0,
    states_dtype = torch.bfloat16,
) -> torch.Tensor:
    """
    TileLang kernel factory for the varlen Mamba3 bwd-fwd pass.

    The bwd-fwd pass re-runs the forward recurrence in *forward* chunk order
    to (a) cache the per-chunk recurrent STATES for use by the bwd-bwd pass,
    (b) cache QK_DOT (per-step Q·K diagonal blocks), and (c) accumulate
    DMIMO_O and DMIMO_Z.

    Varlen additions vs. ``mamba_mimo_bwd_fwd`` in
    ``mamba3_mimo_bwd.py``:

    * Grid: ``T.Kernel(H, NS, B)`` — extra ``i_ns`` dimension for packed
      sequences.
    * Per-sequence bounds are derived from ``CU_SEQLENS[i_ns]``.
    * DMIMO_O / DMIMO_Z shape: ``[B, H, NS, R, P]`` (extra NS dim).
    * STATES shape: ``[B, H, max_nchunks, N, P]`` with
      ``max_nchunks = (S // chunk_size) + NS``.
    """
    S = T.dynamic("S")
    NS = T.dynamic("NS")

    accum_dtype = 'float32'
    max_nchunks = (S // chunk_size) + NS
    fused_chunk_size = chunk_size * R

    if reduceO:
        DOUT_shape = (B, S, H, P)
    else:
        DOUT_shape = (B, S, R, H, P)
    DOUT_PRE_RMS_shape = (B, H, S * R, P) if fuse_pregate_headwise_rms_norm else (1,)

    dtype_states = str(states_dtype).replace("torch.", "")

    @T.prim_func
    def mamba_mimo_bwd_fwd_kernel(
            DOUT: T.Tensor(DOUT_shape, dtype),  # type: ignore
            Q: T.Tensor([B, S, R, G, N], dtype),  # type: ignore
            K: T.Tensor([B, S, R, G, N], dtype),  # type: ignore
            V: T.Tensor([B, S, H, P], dtype),  # type: ignore
            Q_BIAS: T.Tensor([H, R, N], T.float32),  # type: ignore
            K_BIAS: T.Tensor([H, R, N], T.float32),  # type: ignore
            MIMO_V: T.Tensor([H, R, P], T.float32),  # type: ignore
            MIMO_O: T.Tensor([H, R, P], T.float32),  # type: ignore
            OUT_NORM_WEIGHT: T.Tensor([H, P], T.float32),  # type: ignore
            DMIMO_O: T.Tensor([B, H, NS, R, P], T.float32),  # type: ignore
            DOUT_NORM_WEIGHT: T.Tensor([B, H, NS, R, P], T.float32),  # type: ignore
            DOUT_PRE_RMS: T.Tensor(DOUT_PRE_RMS_shape, dtype),  # type: ignore
            STATES: T.Tensor([B, H, max_nchunks, N, P], dtype_states),  # type: ignore
            Z: T.Tensor([B, S, H, P], dtype),  # type: ignore
            MIMO_Z: T.Tensor([H, R, P], T.float32),  # type: ignore
            DZ: T.Tensor([B, S, H, P], dtype),  # type: ignore
            DMIMO_Z: T.Tensor([B, H, NS, R, P], T.float32),  # type: ignore
            ANGLES: T.Tensor([B, S, H, N // rotary_dim_divisor], T.float32),  # type: ignore
            DA_CS: T.Tensor([B, H, S], T.float32),  # type: ignore
            DA_CS_REV: T.Tensor([B, H, S], T.float32),  # type: ignore
            DT: T.Tensor([B, H, S], T.float32),  # type: ignore
            TRAP: T.Tensor([B, H, S], dtype),  # type: ignore
            D: T.Tensor([H], T.float32),  # type: ignore
            QK_DOT: T.Tensor([B, H, S, R, R], dtype),  # type: ignore
            SEGSUM: T.Tensor([B, H, max_nchunks, chunk_size, chunk_size], T.float32),  # type: ignore
            NS_ANCHOR: T.Tensor([NS], dtype=T.int32),  # type: ignore
            CU_SEQLENS: T.Tensor([NS + 1], dtype=T.int32),  # type: ignore
            ):
        """
        Varlen bwd-fwd kernel: re-runs the forward recurrence per sequence.

        Inputs:
            - DOUT: upstream gradient.
            - Q, K, V: packed activations.
            - Q_BIAS, K_BIAS, MIMO_V, MIMO_O, MIMO_Z: projection params.
            - Z, D: optional gating / skip-connection tensors.
            - ANGLES: rotary angles.
            - DA_CS, DA_CS_REV, DT, TRAP, SEGSUM: discretization tensors.
            - NS_ANCHOR: unused int32 view with shape ``[NS]``. It lets
              TileLang infer ``NS`` when optional ``[B,H,NS,R,P]`` outputs are None.
            - CU_SEQLENS: int32 prefix-sum of per-sequence lengths,
              shape ``[NS+1]``.

        Outputs (written in place):
            - DMIMO_O: gradient of the Phi projection, shape
              ``[B, H, NS, R, P]`` (sum-reduced over NS by the wrapper).
            - DMIMO_Z: gradient of the Zeta projection, same shape.
            - DOUT_NORM_WEIGHT: per-sequence partial norm-weight gradient.
            - DOUT_PRE_RMS: packed ``dL/d(raw_y)`` for bwd-bwd.
            - STATES: per-chunk recurrent states cached for the bwd-bwd
              pass, shape ``[B, H, max_nchunks, N, P]``.
            - QK_DOT: per-step Q·K diagonal blocks, shape
              ``[B, H, S, R, R]``.
            - DZ: gradient of Z (if hasZ).
        """

        with T.Kernel(H, NS, B, threads=threads) as (i_h, i_ns, i_b):
            i_h_qk = i_h // (H // G)

            q_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            k_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            PsiV_shared = T.alloc_shared([fused_chunk_size, P], dtype)
            qs_shared = T.alloc_shared([fused_chunk_size, P], dtype)
            o_shared = T.alloc_shared([chunk_size, P], dtype)
            v_shared = T.alloc_shared([chunk_size, P], dtype)
            states_accum_cast_shared = T.alloc_shared([N, P], dtype)
            qk_dot_full_shared = T.alloc_shared([fused_chunk_size, fused_chunk_size], dtype)

            if reduceO:
                dPhi_shared = T.alloc_shared([R, P], accum_dtype)
                T.clear(dPhi_shared)
            if fuse_pregate_headwise_rms_norm:
                dOutNorm_shared = T.alloc_shared([R, P], accum_dtype)
                T.clear(dOutNorm_shared)

            dout_shared = T.alloc_shared([chunk_size, P], dtype)
            z_shared = T.alloc_shared([chunk_size, P], dtype)
            dZeta_shared = T.alloc_shared([R, P], accum_dtype)
            T.clear(dZeta_shared)

            T.annotate_layout({
                q_shared: tilelang.layout.make_swizzled_layout(q_shared),
                k_shared: tilelang.layout.make_swizzled_layout(k_shared),
                PsiV_shared: tilelang.layout.make_swizzled_layout(PsiV_shared),
                qs_shared: tilelang.layout.make_swizzled_layout(qs_shared),
                o_shared: tilelang.layout.make_swizzled_layout(o_shared),
                states_accum_cast_shared: tilelang.layout.make_swizzled_layout(states_accum_cast_shared),
                qk_dot_full_shared: tilelang.layout.make_swizzled_layout(qk_dot_full_shared),
                dout_shared: tilelang.layout.make_swizzled_layout(dout_shared),
                z_shared: tilelang.layout.make_swizzled_layout(z_shared),
            })
            T.use_swizzle(10, "row")
            T.no_set_max_nreg()

            states_frag = T.alloc_fragment([N, P], accum_dtype)
            T.clear(states_frag)

            if reduceO:
                phi_frag_intrachunk = T.alloc_fragment([R, P], dtype=dtype)
                T.copy(MIMO_O[i_h, :, :], phi_frag_intrachunk)
            Psi_frag = T.alloc_fragment([R, P], dtype)
            T.copy(MIMO_V[i_h, :, :], Psi_frag)

            q_bias_frag = T.alloc_fragment([R, N], dtype)
            k_bias_frag = T.alloc_fragment([R, N], dtype)
            T.copy(Q_BIAS[i_h, :, :], q_bias_frag)
            T.copy(K_BIAS[i_h, :, :], k_bias_frag)

            # --- Per-sequence bounds ---
            start_seq_ind = T.alloc_var(T.int32)
            start_chunk_ind = T.alloc_var(T.int32)
            seq_len = T.alloc_var(T.int32)
            seq_end = T.alloc_var(T.int32)
            full_nchunks = T.alloc_var(T.int32)
            tail_len = T.alloc_var(T.int32)
            if NS > 1:
                start_seq_ind = CU_SEQLENS[i_ns]
                start_chunk_ind = (start_seq_ind // chunk_size) + i_ns
                seq_len = CU_SEQLENS[i_ns + 1] - CU_SEQLENS[i_ns]
                seq_end = start_seq_ind + seq_len
                full_nchunks = seq_len // chunk_size
                tail_len = seq_len % chunk_size
            else:
                start_seq_ind = 0
                start_chunk_ind = 0
                seq_len = S
                seq_end = S
                full_nchunks = S // chunk_size
                tail_len = S % chunk_size
            if tail_len > 0:
                full_nchunks += 1

            for i in T.Pipelined(0, full_nchunks, num_stages=num_stages):
                chunk_start = start_seq_ind + i * chunk_size
                fused_chunk_start = chunk_start * R
                global_chunk_idx = start_chunk_ind + i
                eff_tail = T.alloc_var(T.int32)
                eff_tail = chunk_size
                if i == full_nchunks - 1:
                    if tail_len > 0:
                        eff_tail = tail_len
                da_cs_end_idx = T.alloc_var(T.int32)
                da_cs_end_idx = chunk_start + chunk_size - 1
                if eff_tail < chunk_size:
                    da_cs_end_idx = seq_end - 1

                # --- Discretization Factors ---
                trap_shifted_frag = T.alloc_fragment([chunk_size], T.float32)
                dt_shifted_frag = T.alloc_fragment([chunk_size], dtype)
                shifted_gamma_frag = T.alloc_fragment([chunk_size], dtype)
                if i == full_nchunks - 1:
                    # Last chunk: shifted positions may exceed seq_end.
                    for cs in T.Parallel(chunk_size):
                        trap_shifted_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            TRAP[i_b, i_h, chunk_start + cs + 1], 0.0)
                        dt_shifted_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            DT[i_b, i_h, chunk_start + cs + 1], 0.0)
                    for cs in T.Parallel(chunk_size):
                        shifted_gamma_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            dt_shifted_frag[cs] * (T.sigmoid(-trap_shifted_frag[cs])), 0.0)
                else:
                    # Non-last chunk: chunk_start+1 .. chunk_start+chunk_size are all in bounds.
                    T.copy(TRAP[i_b, i_h, chunk_start + 1:chunk_start + 1 + chunk_size], trap_shifted_frag)
                    T.copy(DT[i_b, i_h, chunk_start + 1:chunk_start + 1 + chunk_size], dt_shifted_frag)
                    for cs in T.Parallel(chunk_size):
                        shifted_gamma_frag[cs] = dt_shifted_frag[cs] * (T.sigmoid(-trap_shifted_frag[cs]))
                shifted_gamma_shared = T.alloc_shared([chunk_size], dtype)
                T.copy(shifted_gamma_frag, shifted_gamma_shared)

                trap_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(TRAP[i_b, i_h, chunk_start:chunk_start + chunk_size], trap_frag)
                dt_frag = T.alloc_fragment([chunk_size], dtype)
                T.copy(DT[i_b, i_h, chunk_start:chunk_start + chunk_size], dt_frag)
                gamma_frag = T.alloc_fragment([chunk_size], T.float32)
                for cs in T.Parallel(chunk_size):
                    gamma_frag[cs] = dt_frag[cs] * T.sigmoid(trap_frag[cs])
                trap_scale_frag = T.alloc_fragment([chunk_size], dtype)
                for cs in T.Parallel(chunk_size):
                    trap_scale_frag[cs] = gamma_frag[cs] + shifted_gamma_shared[cs]
                trap_scale_shared = T.alloc_shared([chunk_size], dtype)
                T.copy(trap_scale_frag, trap_scale_shared)

                # --- Up-Project V and Prepare Biased Q/K ---
                PsiV_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                T.copy(V[i_b, chunk_start:chunk_start + chunk_size, i_h, :], v_shared)
                for cs, r, p in T.Parallel(chunk_size, R, P):
                    PsiV_frag[cs, r, p] = v_shared[cs, p] * Psi_frag[r, p]
                PsiV_reshaped_frag = T.view(PsiV_frag, shape=[fused_chunk_size, P])
                T.copy(PsiV_reshaped_frag, PsiV_shared)

                q_reshaped_shared = T.view(q_shared, shape=[chunk_size, R, N])
                T.copy(Q[i_b, chunk_start:chunk_start + chunk_size, :, i_h_qk, :], q_reshaped_shared)
                q_frag = T.alloc_fragment([chunk_size, R, N], dtype)
                T.copy(q_reshaped_shared, q_frag)
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    q_frag[cs, r, n] += q_bias_frag[r, n]
                T.copy(q_frag, q_reshaped_shared)

                k_reshaped_shared = T.view(k_shared, shape=[chunk_size, R, N])
                T.copy(K[i_b, chunk_start:chunk_start + chunk_size, :, i_h_qk, :], k_reshaped_shared)
                k_frag = T.alloc_fragment([chunk_size, R, N], dtype)
                T.copy(k_reshaped_shared, k_frag)
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    k_frag[cs, r, n] += k_bias_frag[r, n]
                T.copy(k_frag, k_reshaped_shared)

                # --- QK dot ---
                qk_dot_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], dtype=accum_dtype)
                T.gemm(q_shared, k_shared, qk_dot_frag, transpose_B=True, clear_accum=True)
                T.copy(qk_dot_frag, qk_dot_full_shared)
                # Write QK_DOT; for full chunks eff_tail == chunk_size so cs < eff_tail is
                # always true and the predication has no effect on non-tail iterations.
                for cs, r_out, r_in in T.Parallel(chunk_size, R, R):
                    if cs < eff_tail:
                        QK_DOT[i_b, i_h, chunk_start + cs, r_out, r_in] = \
                            qk_dot_full_shared[cs * R + r_out, cs * R + r_in]

                # --- Rotary Q/K ---
                q_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                q_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    q_first_half_frag[cs, r, n] = q_shared[cs * R + r, n]
                    q_second_half_frag[cs, r, n] = q_shared[cs * R + r, N // 2 + n]
                angles_frag = T.alloc_fragment([chunk_size, N // rotary_dim_divisor], T.float32)
                T.copy(ANGLES[i_b, chunk_start:chunk_start + chunk_size, i_h, :], angles_frag)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    q_shared[cs * R + r, n] = T.cos(angles_frag[cs, n]) * q_first_half_frag[cs, r, n] - T.sin(angles_frag[cs, n]) * q_second_half_frag[cs, r, n]
                    q_shared[cs * R + r, N // 2 + n] = T.sin(angles_frag[cs, n]) * q_first_half_frag[cs, r, n] + T.cos(angles_frag[cs, n]) * q_second_half_frag[cs, r, n]

                k_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                k_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    k_first_half_frag[cs, r, n] = k_shared[cs * R + r, n]
                    k_second_half_frag[cs, r, n] = k_shared[cs * R + r, N // 2 + n]
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    k_shared[cs * R + r, n] = T.cos(angles_frag[cs, n]) * k_first_half_frag[cs, r, n] - T.sin(angles_frag[cs, n]) * k_second_half_frag[cs, r, n]
                    k_shared[cs * R + r, N // 2 + n] = T.sin(angles_frag[cs, n]) * k_first_half_frag[cs, r, n] + T.cos(angles_frag[cs, n]) * k_second_half_frag[cs, r, n]

                k_trap_scaled_frag = T.alloc_fragment([fused_chunk_size, N], dtype)
                T.copy(k_shared, k_trap_scaled_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    k_trap_scaled_frag[csr, n] *= trap_scale_shared[csr // R]
                T.copy(k_trap_scaled_frag, k_shared)

                # --- Interchunk + Intrachunk Output ---
                q_state_out_frag = T.alloc_fragment([fused_chunk_size, P], dtype=accum_dtype)
                T.copy(states_frag, states_accum_cast_shared)
                T.gemm(q_shared, states_accum_cast_shared, q_state_out_frag, clear_accum=True)

                qk_intrachunk_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], dtype=accum_dtype)
                T.gemm(q_shared, k_shared, qk_intrachunk_frag, transpose_B=True, clear_accum=True)

                da_cs__or__exp_da_cs_shared = T.alloc_shared([chunk_size], T.float32)
                T.copy(DA_CS[i_b, i_h, chunk_start:chunk_start + chunk_size], da_cs__or__exp_da_cs_shared)
                for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                    qk_intrachunk_frag[csr_i, csr_j] = T.if_then_else(
                        csr_i // R > csr_j // R,
                        qk_intrachunk_frag[csr_i, csr_j] * T.exp(SEGSUM[i_b, i_h, global_chunk_idx, csr_i // R, csr_j // R]),
                        0.0)
                qk_intrachunk_masked_shared = T.alloc_shared([fused_chunk_size, fused_chunk_size], dtype=dtype)
                for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                    qk_intrachunk_masked_shared[csr_i, csr_j] = qk_intrachunk_frag[csr_i, csr_j]

                for cs in T.Parallel(chunk_size):
                    da_cs__or__exp_da_cs_shared[cs] = T.exp(da_cs__or__exp_da_cs_shared[cs])
                exp_da_cs_frag = T.alloc_fragment([chunk_size], dtype=T.float32)
                T.copy(da_cs__or__exp_da_cs_shared, exp_da_cs_frag)
                for csr, p in T.Parallel(fused_chunk_size, P):
                    q_state_out_frag[csr, p] *= exp_da_cs_frag[csr // R]

                o_mimo_accum_frag = T.alloc_fragment([fused_chunk_size, P], dtype=accum_dtype)
                T.gemm(qk_intrachunk_masked_shared, PsiV_shared, o_mimo_accum_frag, clear_accum=True)
                for cs, p in T.Parallel(fused_chunk_size, P):
                    o_mimo_accum_frag[cs, p] += q_state_out_frag[cs, p]

                # --- Diagonal Terms ---
                qkdot_psiv_frag = T.alloc_fragment([chunk_size, R, P], dtype=dtype)
                T.clear(qkdot_psiv_frag)
                for cs, r_out, p in T.Parallel(chunk_size, R, P):
                    for r_in in T.serial(R):
                        qkdot_psiv_frag[cs, r_out, p] += qk_dot_full_shared[cs * R + r_out, cs * R + r_in] * PsiV_shared[cs * R + r_in, p]
                    qkdot_psiv_frag[cs, r_out, p] *= gamma_frag[cs]
                qkdot_psiv_reshaped_frag = T.view(qkdot_psiv_frag, shape=[fused_chunk_size, P])
                for csr, p in T.Parallel(fused_chunk_size, P):
                    o_mimo_accum_frag[csr, p] += qkdot_psiv_reshaped_frag[csr, p]

                if hasD:
                    D_var = T.alloc_var(T.float32)
                    T.copy(D[i_h], D_var)
                    PsiV_D_frag = T.alloc_fragment([fused_chunk_size, P], T.float32)
                    T.copy(PsiV_shared, PsiV_D_frag)
                    for csr, p in T.Parallel(fused_chunk_size, P):
                        o_mimo_accum_frag[csr, p] += D_var * PsiV_D_frag[csr, p]

                # --- Projection, optional gate, and pregate RMS side gradients ---
                if reduceO:
                    if not fuse_pregate_headwise_rms_norm:
                        out_prereduced_shared = T.alloc_shared([fused_chunk_size, P], dtype)
                        T.copy(o_mimo_accum_frag, out_prereduced_shared)

                    o_gated_frag = T.alloc_fragment([chunk_size, R, P], T.float32)
                    if fuse_pregate_headwise_rms_norm:
                        # raw_y is the per-rank MIMO output before RMSNorm, gate, and down projection.
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            o_gated_frag[cs, r, p] = o_mimo_accum_frag[cs * R + r, p]
                            o_gated_frag[cs, r, p] *= o_gated_frag[cs, r, p]
                        o_rstd_frag = T.alloc_fragment([chunk_size, R], T.float32)
                        T.reduce_sum(o_gated_frag, o_rstd_frag, dim=-1, clear=True)
                        for cs, r in T.Parallel(chunk_size, R):
                            o_rstd_frag[cs, r] = 1.0 / T.sqrt(
                                o_rstd_frag[cs, r] / P + outproj_norm_eps
                            )

                        T.copy(Z[i_b, chunk_start:chunk_start + chunk_size, i_h, :], z_shared)
                        z_o_frag = T.alloc_fragment([chunk_size, P], T.float32)
                        T.copy(z_shared, z_o_frag)
                        Zeta_o_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(MIMO_Z[i_h, :, :], Zeta_o_frag)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            tmp = z_o_frag[cs, p] * Zeta_o_frag[r, p] * 0.5
                            o_gated_frag[cs, r, p] = tmp * T.tanh(tmp) + tmp
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            o_gated_frag[cs, r, p] *= (
                                o_mimo_accum_frag[cs * R + r, p]
                                * o_rstd_frag[cs, r]
                                * OUT_NORM_WEIGHT[i_h, p]
                            )
                    elif hasZ:
                        T.copy(Z[i_b, chunk_start:chunk_start + chunk_size, i_h, :], z_shared)
                        z_o_frag = T.alloc_fragment([chunk_size, P], T.float32)
                        T.copy(z_shared, z_o_frag)
                        Zeta_o_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(MIMO_Z[i_h, :, :], Zeta_o_frag)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            tmp = z_o_frag[cs, p] * Zeta_o_frag[r, p] * 0.5
                            o_gated_frag[cs, r, p] = tmp * T.tanh(tmp) + tmp
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            o_gated_frag[cs, r, p] *= out_prereduced_shared[cs * R + r, p]
                    else:
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            o_gated_frag[cs, r, p] = out_prereduced_shared[cs * R + r, p]

                    dPhi_frag = T.alloc_fragment([R, P], T.float32)
                    T.copy(dPhi_shared, dPhi_frag)
                    dout_frag = T.alloc_fragment([chunk_size, P], dtype)
                    T.copy(DOUT[i_b, chunk_start:chunk_start + chunk_size, i_h, :], dout_shared)
                    if eff_tail < chunk_size:
                        for cs, p in T.Parallel(chunk_size, P):
                            dout_shared[cs, p] = T.if_then_else(cs < eff_tail, dout_shared[cs, p], 0.0)
                    T.copy(dout_shared, dout_frag)

                    if fuse_pregate_headwise_rms_norm:
                        dPhi_prereduce = T.alloc_fragment([chunk_size, R, P], T.float32)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhi_prereduce[cs, r, p] = o_gated_frag[cs, r, p] * dout_frag[cs, p]
                        T.reduce_sum(dPhi_prereduce, dPhi_frag, dim=0, clear=False)
                        T.copy(dPhi_frag, dPhi_shared)
                    else:
                        for r, p in T.Parallel(R, P):
                            for cs in T.serial(chunk_size):
                                dPhi_frag[r, p] += o_gated_frag[cs, r, p] * dout_frag[cs, p]
                        T.copy(dPhi_frag, dPhi_shared)

                    if fuse_pregate_headwise_rms_norm:
                        Phi_frag = T.alloc_fragment([R, P], dtype)
                        T.copy(MIMO_O[i_h, :, :], Phi_frag)

                        # dnorm is dL/d(xhat), where xhat = raw_y * rstd.
                        dPhiO_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] = (
                                dout_frag[cs, p]
                                * Phi_frag[r, p]
                                * o_mimo_accum_frag[cs * R + r, p]
                                * o_rstd_frag[cs, r]
                                * OUT_NORM_WEIGHT[i_h, p]
                            )

                        z_frag = T.alloc_fragment([chunk_size, P], T.float32)
                        T.copy(z_shared, z_frag)
                        Zeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(MIMO_Z[i_h, :, :], Zeta_frag)
                        dZetaZ_frag = T.alloc_fragment([chunk_size, R, P], T.float32)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] = z_frag[cs, p] * Zeta_frag[r, p]
                            dZetaZ_frag[cs, r, p] = dPhiO_frag[cs, r, p] * T.sigmoid(dZetaZ_frag[cs, r, p]) * \
                                (1 + dZetaZ_frag[cs, r, p] * (T.sigmoid(-dZetaZ_frag[cs, r, p])))
                        dZ_frag_prereduce = T.alloc_fragment([chunk_size, R, P], dtype)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZ_frag_prereduce[cs, r, p] = dZetaZ_frag[cs, r, p] * Zeta_frag[r, p]
                        dZ_frag = T.alloc_fragment([chunk_size, P], dtype)
                        T.reduce_sum(dZ_frag_prereduce, dZ_frag, clear=True, dim=1)
                        if eff_tail < chunk_size:
                            for cs, p in T.Parallel(chunk_size, P):
                                if cs < eff_tail:
                                    DZ[i_b, chunk_start + cs, i_h, p] = dZ_frag[cs, p]
                        else:
                            T.copy(dZ_frag, DZ[i_b, chunk_start:chunk_start + chunk_size, i_h, :])

                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] *= z_frag[cs, p]
                        dZeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(dZeta_shared, dZeta_frag)
                        T.reduce_sum(dZetaZ_frag, dZeta_frag, clear=False, dim=0)
                        T.copy(dZeta_frag, dZeta_shared)

                        weighted_dot_frag = T.alloc_fragment([chunk_size, R, P], T.float32)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            tmp = z_frag[cs, p] * Zeta_frag[r, p] * 0.5
                            gate = tmp * T.tanh(tmp) + tmp
                            xhat = o_mimo_accum_frag[cs * R + r, p] * o_rstd_frag[cs, r]
                            dnorm = dout_frag[cs, p] * Phi_frag[r, p] * gate
                            weighted_dot_frag[cs, r, p] = dnorm * xhat

                        dOutNorm_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(dOutNorm_shared, dOutNorm_frag)
                        T.reduce_sum(weighted_dot_frag, dOutNorm_frag, dim=0, clear=False)
                        T.copy(dOutNorm_frag, dOutNorm_shared)

                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            weighted_dot_frag[cs, r, p] *= OUT_NORM_WEIGHT[i_h, p]
                        rms_dot_frag = T.alloc_fragment([chunk_size, R], T.float32)
                        T.reduce_sum(weighted_dot_frag, rms_dot_frag, dim=-1, clear=True)
                        for cs, r in T.Parallel(chunk_size, R):
                            rms_dot_frag[cs, r] /= P

                        if eff_tail < chunk_size:
                            for cs, r, p in T.Parallel(chunk_size, R, P):
                                if cs < eff_tail:
                                    tmp = z_frag[cs, p] * Zeta_frag[r, p] * 0.5
                                    gate = tmp * T.tanh(tmp) + tmp
                                    xhat = o_mimo_accum_frag[cs * R + r, p] * o_rstd_frag[cs, r]
                                    dnorm = dout_frag[cs, p] * Phi_frag[r, p] * gate
                                    DOUT_PRE_RMS[i_b, i_h, fused_chunk_start + cs * R + r, p] = (
                                        o_rstd_frag[cs, r]
                                        * (
                                            dnorm * OUT_NORM_WEIGHT[i_h, p]
                                            - xhat * rms_dot_frag[cs, r]
                                        )
                                    )
                        else:
                            for cs, r, p in T.Parallel(chunk_size, R, P):
                                tmp = z_frag[cs, p] * Zeta_frag[r, p] * 0.5
                                gate = tmp * T.tanh(tmp) + tmp
                                xhat = o_mimo_accum_frag[cs * R + r, p] * o_rstd_frag[cs, r]
                                dnorm = dout_frag[cs, p] * Phi_frag[r, p] * gate
                                DOUT_PRE_RMS[i_b, i_h, fused_chunk_start + cs * R + r, p] = (
                                    o_rstd_frag[cs, r]
                                    * (
                                        dnorm * OUT_NORM_WEIGHT[i_h, p]
                                        - xhat * rms_dot_frag[cs, r]
                                    )
                                )
                    elif hasZ:
                        Phi_frag = T.alloc_fragment([R, P], dtype)
                        T.copy(MIMO_O[i_h, :, :], Phi_frag)
                        dPhiO_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                        dout_preexpand_frag = T.alloc_fragment([chunk_size, P], dtype)
                        T.copy(dout_shared, dout_preexpand_frag)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] = dout_frag[cs, p] * Phi_frag[r, p]
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] *= out_prereduced_shared[cs * R + r, p]
                        z_frag = T.alloc_fragment([chunk_size, P], T.float32)
                        T.copy(z_shared, z_frag)
                        Zeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(MIMO_Z[i_h, :, :], Zeta_frag)
                        dZetaZ_frag = T.alloc_fragment([chunk_size, R, P], T.float32)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] = z_frag[cs, p] * Zeta_frag[r, p]
                            dZetaZ_frag[cs, r, p] = dPhiO_frag[cs, r, p] * T.sigmoid(dZetaZ_frag[cs, r, p]) * \
                                (1 + dZetaZ_frag[cs, r, p] * (T.sigmoid(-dZetaZ_frag[cs, r, p])))
                        dZ_frag = T.alloc_fragment([chunk_size, P], dtype)
                        T.clear(dZ_frag)
                        for cs, p in T.Parallel(chunk_size, P):
                            for r in T.serial(R):
                                dZ_frag[cs, p] += dZetaZ_frag[cs, r, p] * Zeta_frag[r, p]
                        if eff_tail < chunk_size:
                            for cs, p in T.Parallel(chunk_size, P):
                                if cs < eff_tail:
                                    DZ[i_b, chunk_start + cs, i_h, p] = dZ_frag[cs, p]
                        else:
                            T.copy(dZ_frag, DZ[i_b, chunk_start:chunk_start + chunk_size, i_h, :])
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] *= z_frag[cs, p]
                        dZeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(dZeta_shared, dZeta_frag)
                        T.reduce_sum(dZetaZ_frag, dZeta_frag, clear=False, dim=0)
                        T.copy(dZeta_frag, dZeta_shared)
                else:
                    if hasZ:
                        out_prereduced_shared = T.alloc_shared([fused_chunk_size, P], dtype)
                        T.copy(o_mimo_accum_frag, out_prereduced_shared)
                        T.copy(Z[i_b, chunk_start:chunk_start + chunk_size, i_h, :], z_shared)
                        dPhiO_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] = DOUT[i_b, chunk_start + cs, r, i_h, p]
                        if eff_tail < chunk_size:
                            for cs, r, p in T.Parallel(chunk_size, R, P):
                                dPhiO_frag[cs, r, p] = T.if_then_else(cs < eff_tail, dPhiO_frag[cs, r, p], 0.0)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] *= out_prereduced_shared[cs * R + r, p]
                        z_frag = T.alloc_fragment([chunk_size, P], T.float32)
                        T.copy(z_shared, z_frag)
                        Zeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(MIMO_Z[i_h, :, :], Zeta_frag)
                        dZetaZ_frag = T.alloc_fragment([chunk_size, R, P], T.float32)
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] = z_frag[cs, p] * Zeta_frag[r, p]
                            dZetaZ_frag[cs, r, p] = dPhiO_frag[cs, r, p] * T.sigmoid(dZetaZ_frag[cs, r, p]) * \
                                (1 + dZetaZ_frag[cs, r, p] * (T.sigmoid(-dZetaZ_frag[cs, r, p])))
                        dZ_frag = T.alloc_fragment([chunk_size, P], dtype)
                        T.clear(dZ_frag)
                        for cs, p in T.Parallel(chunk_size, P):
                            for r in T.serial(R):
                                dZ_frag[cs, p] += dZetaZ_frag[cs, r, p] * Zeta_frag[r, p]
                        if eff_tail < chunk_size:
                            for cs, p in T.Parallel(chunk_size, P):
                                if cs < eff_tail:
                                    DZ[i_b, chunk_start + cs, i_h, p] = dZ_frag[cs, p]
                        else:
                            T.copy(dZ_frag, DZ[i_b, chunk_start:chunk_start + chunk_size, i_h, :])
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dZetaZ_frag[cs, r, p] *= z_frag[cs, p]
                        dZeta_frag = T.alloc_fragment([R, P], T.float32)
                        T.copy(dZeta_shared, dZeta_frag)
                        T.reduce_sum(dZetaZ_frag, dZeta_frag, clear=False, dim=0)
                        T.copy(dZeta_frag, dZeta_shared)

                # --- Save and Update Recurrent State ---
                T.copy(states_frag, STATES[i_b, i_h, global_chunk_idx, :, :])

                dA_cs_rev_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(DA_CS_REV[i_b, i_h, chunk_start:chunk_start + chunk_size], dA_cs_rev_frag)
                k_state_frag = T.alloc_fragment([fused_chunk_size, N], dtype)
                T.copy(k_shared, k_state_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    k_state_frag[csr, n] *= T.exp(dA_cs_rev_frag[csr // R])

                da_cs_sum = T.alloc_var(T.float32)
                T.copy(DA_CS[i_b, i_h, da_cs_end_idx], da_cs_sum)
                for n, p in T.Parallel(N, P):
                    states_frag[n, p] *= T.exp(da_cs_sum)
                T.gemm(k_state_frag, PsiV_shared, states_frag, transpose_A=True, clear_accum=False)

            if reduceO:
                T.copy(dPhi_shared, DMIMO_O[i_b, i_h, i_ns, :, :])
            if fuse_pregate_headwise_rms_norm:
                T.copy(dOutNorm_shared, DOUT_NORM_WEIGHT[i_b, i_h, i_ns, :, :])
            if hasZ:
                T.copy(dZeta_shared, DMIMO_Z[i_b, i_h, i_ns, :, :])

    return mamba_mimo_bwd_fwd_kernel

# def get_configs():
#     iter_params = dict(num_stages=[0], threads=[128, 256])
#     return [dict(zip(iter_params, values)) for values in itertools.product(*iter_params.values())]

# @autotune(
#     configs=get_configs(),
#     warmup=3,
#     rep=20,
# )
@tilelang.jit(
    out_idx=[],
    pass_configs={
        tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True,
        tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True,
        tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True,
    })
def mamba_mimo_bwd_bwd(
    B,
    H,
    G,
    N,
    P,
    R,
    hasZ,
    hasD,
    reduceO,
    packed_dout=False,
    isVarlen: bool = True,
    chunk_size: int = 16,
    rotary_dim_divisor: int = 4,
    dtype: str = 'float16',
    threads: int = 256,
    num_stages: int = 0,
    states_dtype = torch.bfloat16,
) -> torch.Tensor:
    """
    TileLang kernel factory for the varlen Mamba3 bwd-bwd pass.

    The bwd-bwd pass iterates chunks in *reverse* order, consuming the
    STATES and QK_DOT cached by the bwd-fwd pass, and writes all remaining
    gradients: DQ, DK, DV, DMIMO_V, DD, DANGLES, DDA, DFACTOR,
    DGAMMA_DIAG, DDA_CS, DDA_CS_REV, DSSDA.

    Varlen additions vs. ``mamba_mimo_bwd_bwd`` in
    ``mamba3_mimo_bwd.py``:

    * Grid: ``T.Kernel(H, NS, B)`` — extra ``i_ns`` dimension.
    * Per-sequence bounds from ``CU_SEQLENS[i_ns]``.
    * DMIMO_V / DD shapes: ``[B, H, NS, R, P]`` / ``[B, H, NS]``.
    * DSSDA shape: ``[B, H, max_nchunks, C, C]`` with
      ``max_nchunks = (S // chunk_size) + NS``.
    """
    S = T.dynamic("S")
    NS = T.dynamic("NS")

    accum_dtype = 'float32'
    max_nchunks = (S // chunk_size) + NS
    fused_chunk_size = chunk_size * R

    if packed_dout:
        DOUT_shape = (B, H, S * R, P)
    elif reduceO:
        DOUT_shape = (B, S, H, P)
    else:
        DOUT_shape = (B, S, R, H, P)

    dtype_states = str(states_dtype).replace("torch.", "")

    @T.prim_func
    def mamba_mimo_bwd_bwd_kernel(
            DOUT: T.Tensor(DOUT_shape, dtype),  # type: ignore
            Q: T.Tensor([B, S, R, G, N], dtype),  # type: ignore
            K: T.Tensor([B, S, R, G, N], dtype),  # type: ignore
            V: T.Tensor([B, S, H, P], dtype),  # type: ignore
            Q_BIAS: T.Tensor([H, R, N], T.float32),  # type: ignore
            K_BIAS: T.Tensor([H, R, N], T.float32),  # type: ignore
            MIMO_V: T.Tensor([H, R, P], T.float32),  # type: ignore
            MIMO_O: T.Tensor([H, R, P], T.float32),  # type: ignore
            DK: T.Tensor([B, S * R, H, N], dtype),  # type: ignore
            DV: T.Tensor([B, S, H, P], dtype),  # type: ignore
            DMIMO_V: T.Tensor([B, H, NS, R, P], T.float32),  # type: ignore
            STATES: T.Tensor([B, H, max_nchunks, N, P], dtype_states),  # type: ignore
            DQ: T.Tensor([B, S * R, H, N], dtype),  # type: ignore
            Z: T.Tensor([B, S, H, P], dtype),  # type: ignore
            MIMO_Z: T.Tensor([H, R, P], T.float32),  # type: ignore
            ANGLES: T.Tensor([B, S, H, N // rotary_dim_divisor], T.float32),  # type: ignore
            DA_CS: T.Tensor([B, H, S], T.float32),  # type: ignore
            DA_CS_REV: T.Tensor([B, H, S], T.float32),  # type: ignore
            DT: T.Tensor([B, H, S], T.float32),  # type: ignore
            TRAP: T.Tensor([B, H, S], dtype),  # type: ignore
            DFACTOR: T.Tensor([B, H, S], T.float32),  # type: ignore
            DGAMMA_DIAG: T.Tensor([B, H, S], T.float32),  # type: ignore
            DANGLES: T.Tensor([B, S, H, N // rotary_dim_divisor], T.float32),  # type: ignore
            D: T.Tensor([H], T.float32),  # type: ignore
            DD: T.Tensor([B, H, NS], T.float32),  # type: ignore
            QK_DOT: T.Tensor([B, H, S, R, R], dtype),  # type: ignore
            DDA: T.Tensor([B, H, S], T.float32),  # type: ignore
            DSSDA: T.Tensor([B, H, max_nchunks, chunk_size, chunk_size], T.float32),  # type: ignore
            DDA_CS_REV: T.Tensor([B, H, S], T.float32),  # type: ignore
            DDA_CS: T.Tensor([B, H, S], T.float32),  # type: ignore
            SEGSUM: T.Tensor([B, H, max_nchunks, chunk_size, chunk_size], T.float32),  # type: ignore
            CU_SEQLENS: T.Tensor([NS + 1], dtype=T.int32),  # type: ignore
            ):
        """
        Varlen bwd-bwd kernel: computes all remaining gradients in reverse
        chunk order.

        Inputs:
            - DOUT: upstream gradient. Shape is ``[B,S,H,P]`` when reduceO,
              ``[B,S,R,H,P]`` when not reduceO, or packed ``[B,H,S*R,P]``
              when packed_dout consumes DOUT_PRE_RMS from bwd-fwd.
            - Q, K, V: packed activations.
            - Q_BIAS, K_BIAS, MIMO_V, MIMO_O, MIMO_Z: projection params.
            - Z, D: optional gating / skip-connection tensors.
            - ANGLES: rotary angles.
            - DA_CS, DA_CS_REV, DT, TRAP, SEGSUM: discretization tensors.
            - STATES: per-chunk states cached by the bwd-fwd pass,
              shape ``[B, H, max_nchunks, N, P]``.
            - QK_DOT: per-step Q·K diagonal blocks cached by bwd-fwd,
              shape ``[B, H, S, R, R]``.
            - CU_SEQLENS: int32 prefix-sum of per-sequence lengths,
              shape ``[NS+1]``.

        Outputs (written in place):
            - DQ, DK, DV: activation gradients.
            - DMIMO_V: gradient of Psi, shape ``[B, H, NS, R, P]``.
            - DD: gradient of D, shape ``[B, H, NS]``.
            - DANGLES: gradient of rotary angles.
            - DDA, DDA_CS, DDA_CS_REV, DSSDA, DFACTOR, DGAMMA_DIAG:
              intermediate discretization gradients consumed by the
              Triton utility kernels in ``mamba3_utils_varlen.py``.
        """

        with T.Kernel(H, NS, B, threads=threads) as (i_h, i_ns, i_b):
            i_h_qk = i_h // (H // G)

            dstates_shared = T.alloc_shared([N, P], dtype)
            dstates_frag = T.alloc_fragment([N, P], accum_dtype)
            dout_shared = T.alloc_shared([chunk_size, P], dtype)
            dPhiO_shared = T.alloc_shared([fused_chunk_size, P], dtype)
            q_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            k_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            v_shared = T.alloc_shared([chunk_size, P], dtype)
            states_shared = T.alloc_shared([N, P], dtype)
            lkq_masked__or__dkq_masked_shared = T.alloc_shared([fused_chunk_size, fused_chunk_size], dtype)
            dPsiV_combined_shared = T.alloc_shared([fused_chunk_size, P], dtype)
            dqk_from_diag_shared = T.alloc_shared([fused_chunk_size, fused_chunk_size], accum_dtype)
            q_pre_rot_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            k_pre_rot_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            dk_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            dq_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            qk_dot_shared = T.alloc_shared([chunk_size, R, R], dtype)
            k_pre_trap_shared = T.alloc_shared([fused_chunk_size, N], dtype)
            dangle_dk__or__dq_shared = T.alloc_shared([fused_chunk_size, N // rotary_dim_divisor], T.float32)

            noswizzle_annot = threads == 256 and (N <= 32 or P >= 128) # NOTE: heuristics for when swizzling annotation causes kernel hang, needs more investigation
            if not noswizzle_annot:
                T.annotate_layout({
                    dstates_shared: tilelang.layout.make_swizzled_layout(dstates_shared),
                    dout_shared: tilelang.layout.make_swizzled_layout(dout_shared),
                    q_shared: tilelang.layout.make_swizzled_layout(q_shared),
                    k_shared: tilelang.layout.make_swizzled_layout(k_shared),
                    v_shared: tilelang.layout.make_swizzled_layout(v_shared),
                    states_shared: tilelang.layout.make_swizzled_layout(states_shared),
                    lkq_masked__or__dkq_masked_shared: tilelang.layout.make_swizzled_layout(lkq_masked__or__dkq_masked_shared),
                    dPsiV_combined_shared: tilelang.layout.make_swizzled_layout(dPsiV_combined_shared),
                    dqk_from_diag_shared: tilelang.layout.make_swizzled_layout(dqk_from_diag_shared),
                    k_pre_rot_shared: tilelang.layout.make_swizzled_layout(k_pre_rot_shared),
                    q_pre_rot_shared: tilelang.layout.make_swizzled_layout(q_pre_rot_shared),
                    dk_shared: tilelang.layout.make_swizzled_layout(dk_shared),
                    dq_shared: tilelang.layout.make_swizzled_layout(dq_shared),
                    k_pre_trap_shared: tilelang.layout.make_swizzled_layout(k_pre_trap_shared),
                    dangle_dk__or__dq_shared: tilelang.layout.make_swizzled_layout(dangle_dk__or__dq_shared),
                })
            T.use_swizzle(10, "row")
            T.no_set_max_nreg()

            T.clear(dstates_frag)
            T.clear(dstates_shared)

            if reduceO:
                Phi_frag = T.alloc_fragment([R, P], dtype)
                T.copy(MIMO_O[i_h, :, :], Phi_frag)
            Psi_frag = T.alloc_fragment([R, P], dtype)
            T.copy(MIMO_V[i_h, :, :], Psi_frag)

            dPsi_acc = T.alloc_fragment([R, P], accum_dtype)
            T.clear(dPsi_acc)

            if hasD:
                dD_frag = T.alloc_fragment([1], accum_dtype)
                T.clear(dD_frag)

            q_bias_frag = T.alloc_fragment([R, N], dtype)
            k_bias_frag = T.alloc_fragment([R, N], dtype)
            T.copy(Q_BIAS[i_h, :, :], q_bias_frag)
            T.copy(K_BIAS[i_h, :, :], k_bias_frag)

            # --- Per-sequence bounds ---
            start_seq_ind = T.alloc_var(T.int32)
            start_chunk_ind = T.alloc_var(T.int32)
            seq_len = T.alloc_var(T.int32)
            seq_end = T.alloc_var(T.int32)
            full_nchunks = T.alloc_var(T.int32)
            tail_len = T.alloc_var(T.int32)
            if NS > 1:
                start_seq_ind = CU_SEQLENS[i_ns]
                start_chunk_ind = (start_seq_ind // chunk_size) + i_ns
                seq_len = CU_SEQLENS[i_ns + 1] - CU_SEQLENS[i_ns]
                seq_end = start_seq_ind + seq_len
                full_nchunks = seq_len // chunk_size
                tail_len = seq_len % chunk_size
            else:
                start_seq_ind = 0
                start_chunk_ind = 0
                seq_len = S
                seq_end = S
                full_nchunks = S // chunk_size
                tail_len = S % chunk_size
            if tail_len > 0:
                full_nchunks += 1

            for chunk_idx_rev in T.Pipelined(0, full_nchunks, num_stages=num_stages):
                chunk_idx = full_nchunks - 1 - chunk_idx_rev
                chunk_start = start_seq_ind + chunk_idx * chunk_size
                fused_chunk_start = chunk_start * R
                global_chunk_idx = start_chunk_ind + chunk_idx
                # Effective tail cutoff: chunk_size for full chunks, tail_len for the last tail chunk.
                # Used to mask all output writes with a single T.if_then_else + T.copy instead of
                # many separate if-else blocks.  For full chunks eff_tail == chunk_size so the
                # masking condition (cs < eff_tail) is trivially true and T.copy runs unguarded.
                eff_tail = T.alloc_var(T.int32)
                eff_tail = chunk_size
                if chunk_idx == full_nchunks - 1:
                    if tail_len > 0:
                        eff_tail = tail_len
                # Index of the last valid DA_CS entry for this chunk (used twice below).
                da_cs_end_idx = T.alloc_var(T.int32)
                da_cs_end_idx = chunk_start + chunk_size - 1
                if eff_tail < chunk_size:
                    da_cs_end_idx = seq_end - 1
                # --- Discretization Factors ---
                # For non-last chunks every shifted position is within sequence bounds, so we
                # can use vectorised T.copy instead of per-element T.if_then_else. This
                # replaces 3*chunk_size predicated GMEM loads with 2 bulk T.copy calls for
                # all but the last chunk in the reverse iteration order (chunk_idx_rev == 0).
                trap_shifted_frag = T.alloc_fragment([chunk_size], T.float32)
                dt_shifted_frag = T.alloc_fragment([chunk_size], dtype)
                shifted_gamma_frag = T.alloc_fragment([chunk_size], dtype)
                if chunk_idx_rev == 0:
                    # Last chunk (first in reverse): shifted positions may exceed seq_end.
                    for cs in T.Parallel(chunk_size):
                        trap_shifted_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            TRAP[i_b, i_h, chunk_start + cs + 1], 0.0)
                        dt_shifted_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            DT[i_b, i_h, chunk_start + cs + 1], 0.0)
                    for cs in T.Parallel(chunk_size):
                        shifted_gamma_frag[cs] = T.if_then_else(
                            cs + 1 < eff_tail,
                            dt_shifted_frag[cs] * T.sigmoid(-trap_shifted_frag[cs]), 0.0)
                else:
                    # Non-last chunk: chunk_start+1 .. chunk_start+chunk_size are all in bounds.
                    T.copy(TRAP[i_b, i_h, chunk_start + 1:chunk_start + 1 + chunk_size], trap_shifted_frag)
                    T.copy(DT[i_b, i_h, chunk_start + 1:chunk_start + 1 + chunk_size], dt_shifted_frag)
                    for cs in T.Parallel(chunk_size):
                        shifted_gamma_frag[cs] = dt_shifted_frag[cs] * T.sigmoid(-trap_shifted_frag[cs])

                trap_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(TRAP[i_b, i_h, chunk_start:chunk_start + chunk_size], trap_frag)
                dt_frag = T.alloc_fragment([chunk_size], dtype)
                T.copy(DT[i_b, i_h, chunk_start:chunk_start + chunk_size], dt_frag)
                gamma_frag = T.alloc_fragment([chunk_size], T.float32)
                for cs in T.Parallel(chunk_size):
                    gamma_frag[cs] = dt_frag[cs] * T.sigmoid(trap_frag[cs])
                gamma_cached_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(gamma_frag, gamma_cached_frag)
                trap_scale_frag = T.alloc_fragment([chunk_size], dtype)
                for cs in T.Parallel(chunk_size):
                    trap_scale_frag[cs] = gamma_frag[cs] + shifted_gamma_frag[cs]
                trap_scale_shared = T.alloc_shared([chunk_size], dtype)
                T.copy(trap_scale_frag, trap_scale_shared)

                # --- DOUT projection (zero-masked at tail) ---
                dPhiO_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                if reduceO:
                    for cs, p in T.Parallel(chunk_size, P):
                        dout_shared[cs, p] = DOUT[i_b, chunk_start + cs, i_h, p]
                    if eff_tail < chunk_size:
                        for cs, p in T.Parallel(chunk_size, P):
                            dout_shared[cs, p] = T.if_then_else(cs < eff_tail, dout_shared[cs, p], 0.0)
                    for cs, r, p in T.Parallel(chunk_size, R, P):
                        dPhiO_frag[cs, r, p] = dout_shared[cs, p] * Phi_frag[r, p]
                else:
                    for cs, r, p in T.Parallel(chunk_size, R, P):
                        if packed_dout:
                            dPhiO_frag[cs, r, p] = DOUT[
                                i_b, i_h, fused_chunk_start + cs * R + r, p
                            ]
                        else:
                            dPhiO_frag[cs, r, p] = DOUT[i_b, chunk_start + cs, r, i_h, p]
                    if eff_tail < chunk_size:
                        for cs, r, p in T.Parallel(chunk_size, R, P):
                            dPhiO_frag[cs, r, p] = T.if_then_else(cs < eff_tail, dPhiO_frag[cs, r, p], 0.0)

                if hasZ:
                    Zeta_frag = T.alloc_fragment([R, P], dtype)
                    T.copy(MIMO_Z[i_h, :, :], Zeta_frag)
                    z_frag = T.alloc_fragment([chunk_size, P], dtype)
                    T.copy(Z[i_b, chunk_start:chunk_start + chunk_size, i_h, :], z_frag)
                    for cs, r, p in T.Parallel(chunk_size, R, P):
                        tmp = z_frag[cs, p] * Zeta_frag[r, p] * 0.5
                        dPhiO_frag[cs, r, p] *= tmp * T.tanh(tmp) + tmp
                T.copy(T.view(dPhiO_frag, shape=[fused_chunk_size, P]), dPhiO_shared)

                T.copy(V[i_b, chunk_start:chunk_start + chunk_size, i_h, :], v_shared)
                if hasD:
                    v_dD_frag = T.alloc_fragment([chunk_size, P], accum_dtype)
                    Psi_dD_frag = T.alloc_fragment([R, P], accum_dtype)
                    T.copy(v_shared, v_dD_frag)
                    T.copy(MIMO_V[i_h, :, :], Psi_dD_frag)
                    for cs, r, p in T.Parallel(chunk_size, R, P):
                        dPhiO_frag[cs, r, p] *= v_dD_frag[cs, p] * Psi_dD_frag[r, p]
                    T.reduce_sum(T.view(dPhiO_frag, shape=[fused_chunk_size * P]), dD_frag, clear=False)

                # --- Load and rotate Q ---
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    q_shared[cs * R + r, n] = Q[i_b, chunk_start + cs, r, i_h_qk, n]
                q_frag = T.alloc_fragment([chunk_size, R, N], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    q_frag[cs, r, n] = q_shared[cs * R + r, n]
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    q_frag[cs, r, n] += q_bias_frag[r, n]
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    q_shared[cs * R + r, n] = q_frag[cs, r, n]
                T.copy(q_shared, q_pre_rot_shared)
                q_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                q_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    q_first_half_frag[cs, r, n] = q_shared[cs * R + r, n]
                    q_second_half_frag[cs, r, n] = q_shared[cs * R + r, N // 2 + n]
                angles_frag = T.alloc_fragment([chunk_size, N // rotary_dim_divisor], T.float32)
                T.copy(ANGLES[i_b, chunk_start:chunk_start + chunk_size, i_h, :], angles_frag)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    q_shared[cs * R + r, n] = T.cos(angles_frag[cs, n]) * q_first_half_frag[cs, r, n] - T.sin(angles_frag[cs, n]) * q_second_half_frag[cs, r, n]
                    q_shared[cs * R + r, N // 2 + n] = T.sin(angles_frag[cs, n]) * q_first_half_frag[cs, r, n] + T.cos(angles_frag[cs, n]) * q_second_half_frag[cs, r, n]

                # --- Load and rotate K ---
                k_reshaped_shared = T.view(k_pre_trap_shared, shape=[chunk_size, R, N])
                T.copy(K[i_b, chunk_start:chunk_start + chunk_size, :, i_h_qk, :], k_reshaped_shared)
                k_frag = T.alloc_fragment([chunk_size, R, N], dtype)
                T.copy(k_reshaped_shared, k_frag)
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    k_frag[cs, r, n] += k_bias_frag[r, n]
                T.copy(k_frag, k_reshaped_shared)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    k_pre_rot_shared[csr, n] = k_pre_trap_shared[csr, n]
                k_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                k_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    k_first_half_frag[cs, r, n] = k_reshaped_shared[cs, r, n]
                    k_second_half_frag[cs, r, n] = k_reshaped_shared[cs, r, N // 2 + n]
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    k_reshaped_shared[cs, r, n] = T.cos(angles_frag[cs, n]) * k_first_half_frag[cs, r, n] - T.sin(angles_frag[cs, n]) * k_second_half_frag[cs, r, n]
                    k_reshaped_shared[cs, r, N // 2 + n] = T.sin(angles_frag[cs, n]) * k_first_half_frag[cs, r, n] + T.cos(angles_frag[cs, n]) * k_second_half_frag[cs, r, n]
                k_trap_scaled_frag = T.alloc_fragment([fused_chunk_size, N], dtype)
                T.copy(k_pre_trap_shared, k_trap_scaled_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    k_trap_scaled_frag[csr, n] *= trap_scale_shared[csr // R]
                T.copy(k_trap_scaled_frag, k_shared)

                # --- dPsiV: interchunk + intrachunk ---
                dPsiV_frag = T.alloc_fragment([fused_chunk_size, P], accum_dtype)
                T.gemm(k_shared, dstates_shared, dPsiV_frag, clear_accum=True)
                dA_cs_rev_frag = T.alloc_fragment([chunk_size], T.float32)
                dA_cs_rev_shared = T.alloc_shared([chunk_size], T.float32)
                T.copy(DA_CS_REV[i_b, i_h, chunk_start:chunk_start + chunk_size], dA_cs_rev_shared)
                T.copy(dA_cs_rev_shared, dA_cs_rev_frag)
                for csr, p in T.Parallel(fused_chunk_size, P):
                    dPsiV_frag[csr, p] *= T.exp(dA_cs_rev_frag[csr // R])

                lkq_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], accum_dtype)
                T.gemm(k_shared, q_shared, lkq_frag, transpose_B=True, clear_accum=True)
                T.copy(lkq_frag, lkq_masked__or__dkq_masked_shared)
                if R == 1:
                    lkq_masked_dtype_buf = T.alloc_fragment([fused_chunk_size, fused_chunk_size], dtype)
                    T.copy(lkq_masked__or__dkq_masked_shared, lkq_masked_dtype_buf)
                    for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                        lkq_masked_dtype_buf[csr_i, csr_j] = T.if_then_else(
                            csr_i // R < csr_j // R,
                            lkq_masked_dtype_buf[csr_i, csr_j] * T.exp(SEGSUM[i_b, i_h, global_chunk_idx, csr_j // R, csr_i // R]),
                            0.0)
                else:
                    for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                        lkq_frag[csr_i, csr_j] = T.if_then_else(
                            csr_i // R < csr_j // R,
                            lkq_frag[csr_i, csr_j] * T.exp(SEGSUM[i_b, i_h, global_chunk_idx, csr_j // R, csr_i // R]),
                            0.0)
                    lkq_masked_dtype_buf = T.alloc_shared([fused_chunk_size, fused_chunk_size], dtype)
                    T.copy(lkq_frag, lkq_masked_dtype_buf)
                T.gemm(lkq_masked_dtype_buf, dPhiO_shared, dPsiV_frag, clear_accum=False)

                # --- Diagonal contributions to dPsiV ---
                dPsiV_D_fused_frag = T.alloc_fragment([fused_chunk_size, P], accum_dtype)
                if hasD:
                    D_frag = T.alloc_var(T.float32)
                    T.copy(D[i_h], D_frag)
                    for csr, p in T.Parallel(fused_chunk_size, P):
                        dPsiV_D_fused_frag[csr, p] = dPsiV_frag[csr, p] + dPhiO_shared[csr, p] * D_frag
                else:
                    T.copy(dPsiV_frag, dPsiV_D_fused_frag)
                qk_dot_frag = T.alloc_fragment([chunk_size, R, R], dtype)
                T.copy(QK_DOT[i_b, i_h, chunk_start:chunk_start + chunk_size, :, :], qk_dot_shared)
                T.copy(qk_dot_shared, qk_dot_frag)
                gamma_dPsiV_frag = T.alloc_fragment([chunk_size], dtype)
                T.copy(gamma_frag, gamma_dPsiV_frag)
                for csr, p in T.Parallel(fused_chunk_size, P):
                    cs = csr // R
                    r_in = csr % R
                    for r_out in T.serial(R):
                        csr_out = cs * R + r_out
                        dPsiV_D_fused_frag[csr, p] += dPhiO_shared[csr_out, p] * qk_dot_frag[cs, r_out, r_in] * gamma_dPsiV_frag[cs]
                T.copy(dPsiV_D_fused_frag, dPsiV_combined_shared)

                # --- dV and dPsi ---
                dv_frag = T.alloc_fragment([chunk_size, P], dtype)
                T.clear(dv_frag)
                for cs, p in T.Parallel(chunk_size, P):
                    for r in T.serial(R):
                        dv_frag[cs, p] += dPsiV_combined_shared[cs * R + r, p] * Psi_frag[r, p]
                if eff_tail < chunk_size:
                    for cs, p in T.Parallel(chunk_size, P):
                        if cs < eff_tail:
                            DV[i_b, chunk_start + cs, i_h, p] = dv_frag[cs, p]
                else:
                    T.copy(dv_frag, DV[i_b, chunk_start:chunk_start + chunk_size, i_h, :])

                dPsi_frag = T.alloc_fragment([R, P], accum_dtype)
                T.copy(dPsi_acc, dPsi_frag)
                v_frag = T.alloc_fragment([chunk_size, P], accum_dtype)
                T.copy(v_shared, v_frag)
                for r, p in T.Parallel(R, P):
                    for cs in T.serial(chunk_size):
                        dPsi_frag[r, p] += dPsiV_combined_shared[cs * R + r, p] * v_frag[cs, p]
                T.copy(dPsi_frag, dPsi_acc)

                PsiV_frag = T.alloc_fragment([chunk_size, R, P], dtype)
                T.clear(PsiV_frag)
                for cs, p in T.Parallel(chunk_size, P):
                    for r in T.serial(R):
                        PsiV_frag[cs, r, p] += v_frag[cs, p] * Psi_frag[r, p]
                PsiV_shared = T.alloc_shared([fused_chunk_size, P], dtype)
                for cs, r, p in T.Parallel(chunk_size, R, P):
                    PsiV_shared[cs * R + r, p] = PsiV_frag[cs, r, p]

                # --- dqk_from_diag ---
                dqk_from_diag_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], accum_dtype)
                T.gemm(dPhiO_shared, PsiV_shared, dqk_from_diag_frag, transpose_B=True, clear_accum=True)
                dgamma_diag_prereduce_frag = T.alloc_fragment([chunk_size, R, R], accum_dtype)
                T.copy(qk_dot_shared, dgamma_diag_prereduce_frag)
                T.copy(dqk_from_diag_frag, dqk_from_diag_shared)
                for cs, r_out, r_in in T.Parallel(chunk_size, R, R):
                    dgamma_diag_prereduce_frag[cs, r_out, r_in] *= dqk_from_diag_shared[cs * R + r_out, cs * R + r_in]
                dgamma_diag_reduced_frag = T.alloc_fragment([chunk_size], accum_dtype)
                T.reduce_sum(T.view(dgamma_diag_prereduce_frag, shape=[chunk_size, R * R]),
                             dgamma_diag_reduced_frag, dim=-1, clear=True)
                if eff_tail < chunk_size:
                    for cs in T.Parallel(chunk_size):
                        if cs < eff_tail:
                            DGAMMA_DIAG[i_b, i_h, chunk_start + cs] = dgamma_diag_reduced_frag[cs]
                else:
                    T.copy(dgamma_diag_reduced_frag, DGAMMA_DIAG[i_b, i_h, chunk_start:chunk_start + chunk_size])
                gamma_qk_frag = T.alloc_fragment([chunk_size], accum_dtype)
                T.copy(gamma_cached_frag, gamma_qk_frag)
                for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                    dqk_from_diag_frag[csr_i, csr_j] *= gamma_qk_frag[csr_i // R]
                T.copy(dqk_from_diag_frag, dqk_from_diag_shared)

                # --- dK ---
                dk_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.gemm(PsiV_shared, dstates_shared, dk_frag, transpose_B=True, clear_accum=True)

                ddA_state_kv_prereduce_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(k_shared, ddA_state_kv_prereduce_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    ddA_state_kv_prereduce_frag[csr, n] *= dk_frag[csr, n]
                ddA_state_kv_prereduce_frag_reshaped = T.view(ddA_state_kv_prereduce_frag, shape=[chunk_size, R * N])
                ddA_state_kv_frag = T.alloc_fragment([chunk_size], accum_dtype)
                T.reduce_sum(ddA_state_kv_prereduce_frag_reshaped, ddA_state_kv_frag, dim=-1, clear=True)
                if eff_tail < chunk_size:
                    for cs in T.Parallel(chunk_size):
                        if cs < eff_tail:
                            DDA_CS_REV[i_b, i_h, chunk_start + cs] = ddA_state_kv_frag[cs]
                else:
                    T.copy(ddA_state_kv_frag, DDA_CS_REV[i_b, i_h, chunk_start:chunk_start + chunk_size])

                dA_cs_rev_dk_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(dA_cs_rev_shared, dA_cs_rev_dk_frag)
                for cs in T.Parallel(chunk_size):
                    dA_cs_rev_dk_frag[cs] = T.exp(dA_cs_rev_dk_frag[cs])
                for csr, n in T.Parallel(fused_chunk_size, N):
                    dk_frag[csr, n] *= dA_cs_rev_dk_frag[csr // R]

                dk_intrachunk_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], accum_dtype)
                T.gemm(PsiV_shared, dPhiO_shared, dk_intrachunk_frag, transpose_B=True, clear_accum=True)

                # DSSDA: contributions are auto-zeroed at tail because dPhiO is zero-masked
                kq_frag = T.alloc_fragment([fused_chunk_size, fused_chunk_size], dtype)
                T.copy(lkq_masked__or__dkq_masked_shared, kq_frag)
                for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                    kq_frag[csr_i, csr_j] *= dk_intrachunk_frag[csr_i, csr_j]
                kq_frag_reshaped = T.view(kq_frag, shape=[fused_chunk_size, chunk_size, R])
                interchunk_dda_prereduce_frag = T.alloc_fragment([fused_chunk_size, chunk_size], accum_dtype)
                T.reduce_sum(kq_frag_reshaped, interchunk_dda_prereduce_frag, dim=-1, clear=True)
                interchunk_dda_prereduce_frag_reshaped = T.view(interchunk_dda_prereduce_frag, shape=[chunk_size, R, chunk_size])
                interchunk_dda_frag = T.alloc_fragment([chunk_size, chunk_size], accum_dtype)
                T.reduce_sum(interchunk_dda_prereduce_frag_reshaped, interchunk_dda_frag, dim=1, clear=True)
                T.copy(interchunk_dda_frag, DSSDA[i_b, i_h, global_chunk_idx, :, :])

                for csr_i, csr_j in T.Parallel(fused_chunk_size, fused_chunk_size):
                    dk_intrachunk_frag[csr_i, csr_j] = T.if_then_else(
                        csr_i // R < csr_j // R,
                        dk_intrachunk_frag[csr_i, csr_j] * T.exp(SEGSUM[i_b, i_h, global_chunk_idx, csr_j // R, csr_i // R]),
                        0.0)
                T.copy(dk_intrachunk_frag, lkq_masked__or__dkq_masked_shared)
                T.copy(dk_frag, dk_shared)
                dk_nodiag_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(dk_shared, dk_nodiag_frag)
                T.gemm(lkq_masked__or__dkq_masked_shared, q_shared, dk_nodiag_frag, clear_accum=False)

                k_factor_frag = T.alloc_fragment([chunk_size, R, N], accum_dtype)
                T.copy(k_pre_trap_shared, T.view(k_factor_frag, shape=[fused_chunk_size, N]))
                dfactor_prereduce_frag = T.alloc_fragment([chunk_size, R, N], accum_dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N):
                    dfactor_prereduce_frag[cs, r, n] = k_factor_frag[cs, r, n] * dk_nodiag_frag[cs * R + r, n]
                dfactor_frag = T.alloc_fragment([chunk_size], accum_dtype)
                T.reduce_sum(T.view(dfactor_prereduce_frag, shape=[chunk_size, R * N]), dfactor_frag, dim=-1, clear=True)
                if eff_tail < chunk_size:
                    for cs in T.Parallel(chunk_size):
                        if cs < eff_tail:
                            DFACTOR[i_b, i_h, chunk_start + cs] = dfactor_frag[cs]
                else:
                    T.copy(dfactor_frag, DFACTOR[i_b, i_h, chunk_start:chunk_start + chunk_size])

                trap_scale_dk_frag = T.alloc_fragment([chunk_size], dtype)
                T.copy(trap_scale_shared, trap_scale_dk_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    dk_nodiag_frag[csr, n] *= trap_scale_dk_frag[csr // R]
                T.copy(dk_nodiag_frag, dk_shared)

                # --- State-passing ddA + interchunk dQ ---
                states_frag = T.alloc_fragment([N, P], T.float32)
                # Match dense bwd_bwd: load cached state through fp32 before
                # staging into the shared buffer used by GEMM.
                T.copy(STATES[i_b, i_h, global_chunk_idx, :, :], states_frag)
                T.copy(states_frag, states_shared)
                ddA_state_passing = T.alloc_fragment([1], T.float32)
                ddA_state_passing_prereduce_frag = T.alloc_fragment([N, P], T.float32)
                da_cs_sum = T.alloc_var(T.float32)
                T.copy(DA_CS[i_b, i_h, da_cs_end_idx], da_cs_sum)
                for n, p in T.Parallel(N, P):
                    ddA_state_passing_prereduce_frag[n, p] = (
                        states_frag[n, p] * dstates_frag[n, p] * T.exp(da_cs_sum))
                T.reduce_sum(T.view(ddA_state_passing_prereduce_frag, shape=[N * P]),
                             ddA_state_passing, dim=-1, clear=True)
                if eff_tail < chunk_size:
                    for cs in T.Parallel(chunk_size):
                        if cs < eff_tail:
                            DDA[i_b, i_h, chunk_start + cs] = ddA_state_passing[0]
                else:
                    dda_frag = T.alloc_fragment([chunk_size], T.float32)
                    for cs in T.Parallel(chunk_size):
                        dda_frag[cs] = ddA_state_passing[0]
                    T.copy(dda_frag, DDA[i_b, i_h, chunk_start:chunk_start + chunk_size])

                dq_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.gemm(dPhiO_shared, states_shared, dq_frag, transpose_B=True, clear_accum=True)

                dda_cs_prereduce_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(q_shared, dda_cs_prereduce_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    dda_cs_prereduce_frag[csr, n] *= dq_frag[csr, n]
                dda_cs_frag = T.alloc_fragment([chunk_size], accum_dtype)
                T.reduce_sum(T.view(dda_cs_prereduce_frag, shape=[chunk_size, R * N]),
                             dda_cs_frag, dim=-1, clear=True)
                if eff_tail < chunk_size:
                    for cs in T.Parallel(chunk_size):
                        if cs < eff_tail:
                            DDA_CS[i_b, i_h, chunk_start + cs] = dda_cs_frag[cs]
                else:
                    T.copy(dda_cs_frag, DDA_CS[i_b, i_h, chunk_start:chunk_start + chunk_size])

                dA_cs_dq_frag = T.alloc_fragment([chunk_size], T.float32)
                dA_cs_shared = T.alloc_shared([chunk_size], T.float32)
                T.copy(DA_CS[i_b, i_h, chunk_start:chunk_start + chunk_size], dA_cs_shared)
                T.copy(dA_cs_shared, dA_cs_dq_frag)
                for csr, n in T.Parallel(fused_chunk_size, N):
                    dq_frag[csr, n] *= T.exp(dA_cs_dq_frag[csr // R])
                T.copy(dq_frag, dq_shared)
                dq_combined_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(dq_shared, dq_combined_frag)
                T.gemm(lkq_masked__or__dkq_masked_shared, k_shared, dq_combined_frag, transpose_A=True, clear_accum=False)
                T.copy(dq_combined_frag, dq_shared)

                # --- Inverse rotary for dK and dQ + dAngles ---
                angles_dk_frag = T.alloc_fragment([chunk_size, N // rotary_dim_divisor], T.float32)
                T.copy(ANGLES[i_b, chunk_start:chunk_start + chunk_size, i_h, :], angles_dk_frag)
                dk_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                dk_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                k_prerot_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                k_prerot_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dk_first_half_frag[cs, r, n] = dk_shared[cs * R + r, n]
                    dk_second_half_frag[cs, r, n] = dk_shared[cs * R + r, N // 2 + n]
                    k_prerot_first_half_frag[cs, r, n] = k_pre_rot_shared[cs * R + r, n]
                    k_prerot_second_half_frag[cs, r, n] = k_pre_rot_shared[cs * R + r, N // 2 + n]
                dangle_dk_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], T.float32)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dangle_dk_frag[cs, r, n] = \
                        dk_first_half_frag[cs, r, n] * (-k_prerot_first_half_frag[cs, r, n] * T.sin(angles_dk_frag[cs, n]) - k_prerot_second_half_frag[cs, r, n] * T.cos(angles_dk_frag[cs, n])) + \
                        dk_second_half_frag[cs, r, n] * (k_prerot_first_half_frag[cs, r, n] * T.cos(angles_dk_frag[cs, n]) - k_prerot_second_half_frag[cs, r, n] * T.sin(angles_dk_frag[cs, n]))
                T.copy(T.view(dangle_dk_frag, shape=[fused_chunk_size, N // rotary_dim_divisor]), dangle_dk__or__dq_shared)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dk_shared[cs * R + r, n] = T.cos(angles_dk_frag[cs, n]) * dk_first_half_frag[cs, r, n] + T.sin(angles_dk_frag[cs, n]) * dk_second_half_frag[cs, r, n]
                    dk_shared[cs * R + r, N // 2 + n] = -T.sin(angles_dk_frag[cs, n]) * dk_first_half_frag[cs, r, n] + T.cos(angles_dk_frag[cs, n]) * dk_second_half_frag[cs, r, n]

                dk_combined_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(dk_shared, dk_combined_frag)
                q_dk_frag = T.alloc_fragment([fused_chunk_size, N], accum_dtype)
                T.copy(q_pre_rot_shared, q_dk_frag)
                q_dk_frag_reshaped = T.view(q_dk_frag, [chunk_size, R, N])
                for csr_in, n in T.Parallel(fused_chunk_size, N):
                    cs = csr_in // R
                    for r_out in T.serial(R):
                        csr_out = cs * R + r_out
                        dk_combined_frag[csr_in, n] += dqk_from_diag_shared[csr_out, csr_in] * q_dk_frag_reshaped[cs, r_out, n]
                if eff_tail < chunk_size:
                    for csr, n in T.Parallel(fused_chunk_size, N):
                        if csr < eff_tail * R:
                            DK[i_b, fused_chunk_start + csr, i_h, n] = dk_combined_frag[csr, n]
                else:
                    T.copy(dk_combined_frag, DK[i_b, fused_chunk_start:fused_chunk_start + fused_chunk_size, i_h, :])

                # dQ inverse rotary
                angles_dq_frag = T.alloc_fragment([chunk_size, N // rotary_dim_divisor], T.float32)
                T.copy(ANGLES[i_b, chunk_start:chunk_start + chunk_size, i_h, :], angles_dq_frag)
                dq_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                dq_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dq_first_half_frag[cs, r, n] = dq_shared[cs * R + r, n]
                    dq_second_half_frag[cs, r, n] = dq_shared[cs * R + r, N // 2 + n]
                q_prerot_first_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                q_prerot_second_half_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], dtype)
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    q_prerot_first_half_frag[cs, r, n] = q_pre_rot_shared[cs * R + r, n]
                    q_prerot_second_half_frag[cs, r, n] = q_pre_rot_shared[cs * R + r, N // 2 + n]
                dangle_dq_frag = T.alloc_fragment([chunk_size, R, N // rotary_dim_divisor], T.float32)
                T.copy(dangle_dk__or__dq_shared, T.view(dangle_dq_frag, shape=[fused_chunk_size, N // rotary_dim_divisor]))
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dangle_dq_frag[cs, r, n] += \
                        dq_first_half_frag[cs, r, n] * (-q_prerot_first_half_frag[cs, r, n] * T.sin(angles_dq_frag[cs, n]) - q_prerot_second_half_frag[cs, r, n] * T.cos(angles_dq_frag[cs, n])) + \
                        dq_second_half_frag[cs, r, n] * (q_prerot_first_half_frag[cs, r, n] * T.cos(angles_dq_frag[cs, n]) - q_prerot_second_half_frag[cs, r, n] * T.sin(angles_dq_frag[cs, n]))
                dangle_frag_reduced = T.alloc_fragment([chunk_size, N // rotary_dim_divisor], T.float32)
                T.clear(dangle_frag_reduced)
                for cs, n in T.Parallel(chunk_size, N // rotary_dim_divisor):
                    for r in T.serial(R):
                        dangle_frag_reduced[cs, n] += dangle_dq_frag[cs, r, n]
                if eff_tail < chunk_size:
                    for cs, n in T.Parallel(chunk_size, N // rotary_dim_divisor):
                        if cs < eff_tail:
                            DANGLES[i_b, chunk_start + cs, i_h, n] = dangle_frag_reduced[cs, n]
                else:
                    T.copy(dangle_frag_reduced, DANGLES[i_b, chunk_start:chunk_start + chunk_size, i_h, :])
                for cs, r, n in T.Parallel(chunk_size, R, N // rotary_dim_divisor):
                    dq_shared[cs * R + r, n] = T.cos(angles_dk_frag[cs, n]) * dq_first_half_frag[cs, r, n] + T.sin(angles_dk_frag[cs, n]) * dq_second_half_frag[cs, r, n]
                    dq_shared[cs * R + r, N // 2 + n] = -T.sin(angles_dk_frag[cs, n]) * dq_first_half_frag[cs, r, n] + T.cos(angles_dk_frag[cs, n]) * dq_second_half_frag[cs, r, n]
                T.copy(dq_shared, dq_frag)
                for csr_out, n in T.Parallel(fused_chunk_size, N):
                    cs = csr_out // R
                    for r_in in T.serial(R):
                        csr_in = cs * R + r_in
                        dq_frag[csr_out, n] += dqk_from_diag_shared[csr_out, csr_in] * k_pre_rot_shared[csr_in, n]
                if eff_tail < chunk_size:
                    for csr, n in T.Parallel(fused_chunk_size, N):
                        if csr < eff_tail * R:
                            DQ[i_b, fused_chunk_start + csr, i_h, n] = dq_frag[csr, n]
                else:
                    T.copy(dq_frag, DQ[i_b, fused_chunk_start:fused_chunk_start + fused_chunk_size, i_h, :])

                # --- Update reverse-passed state gradient ---
                da_cs_sum_dstates = T.alloc_var(T.float32)
                T.copy(DA_CS[i_b, i_h, da_cs_end_idx], da_cs_sum_dstates)
                for n, p in T.Parallel(N, P):
                    dstates_frag[n, p] *= T.exp(da_cs_sum_dstates)
                dPhiO_scaled_frag = T.alloc_fragment([fused_chunk_size, P], dtype)
                T.copy(dPhiO_shared, dPhiO_scaled_frag)
                dA_cs_dPhiO_frag = T.alloc_fragment([chunk_size], T.float32)
                T.copy(dA_cs_shared, dA_cs_dPhiO_frag)
                for csr, p in T.Parallel(fused_chunk_size, P):
                    dPhiO_scaled_frag[csr, p] *= T.exp(dA_cs_dPhiO_frag[csr // R])
                T.gemm(q_shared, dPhiO_scaled_frag, dstates_frag, transpose_A=True, clear_accum=False)
                T.copy(dstates_frag, dstates_shared)

            T.copy(dPsi_acc, DMIMO_V[i_b, i_h, i_ns, :, :])
            if hasD:
                T.copy(dD_frag, DD[i_b, i_h, i_ns])

    return mamba_mimo_bwd_bwd_kernel


def mamba_mimo_bwd_combined_varlen(
        dout,
        q,
        k,
        v,
        q_bias,
        k_bias,
        mimo_v,
        mimo_o,
        z,
        mimo_z,
        angles,
        dA_cs,
        dA_cs_rev,
        dt,
        trap,
        D,
        segsum,
        chunk_size,
        rotary_dim_divisor,
        dtype,
        cu_seqlens=None,
        bf_threads=128,
        bf_num_stages=0,
        bb_threads=256,
        bb_num_stages=0,
        states_dtype=None,
        fuse_pregate_headwise_rms_norm=False,
        outproj_norm_weight=None,
        outproj_norm_eps=1e-5,
        ):
    """
    Varlen combined backward pass (bwd-fwd + bwd-bwd).

    Dispatches to ``mamba_mimo_bwd_combined`` (non-varlen) when
    ``cu_seqlens`` is None; otherwise runs the two varlen TileLang kernels
    followed by the Triton utility kernels from
    ``mamba3_mimo_utils.py``.

    Args:
        dout:             Upstream gradient, shape ``[B, S, H, P]`` or
                          ``[B, S, R, H, P]`` depending on reduceO.
        q, k, v:          Packed activations.
        q_bias, k_bias:   Per-head per-R biases, shape ``[H, R, N]``.
        mimo_v:           Psi projection, shape ``[H, R, P]``.
        mimo_o:           Phi projection, shape ``[H, R, P]``, or None.
        z:                Optional gating tensor, shape ``[B, S, H, P]``.
        mimo_z:           Optional Zeta projection, shape ``[H, R, P]``.
        angles:           Rotary angles.
        dA_cs:            Forward cumsum of discretised A.
        dA_cs_rev:        Reverse cumsum of discretised A.
        dt:               Discretised time step.
        trap:             Pre-sigmoid trapezoidal modulator.
        D:                Optional skip-connection vector, shape ``[H]``.
        segsum:           Per-chunk lower-triangular log-decay matrix.
        chunk_size:       Tokens per chunk.
        rotary_dim_divisor: Divisor for the rotary head dimension.
        dtype:            Compute dtype (torch.dtype or string).
        cu_seqlens:       Optional int32 tensor of shape ``[NS+1]``.
                          None → dense (non-varlen) path.
        bf_threads:       Threads for the bwd-fwd kernel.
        bf_num_stages:    Pipeline stages for the bwd-fwd kernel.
        bb_threads:       Threads for the bwd-bwd kernel.
        bb_num_stages:    Pipeline stages for the bwd-bwd kernel.
        states_dtype:     Optional dtype for cached recurrent states.
        fuse_pregate_headwise_rms_norm:
                          If True, bwd-fwd also backpropagates through
                          RMSNorm(raw_y) * silu(z) and writes DOUT_PRE_RMS.
        outproj_norm_weight:
                          RMSNorm weight, shape ``[H, P]`` or flattened.
        outproj_norm_eps: RMSNorm epsilon.

    Returns:
        Tuple of gradients:
            (dq, dk, dv, ddA, ddt, dtrap,
             dq_bias, dk_bias, dmimo_v, dmimo_z, dmimo_o,
             dangles, dD, dz, dout_norm_weight)
    """
    if cu_seqlens is None:
        return mamba_mimo_bwd_combined(
            dout, q, k, v, q_bias, k_bias, mimo_v, mimo_o,
            z, mimo_z, angles, dA_cs, dA_cs_rev, dt, trap, D,
            segsum, chunk_size, rotary_dim_divisor, dtype,
            bf_threads, bf_num_stages, bb_threads, bb_num_stages,
            states_dtype=states_dtype if states_dtype is not None else v.dtype,
            fuse_pregate_headwise_rms_norm=fuse_pregate_headwise_rms_norm,
            outproj_norm_weight=outproj_norm_weight,
            outproj_norm_eps=outproj_norm_eps,
        )

    B, S, R, G, N = q.shape
    H, P = v.shape[-2], v.shape[-1]
    NS = cu_seqlens.shape[0] - 1
    reduceO = mimo_o is not None
    max_nchunks = (S // chunk_size) + NS
    states_dtype = v.dtype if states_dtype is None else states_dtype

    if fuse_pregate_headwise_rms_norm:
        if not reduceO:
            raise ValueError("fuse_pregate_headwise_rms_norm=True requires mimo_o.")
        if z is None or mimo_z is None:
            raise ValueError("fuse_pregate_headwise_rms_norm=True requires z and mimo_z.")

    if not fuse_pregate_headwise_rms_norm:
        outproj_norm_weight_arg = None
    elif outproj_norm_weight is None:
        outproj_norm_weight_arg = torch.ones((H, P), dtype=torch.float32, device=q.device)
    else:
        if outproj_norm_weight.ndim == 1:
            if outproj_norm_weight.numel() != H * P:
                raise ValueError(
                    f"Expected flattened outproj_norm_weight to have {H * P} elements, "
                    f"got {outproj_norm_weight.numel()}."
                )
            outproj_norm_weight_arg = outproj_norm_weight.reshape(H, P).contiguous()
        elif outproj_norm_weight.shape == (H, P):
            outproj_norm_weight_arg = outproj_norm_weight.contiguous()
        else:
            raise ValueError(
                f"Expected outproj_norm_weight to have shape ({H}, {P}) or ({H * P},), "
                f"got {tuple(outproj_norm_weight.shape)}."
            )
        outproj_norm_weight_arg = outproj_norm_weight_arg.to(device=q.device, dtype=torch.float32)

    if isinstance(dtype, torch.dtype):
        dtype_str = str(dtype).replace("torch.", "")
    else:
        dtype_str = dtype

    dmimo_o = torch.empty([B, H, NS, R, P], dtype=mimo_v.dtype, device=mimo_v.device) if reduceO else None
    dout_norm_weight = (
        torch.empty([B, H, NS, R, P], dtype=torch.float32, device=v.device)
        if fuse_pregate_headwise_rms_norm else None
    )
    dout_norm_weight_arg = dout_norm_weight
    dout_pre_rms = (
        torch.empty([B, H, S * R, P], dtype=dout.dtype, device=dout.device)
        if fuse_pregate_headwise_rms_norm
        else None
    )
    states = torch.empty([B, H, max_nchunks, N, P], dtype=states_dtype, device=v.device)

    if z is not None:
        dz_tilelang = torch.empty_like(v)
        dmimo_z = torch.empty([B, H, NS, R, P], dtype=mimo_v.dtype, device=mimo_v.device)
    else:
        dz_tilelang = None
        dmimo_z = None
    qk_dot = torch.zeros([B, H, S, R, R], dtype=q.dtype, device=q.device)

    bwd_fwd_kernel = mamba_mimo_bwd_fwd(
        B, H, G,
        N, P, R,
        z is not None, D is not None, reduceO,
        fuse_pregate_headwise_rms_norm,
        isVarlen=cu_seqlens is not None,
        chunk_size=chunk_size, rotary_dim_divisor=rotary_dim_divisor, dtype=dtype_str,
        outproj_norm_eps=outproj_norm_eps,
        threads=bf_threads, num_stages=bf_num_stages,
        states_dtype=states_dtype)
    ns_anchor = cu_seqlens[:-1]
    bwd_fwd_kernel(
        dout, q, k, v, q_bias, k_bias, mimo_v, mimo_o,
        outproj_norm_weight_arg,
        dmimo_o, dout_norm_weight_arg, dout_pre_rms, states,
        z, mimo_z, dz_tilelang, dmimo_z,
        angles, dA_cs, dA_cs_rev, dt, trap, D,
        qk_dot, segsum, ns_anchor, cu_seqlens,
    )
    if reduceO and not fuse_pregate_headwise_rms_norm:
        dmimo_o = dmimo_o.sum(dim=(0, 2))

    dq_tilelang = torch.empty([B, S, R, H, N], dtype=q.dtype, device=q.device)
    dk_tilelang = torch.empty([B, S, R, H, N], dtype=k.dtype, device=k.device)
    dv_tilelang = torch.empty_like(v)
    dmimo_v = torch.empty([B, H, NS, R, P], dtype=mimo_v.dtype, device=mimo_v.device)
    dD = torch.empty([B, H, NS], dtype=D.dtype, device=D.device) if D is not None else None
    dangles = torch.zeros([B, S, H, N // rotary_dim_divisor], dtype=angles.dtype, device=angles.device)
    dfactor = torch.zeros([B, H, S], dtype=torch.float32, device=trap.device)
    dgamma_diag = torch.zeros([B, H, S], dtype=torch.float32, device=trap.device)
    ddA = torch.zeros([B, H, S], dtype=torch.float32, device=dt.device)
    dSSdA = torch.zeros([B, H, max_nchunks, chunk_size, chunk_size], dtype=torch.float32, device=dt.device)
    ddA_cs_rev = torch.zeros([B, H, S], dtype=torch.float32, device=dt.device)
    ddA_cs = torch.zeros([B, H, S], dtype=torch.float32, device=dt.device)

    bwd_bwd_dout = dout_pre_rms if fuse_pregate_headwise_rms_norm else dout
    bwd_bwd_reduceO = reduceO and not fuse_pregate_headwise_rms_norm
    bwd_bwd_hasZ = (z is not None) and not fuse_pregate_headwise_rms_norm
    bwd_bwd_packed_dout = fuse_pregate_headwise_rms_norm
    bwd_bwd_kernel = mamba_mimo_bwd_bwd(
        B, H, G,
        N, P, R,
        bwd_bwd_hasZ, D is not None, bwd_bwd_reduceO,
        bwd_bwd_packed_dout,
        isVarlen=cu_seqlens is not None,
        chunk_size=chunk_size, rotary_dim_divisor=rotary_dim_divisor, dtype=dtype_str,
        threads=bb_threads, num_stages=bb_num_stages,
        states_dtype=states_dtype)
    bwd_bwd_kernel(
        bwd_bwd_dout, q, k, v, q_bias, k_bias, mimo_v, mimo_o,
        dk_tilelang.view(B, S * R, H, N),
        dv_tilelang, dmimo_v, states,
        dq_tilelang.view(B, S * R, H, N),
        z, mimo_z, angles, dA_cs, dA_cs_rev, dt, trap,
        dfactor, dgamma_diag, dangles, D, dD,
        qk_dot, ddA, dSSdA, ddA_cs_rev, ddA_cs,
        segsum, cu_seqlens,
    )

    dq_tilelang, dk_tilelang, dq_bias_tilelang, dk_bias_tilelang = (
        reduce_grouped_qk_grads_and_bias_triton(dq_tilelang, dk_tilelang, G)
    )
    dmimo_v = dmimo_v.sum(dim=(0, 2))
    if fuse_pregate_headwise_rms_norm:
        dmimo_o = dmimo_o.sum(dim=(0, 2))
        dmimo_z = dmimo_z.sum(dim=(0, 2))
        dout_norm_weight = dout_norm_weight.sum(dim=(0, 2, 3))
    else:
        dmimo_z = dmimo_z.sum(dim=(0, 2)) if dmimo_z is not None else None
    dD = dD.sum(dim=(0, 2)) if dD is not None else None

    ddt, dtrap = bwd_dtrap_ddt_triton_varlen(trap, dt, dfactor, dgamma_diag, chunk_size, cu_seqlens)

    ddA += bwd_dadt_fused_triton_varlen(
        dSSdA, segsum, ddA_cs, ddA_cs_rev, dA_cs, dA_cs_rev, chunk_size, cu_seqlens
    )

    return (dq_tilelang, dk_tilelang, dv_tilelang, 
            ddA, ddt, dtrap, dq_bias_tilelang, dk_bias_tilelang,
            dmimo_v, dmimo_z, dmimo_o, dangles, 
            dD, dz_tilelang, dout_norm_weight)
