Skip to main content

burn_mamba/mamba3/single_ssd/ssd/serial_recalculated/
combined_backward.rs

1//! # Recompute-based gradient math for the Mamba-3 single-SSD
2//!
3//! The analytic backward of the single-pass MIMO-first scan.  Forward
4//! intermediates (K1–K4) are recomputed from the saved leaf inputs, then a
5//! reverse per-chunk loop fuses the K5 state-to-output (BLUE), the strict
6//! lower-triangular intra-chunk (LOWER), and the K4 state-passing backwards; the
7//! γ-weighted same-step (DIAG) term is computed batched (no recurrence, tiny
8//! `m × m` tensors).  Because this pathway applies the trapezoid weights
9//! internally, it additionally returns `d_gamma` and `d_scale`.  The shared K3
10//! extended helper (and K1/K2/K4) are reused from the double-SSD module.
11//!
12//! Everything operates on backend **primitives** through the rank-tagged `F`
13//! wrapper: the custom [`Backward`](burn::backend::autodiff::ops::Backward) node
14//! runs with a generic backend `B`, so the high-level `Tensor` is unavailable
15//! and the math uses `B`'s `float_*` ops.
16
17#![allow(non_snake_case)]
18
19use crate::mamba3::double_ssd::ssd::serial_recalculated::combined_backward::k3_ssd_chunk_state_extended;
20use crate::mamba3::double_ssd::ssd::serial_recalculated::{
21    k1_ssd_chunk_cumsum, k2_ssd_bmm, k4_ssd_state_passing,
22};
23use crate::mamba3::single_ssd::ssd::serial_recalculated::diag::{
24    DiagGrads, y_diag_correction_backward,
25};
26use burn_stack::utils::fprim::{F, san};
27use burn::backend::Backend;
28use burn::tensor::s;
29
30/// Per-input gradients produced by [`combined_backward`] for the Single-SSD.
31/// Adds `d_gamma_bnlh` and `d_scale_bnlh` over the double-ssd form
32/// [`crate::mamba3::double_ssd::ssd::serial_recalculated::combined_backward::CombinedGrads`].
33#[non_exhaustive]
34pub struct CombinedSingleSsdGrads<B: Backend> {
35    /// Gradient of the raw input `v`.
36    pub d_v_bnlmhp: F<B, 6>,
37    /// Gradient of `Δ·A` (`da`).
38    pub d_da_bnlh: F<B, 4>,
39    /// Gradient of the input projection `B`.
40    pub d_b_bnlmhr: F<B, 6>,
41    /// Gradient of the output projection `C`.
42    pub d_c_bnlmhr: F<B, 6>,
43    /// Gradient of the same-step trapezoid weight `γ`.
44    pub d_gamma_bnlh: F<B, 4>,
45    /// Gradient of the key scale `scale = γ + (1−λ₊₁)·Δ₊₁`.
46    pub d_scale_bnlh: F<B, 4>,
47    /// Gradient of the initial SSM state.
48    pub d_initial_state_bhpr: F<B, 4>,
49}
50
51/// Memory-efficient backward for the Mamba-3 MIMO-first chunkwise Single-SSD.
52///
53/// Recomputes the forward intermediates (K1–K4) from the saved inputs, then:
54/// - runs a reverse per-chunk loop that fuses the K5 BLUE (state-to-output) and
55///   the strict lower-triangular LOWER (intra-chunk) backward with the K4
56///   state-passing backward, and
57/// - computes the γ-weighted same-step DIAG backward batched (it has no
58///   recurrence, and the `m × m` working tensors are tiny).
59///
60/// K3/K2/K1 backwards run as single batched ops once the loop has collected all
61/// per-chunk slices.
62///
63/// # Arguments
64/// - `d_y_bnlmhp` — upstream gradient of the SSD output
65/// - `d_final_bhpr` — upstream gradient of the final SSM state
66/// - `v_bnlmhp`, `da_bnlh`, `b_bnlmhr`, `c_bnlmhr`, `gamma_bnlh`, `scale_bnlh`,
67///   `initial_state_bhpr` — the seven saved forward inputs
68/// - `siso_specialization` — the forward's γ-correction branch choice, replayed
69///   here so the backward matches it (performance-only; both agree)
70///
71/// # Returns
72/// One [`CombinedSingleSsdGrads`] with gradients for all 7 inputs.
73#[allow(clippy::too_many_arguments)]
74pub fn combined_backward<B: Backend>(
75    d_y_bnlmhp: F<B, 6>,
76    d_final_bhpr: F<B, 4>,
77    //
78    v_bnlmhp: F<B, 6>,
79    da_bnlh: F<B, 4>,
80    b_bnlmhr: F<B, 6>,
81    c_bnlmhr: F<B, 6>,
82    gamma_bnlh: F<B, 4>,
83    scale_bnlh: F<B, 4>,
84    initial_state_bhpr: F<B, 4>,
85    siso_specialization: bool,
86) -> CombinedSingleSsdGrads<B> {
87    let [batch, nchunks, chunk_len, mimo_rank, nheads, per_head_dim] = v_bnlmhp.dims();
88    let [.., state_rank] = b_bnlmhr.dims();
89    let device = v_bnlmhp.device();
90    let dtype = v_bnlmhp.dtype();
91
92    san(&d_y_bnlmhp);
93    san(&d_final_bhpr);
94    san(&v_bnlmhp);
95    san(&da_bnlh);
96    san(&b_bnlmhr);
97    san(&c_bnlmhr);
98    san(&gamma_bnlh);
99    san(&scale_bnlh);
100    san(&initial_state_bhpr);
101
102    // ═══════════════════════════════════════════════════════════════════════
103    // RECOMPUTE FORWARD INTERMEDIATES (K1–K4, single-ssd form)
104    // ═══════════════════════════════════════════════════════════════════════
105
106    // K1
107    let (da_cumsum_bhnl, da_chunk_end_bhn) = k1_ssd_chunk_cumsum(da_bnlh.clone());
108    san(&da_cumsum_bhnl);
109
110    // K2 — CB matrix (unscaled), used by LOWER.
111    let cb_bnhLMLM = k2_ssd_bmm(c_bnlmhr.clone(), b_bnlmhr.clone());
112    san(&cb_bnhLMLM);
113
114    // K3 — chunk state from K_scaled = scaleₜ·B.
115    let scale_bnlh11 = scale_bnlh.clone().unsqueeze_dims::<6>(&[3, 5]);
116    let k_scaled_bnlmhr = b_bnlmhr.clone() * scale_bnlh11.clone();
117    let (intra_chunk_state_bnhpr, k3_decay_bhnLM, k3_decayed_v_bnLMhp) =
118        k3_ssd_chunk_state_extended(
119            v_bnlmhp.clone(),
120            k_scaled_bnlmhr.clone(),
121            da_cumsum_bhnl.clone(),
122        );
123
124    // K4 — chunk-input state stream consumed by BLUE.
125    let (chunk_input_state_bnhpr, _final_state_bhpr) = k4_ssd_state_passing(
126        intra_chunk_state_bnhpr,
127        da_chunk_end_bhn.clone(),
128        initial_state_bhpr,
129    );
130
131    // Fused-position cumulative decay.
132    let da_cumsum_bhnLM = da_cumsum_bhnl
133        .clone()
134        .unsqueeze_dim::<5>(4)
135        .expand([batch, nheads, nchunks, chunk_len, mimo_rank])
136        .reshape([batch, nheads, nchunks, chunk_len * mimo_rank]);
137
138    // d_y in (batch, nchunks, nheads, chunk_len·mimo_rank, per_head_dim) ordering.
139    let d_y_bnhLMp = d_y_bnlmhp
140        .clone()
141        .reshape([batch, nchunks, chunk_len * mimo_rank, nheads, per_head_dim])
142        .swap_dims(2, 3);
143    san(&d_y_bnhLMp);
144
145    // ═══════════════════════════════════════════════════════════════════════
146    // DIAG BACKWARD (batched — no recurrence; m × m working set is tiny)
147    //
148    // Forward (per (b,n,l,h)):
149    //   qk_dot[m_out, m_in] = Σ_r C[m_out, r] · B[m_in, r]
150    //   y_diag[m_out, p]    = γ · Σ_{m_in} qk_dot[m_out, m_in] · V[m_in, p]
151    // ═══════════════════════════════════════════════════════════════════════
152    let DiagGrads {
153        d_v_bnlmhp: d_v_diag_bnlmhp,
154        d_c_bnlmhr: d_c_diag_bnlmhr,
155        d_b_bnlmhr: d_b_diag_bnlmhr,
156        d_gamma_bnlh,
157    } = y_diag_correction_backward(
158        d_y_bnlmhp.clone(),
159        v_bnlmhp.clone(),
160        b_bnlmhr.clone(),
161        c_bnlmhr.clone(),
162        gamma_bnlh.clone(),
163        siso_specialization,
164    );
165    san(&d_gamma_bnlh);
166
167    // Reusable [chunk_len, chunk_len] -inf strict-upper mask (triu(0): on+above
168    // diagonal → -inf) for the LOWER (strict lower-triangular) path.
169    let neg_inf_strict_ll: F<B, 2> =
170        F::<B, 2>::full([chunk_len, chunk_len], f32::NEG_INFINITY, &device, dtype).triu(0);
171
172    // ═══════════════════════════════════════════════════════════════════════
173    // REVERSE PER-CHUNK LOOP — K5 (BLUE + LOWER) + K4 fused
174    // ═══════════════════════════════════════════════════════════════════════
175    let mut vec_lower_d_v_bhLMp: Vec<F<B, 4>> = Vec::with_capacity(nchunks);
176    let mut vec_blue_d_c_bhLMr: Vec<F<B, 4>> = Vec::with_capacity(nchunks);
177    let mut vec_d_cb_bhLMLM: Vec<F<B, 4>> = Vec::with_capacity(nchunks);
178    let mut vec_blue_d_da_bhl: Vec<F<B, 3>> = Vec::with_capacity(nchunks);
179    let mut vec_lower_d_da_bhl: Vec<F<B, 3>> = Vec::with_capacity(nchunks);
180    let mut vec_lower_d_scale_bhl: Vec<F<B, 3>> = Vec::with_capacity(nchunks);
181    let mut vec_d_intra_bhpr: Vec<F<B, 4>> = Vec::with_capacity(nchunks);
182    let mut vec_d_da_end_bh: Vec<F<B, 2>> = Vec::with_capacity(nchunks);
183
184    let mut d_running_state_bhpr: F<B, 4> = d_final_bhpr;
185
186    for i_chunk in (0..nchunks).rev() {
187        // ── Per-chunk slices (fused chunk_len · mimo_rank) ─────────────
188        let v_bhLMp: F<B, 4> = v_bnlmhp
189            .clone()
190            .slice(s![.., i_chunk, .., .., .., ..])
191            .squeeze_dim::<5>(1)
192            .reshape([batch, chunk_len * mimo_rank, nheads, per_head_dim])
193            .swap_dims(1, 2);
194
195        let c_bhLMr: F<B, 4> = c_bnlmhr
196            .clone()
197            .slice(s![.., i_chunk, .., .., .., ..])
198            .squeeze_dim::<5>(1)
199            .reshape([batch, chunk_len * mimo_rank, nheads, state_rank])
200            .swap_dims(1, 2);
201
202        let cb_bhLMLM: F<B, 4> = cb_bnhLMLM
203            .clone()
204            .slice(s![.., i_chunk, .., .., ..])
205            .squeeze_dim::<4>(1);
206
207        let da_cumsum_bhLM: F<B, 3> = da_cumsum_bhnLM
208            .clone()
209            .slice(s![.., .., i_chunk, ..])
210            .squeeze_dim::<3>(2);
211
212        // scaleₜ per fused source position: scale[s_time] broadcast over s_m.
213        let scale_bhLM: F<B, 3> = scale_bnlh
214            .clone()
215            .slice(s![.., i_chunk, .., ..]) // [b, l, h]
216            .squeeze_dim::<3>(1)
217            .swap_dims(1, 2) // [b, h, l]
218            .unsqueeze_dim::<4>(3) // [b, h, l, 1]
219            .expand([batch, nheads, chunk_len, mimo_rank])
220            .reshape([batch, nheads, chunk_len * mimo_rank]);
221
222        let chunk_input_state_bhpr: F<B, 4> = chunk_input_state_bnhpr
223            .clone()
224            .slice(s![.., i_chunk, .., .., ..])
225            .squeeze_dim::<4>(1);
226        san(&chunk_input_state_bhpr);
227
228        let d_y_bhLMp: F<B, 4> = d_y_bnhLMp
229            .clone()
230            .slice(s![.., i_chunk, .., .., ..])
231            .squeeze_dim::<4>(1);
232
233        // ── BLUE backward (identical to double-ssd form) ─────────────────
234        let exp_da_cumsum_bhLM: F<B, 3> = da_cumsum_bhLM.clone().exp();
235        let exp_da_cumsum_bhLMp: F<B, 4> = exp_da_cumsum_bhLM
236            .clone()
237            .unsqueeze_dim::<4>(3)
238            .expand([batch, nheads, chunk_len * mimo_rank, per_head_dim]);
239        let d_ch_bhLMp: F<B, 4> = d_y_bhLMp.clone() * exp_da_cumsum_bhLMp.clone();
240        san(&d_ch_bhLMp);
241
242        let d_chunk_input_state_bhpr: F<B, 4> = c_bhLMr
243            .clone()
244            .transpose() // c_bhrLM
245            .matmul(d_ch_bhLMp.clone()) // bhrp
246            .transpose(); // bhpr
247        san(&d_chunk_input_state_bhpr);
248
249        let d_c_blue_bhLMr: F<B, 4> = d_ch_bhLMp.clone().matmul(chunk_input_state_bhpr.clone());
250        vec_blue_d_c_bhLMr.push(d_c_blue_bhLMr);
251
252        let ch_bhLMp: F<B, 4> = c_bhLMr
253            .clone()
254            .matmul(chunk_input_state_bhpr.clone().transpose());
255        let d_da_blue_bhLM: F<B, 3> = (d_y_bhLMp.clone() * ch_bhLMp * exp_da_cumsum_bhLMp)
256            .sum_dim(3)
257            .squeeze_dim::<3>(3);
258        let d_da_blue_bhl: F<B, 3> = d_da_blue_bhLM
259            .reshape([batch, nheads, chunk_len, mimo_rank])
260            .sum_dim(3)
261            .squeeze_dim::<3>(3);
262        vec_blue_d_da_bhl.push(d_da_blue_bhl);
263
264        // ── LOWER backward (strict lower-tri + per-column scale) ────────
265        let da_target_bhLMLM: F<B, 4> = da_cumsum_bhLM.clone().unsqueeze_dim::<4>(3).expand([
266            batch,
267            nheads,
268            chunk_len * mimo_rank,
269            chunk_len * mimo_rank,
270        ]);
271        let da_source_bhLMLM: F<B, 4> = da_cumsum_bhLM.unsqueeze_dim::<4>(2).expand([
272            batch,
273            nheads,
274            chunk_len * mimo_rank,
275            chunk_len * mimo_rank,
276        ]);
277        let diff_bhLMLM = da_target_bhLMLM - da_source_bhLMLM;
278
279        // Strict-lower MIMO mask: -inf where s_time ≥ t_time — interleaved
280        // expansion of the [l, l] strict-upper (triu(0)) base mask.
281        let neg_inf_mimo_bhLMLM: F<B, 4> = neg_inf_strict_ll
282            .clone()
283            .unsqueeze_dims::<4>(&[0, 1])
284            .expand([batch, nheads, chunk_len, chunk_len])
285            .unsqueeze_dim::<5>(3)
286            .expand([batch, nheads, chunk_len, mimo_rank, chunk_len])
287            .reshape([batch, nheads, chunk_len * mimo_rank, chunk_len])
288            .unsqueeze_dim::<5>(4)
289            .expand([batch, nheads, chunk_len * mimo_rank, chunk_len, mimo_rank])
290            .reshape([batch, nheads, chunk_len * mimo_rank, chunk_len * mimo_rank]);
291        let decay_strict_bhLMLM = (diff_bhLMLM + neg_inf_mimo_bhLMLM).exp();
292        san(&decay_strict_bhLMLM);
293
294        let scale_col_bhLMLM: F<B, 4> = scale_bhLM
295            .unsqueeze_dim::<4>(2) // [b,h,1,LMs]
296            .expand([batch, nheads, chunk_len * mimo_rank, chunk_len * mimo_rank]);
297
298        // w = cb · decay_strict · scale_col
299        let prod_bhLMLM = cb_bhLMLM.clone() * decay_strict_bhLMLM.clone();
300        let w_bhLMLM = prod_bhLMLM.clone() * scale_col_bhLMLM.clone();
301
302        // d_w = d_y · vᵀ
303        let d_w_bhLMLM: F<B, 4> = d_y_bhLMp.clone().matmul(v_bhLMp.clone().transpose());
304        san(&d_w_bhLMLM);
305
306        // d_v_lower = wᵀ · d_y
307        let d_v_lower_bhLMp: F<B, 4> = w_bhLMLM.transpose().matmul(d_y_bhLMp.clone());
308        san(&d_v_lower_bhLMp);
309        vec_lower_d_v_bhLMp.push(d_v_lower_bhLMp);
310
311        // d_prod = d_w · scale_col ; d_scale_at = d_w · prod
312        let d_prod_bhLMLM = d_w_bhLMLM.clone() * scale_col_bhLMLM;
313        let d_scale_at_bhLMLM = d_w_bhLMLM * prod_bhLMLM;
314
315        // d_cb_lower = d_prod · decay_strict
316        let d_cb_lower_bhLMLM = d_prod_bhLMLM.clone() * decay_strict_bhLMLM.clone();
317        vec_d_cb_bhLMLM.push(d_cb_lower_bhLMLM);
318
319        // d_decay_strict = d_prod · cb ; d_diff = d_decay_strict · decay_strict
320        let d_decay_strict_bhLMLM = d_prod_bhLMLM * cb_bhLMLM;
321        let d_diff_bhLMLM = d_decay_strict_bhLMLM * decay_strict_bhLMLM;
322
323        let d_da_target_bhLM: F<B, 3> = d_diff_bhLMLM.clone().sum_dim(3).squeeze_dim::<3>(3);
324        let d_da_source_bhLM: F<B, 3> = d_diff_bhLMLM.sum_dim(2).squeeze_dim::<3>(2);
325        let d_da_lower_bhLM = d_da_target_bhLM - d_da_source_bhLM;
326        let d_da_lower_bhl: F<B, 3> = d_da_lower_bhLM
327            .reshape([batch, nheads, chunk_len, mimo_rank])
328            .sum_dim(3)
329            .squeeze_dim::<3>(3);
330        vec_lower_d_da_bhl.push(d_da_lower_bhl);
331
332        // d_scale[s_time] = Σ_{LMt, s_m} d_scale_at[LMt, LMs]
333        let d_scale_lower_bhl: F<B, 3> = d_scale_at_bhLMLM
334            .sum_dim(2) // sum over target LMt → [b,h,1,LMs]
335            .squeeze_dim::<3>(2) // [b,h,LMs]
336            .reshape([batch, nheads, chunk_len, mimo_rank])
337            .sum_dim(3) // sum over source mimo → [b,h,l,1]
338            .squeeze_dim::<3>(3); // [b,h,l]
339        vec_lower_d_scale_bhl.push(d_scale_lower_bhl);
340
341        // ── K4 backward step ───────────────────────────────────────────
342        vec_d_intra_bhpr.push(d_running_state_bhpr.clone());
343
344        let decay_chunk_bhpr: F<B, 4> = da_chunk_end_bhn
345            .clone()
346            .slice(s![.., .., i_chunk])
347            .exp()
348            .unsqueeze_dim::<4>(3)
349            .expand([batch, nheads, per_head_dim, state_rank]);
350        san(&decay_chunk_bhpr);
351
352        let d_decay_chunk_bhpr = d_running_state_bhpr.clone() * chunk_input_state_bhpr;
353        let d_da_chunk_end_bh: F<B, 2> = (d_decay_chunk_bhpr * decay_chunk_bhpr.clone())
354            .reshape([batch, nheads, per_head_dim * state_rank])
355            .sum_dim(2)
356            .squeeze_dim::<2>(2);
357        vec_d_da_end_bh.push(d_da_chunk_end_bh);
358
359        d_running_state_bhpr = decay_chunk_bhpr * d_running_state_bhpr + d_chunk_input_state_bhpr;
360        san(&d_running_state_bhpr);
361    }
362    let d_initial_state_bhpr = d_running_state_bhpr;
363
364    // ── Restore natural (forward) chunk order ─────────────────────────────
365    vec_lower_d_v_bhLMp.reverse();
366    vec_blue_d_c_bhLMr.reverse();
367    vec_d_cb_bhLMLM.reverse();
368    vec_blue_d_da_bhl.reverse();
369    vec_lower_d_da_bhl.reverse();
370    vec_lower_d_scale_bhl.reverse();
371    vec_d_intra_bhpr.reverse();
372    vec_d_da_end_bh.reverse();
373
374    // ── Stack per-chunk slices back into batched tensors ──────────────────
375    let d_v_lower_bnhLMp: F<B, 5> = F::stack(vec_lower_d_v_bhLMp, 1);
376    let d_c_blue_bnhLMr: F<B, 5> = F::stack(vec_blue_d_c_bhLMr, 1);
377    let d_cb_bnhLMLM: F<B, 5> = F::stack(vec_d_cb_bhLMLM, 1);
378    let d_da_blue_bhnl: F<B, 4> = F::stack(vec_blue_d_da_bhl, 2);
379    let d_da_lower_bhnl: F<B, 4> = F::stack(vec_lower_d_da_bhl, 2);
380    let d_scale_lower_bhnl: F<B, 4> = F::stack(vec_lower_d_scale_bhl, 2);
381    let d_intra_chunk_state_bnhpr: F<B, 5> = F::stack(vec_d_intra_bhpr, 1);
382    let d_da_end_bhn: F<B, 3> = F::stack(vec_d_da_end_bh, 2);
383    let d_da_cumsum_k4_bhnl: F<B, 4> = {
384        let zeros = F::<B, 4>::zeros([batch, nheads, nchunks, chunk_len - 1], &device, dtype);
385        let d_da_end_bhn1 = d_da_end_bhn.unsqueeze_dim::<4>(3);
386        F::cat(vec![zeros, d_da_end_bhn1], 3)
387    };
388
389    // ═══════════════════════════════════════════════════════════════════════
390    // K3 BACKWARD (batched) — K_scaled = scaleₜ·B
391    //
392    // intra_state = decayed_vᵀ @ K_scaled, with decayed_v = decay·V.
393    // d_K_scaled = decayed_vᵀ-contraction ; then split into d_b_k3 (·scale) and
394    // d_scale_k3 (Σ_{m,r} ·B).
395    // ═══════════════════════════════════════════════════════════════════════
396    let v_bnLMhp =
397        v_bnlmhp
398            .clone()
399            .reshape([batch, nchunks, chunk_len * mimo_rank, nheads, per_head_dim]);
400    let k_scaled_bnLMhr =
401        k_scaled_bnlmhr.reshape([batch, nchunks, chunk_len * mimo_rank, nheads, state_rank]);
402    let k_scaled_bnhLMr = k_scaled_bnLMhr.swap_dims(2, 3);
403    let decayed_v_bnhpLM = k3_decayed_v_bnLMhp.permute([0, 1, 3, 4, 2]);
404
405    let d_decayed_v_bnhpLM: F<B, 5> = d_intra_chunk_state_bnhpr
406        .clone()
407        .matmul(k_scaled_bnhLMr.clone().transpose()); // k_scaled_bnhrLM
408    let d_k_scaled_bnhLMr: F<B, 5> = decayed_v_bnhpLM
409        .transpose() // decayed_v_bnhLMp
410        .matmul(d_intra_chunk_state_bnhpr);
411
412    let d_decayed_v_bnLMhp = d_decayed_v_bnhpLM.permute([0, 1, 4, 2, 3]);
413    let d_decay_bhnLM: F<B, 4> = (d_decayed_v_bnLMhp.clone() * v_bnLMhp)
414        .sum_dim(4)
415        .squeeze_dim::<4>(4)
416        .permute([0, 3, 1, 2]);
417
418    let k3_decay_bnLMh1 = k3_decay_bhnLM
419        .clone()
420        .permute([0, 2, 3, 1])
421        .unsqueeze_dim::<5>(4);
422    let d_v_k3_bnLMhp: F<B, 5> = d_decayed_v_bnLMhp * k3_decay_bnLMh1;
423    let d_v_k3_bnlmhp: F<B, 6> =
424        d_v_k3_bnLMhp.reshape([batch, nchunks, chunk_len, mimo_rank, nheads, per_head_dim]);
425
426    // d(cumA_last − cumA) = d_decay · decay
427    let d_decay_times_decay_bhnLM = d_decay_bhnLM * k3_decay_bhnLM;
428    let d_a_cumsum_last_bhn: F<B, 3> = d_decay_times_decay_bhnLM
429        .clone()
430        .sum_dim(3)
431        .squeeze_dim::<3>(3);
432    let d_da_cumsum_bhnLM = -d_decay_times_decay_bhnLM;
433
434    let d_da_cumsum_k3_from_fused_bhnl: F<B, 4> = d_da_cumsum_bhnLM
435        .reshape([batch, nheads, nchunks, chunk_len, mimo_rank])
436        .sum_dim(4)
437        .squeeze_dim::<4>(4);
438    let d_da_cumsum_k3_from_last_bhnl: F<B, 4> = {
439        let zeros = F::<B, 4>::zeros([batch, nheads, nchunks, chunk_len - 1], &device, dtype);
440        let d_last = d_a_cumsum_last_bhn.unsqueeze_dim::<4>(3);
441        F::cat(vec![zeros, d_last], 3)
442    };
443    let d_da_cumsum_k3_bhnl = d_da_cumsum_k3_from_fused_bhnl + d_da_cumsum_k3_from_last_bhnl;
444
445    // d_K_scaled → bnlmhr, then split into d_b_k3 (·scale) and d_scale_k3 (Σ·B).
446    let d_k_scaled_bnlmhr: F<B, 6> = d_k_scaled_bnhLMr
447        .swap_dims(2, 3) // bnLMhr
448        .reshape([batch, nchunks, chunk_len, mimo_rank, nheads, state_rank]);
449    let d_b_k3_bnlmhr: F<B, 6> = d_k_scaled_bnlmhr.clone() * scale_bnlh11;
450    let d_scale_k3_bnlh: F<B, 4> = (d_k_scaled_bnlmhr * b_bnlmhr.clone())
451        .sum_dim(5) // sum over state_rank → [b,n,l,m,h,1]
452        .squeeze_dim::<5>(5) // [b,n,l,m,h]
453        .sum_dim(3) // sum over mimo_rank → [b,n,l,1,h]
454        .squeeze_dim::<4>(3); // [b,n,l,h]
455
456    // ═══════════════════════════════════════════════════════════════════════
457    // K2 BACKWARD (batched) — cb = C @ Bᵀ
458    // ═══════════════════════════════════════════════════════════════════════
459    let b_bnLMhr =
460        b_bnlmhr
461            .clone()
462            .reshape([batch, nchunks, chunk_len * mimo_rank, nheads, state_rank]);
463    let c_bnhLMr = c_bnlmhr
464        .clone()
465        .reshape([batch, nchunks, chunk_len * mimo_rank, nheads, state_rank])
466        .swap_dims(2, 3);
467    let b_for_k2_bnhLMr = b_bnLMhr.swap_dims(2, 3);
468
469    let d_c_k2_bnhLMr: F<B, 5> = d_cb_bnhLMLM.clone().matmul(b_for_k2_bnhLMr);
470    let d_b_k2_bnhrLM: F<B, 5> = c_bnhLMr.transpose().matmul(d_cb_bnhLMLM);
471
472    let d_c_k2_bnlmhr: F<B, 6> = d_c_k2_bnhLMr
473        .swap_dims(2, 3)
474        .reshape([batch, nchunks, chunk_len, mimo_rank, nheads, state_rank]);
475    let d_b_k2_bnlmhr: F<B, 6> = d_b_k2_bnhrLM
476        .permute([0, 1, 4, 2, 3])
477        .reshape([batch, nchunks, chunk_len, mimo_rank, nheads, state_rank]);
478
479    // ── Unstack d_c_blue / d_v_lower and reshape back ─────────────────────
480    let d_c_blue_bnlmhr: F<B, 6> = d_c_blue_bnhLMr
481        .swap_dims(2, 3)
482        .reshape([batch, nchunks, chunk_len, mimo_rank, nheads, state_rank]);
483    let d_v_lower_bnlmhp: F<B, 6> = d_v_lower_bnhLMp.swap_dims(2, 3).reshape([
484        batch,
485        nchunks,
486        chunk_len,
487        mimo_rank,
488        nheads,
489        per_head_dim,
490    ]);
491
492    // ═══════════════════════════════════════════════════════════════════════
493    // K1 BACKWARD + SUM CONTRIBUTIONS
494    // ═══════════════════════════════════════════════════════════════════════
495    let d_da_cumsum_bhnl =
496        d_da_blue_bhnl + d_da_lower_bhnl + d_da_cumsum_k3_bhnl + d_da_cumsum_k4_bhnl;
497    san(&d_da_cumsum_bhnl);
498
499    // K1 inverse: suffix sum.
500    let d_da_bhnl = {
501        let d_total_bhnl = d_da_cumsum_bhnl
502            .clone()
503            .sum_dim(3)
504            .expand([batch, nheads, nchunks, chunk_len]);
505        let prefix_bhnl = d_da_cumsum_bhnl.cumsum(3);
506        let zeros_bhn1 = F::<B, 4>::zeros([batch, nheads, nchunks, 1], &device, dtype);
507        let prefix_shifted_bhnl =
508            F::cat(vec![zeros_bhn1, prefix_bhnl.narrow(3, 0, chunk_len - 1)], 3);
509        d_total_bhnl - prefix_shifted_bhnl
510    };
511    let d_da_bnlh = d_da_bhnl.permute([0, 2, 3, 1]);
512
513    // ── Combine per-input gradient contributions ──────────────────────────
514    let d_v_bnlmhp = d_v_k3_bnlmhp + d_v_lower_bnlmhp + d_v_diag_bnlmhp;
515    let d_b_bnlmhr = d_b_k2_bnlmhr + d_b_k3_bnlmhr + d_b_diag_bnlmhr;
516    let d_c_bnlmhr = d_c_k2_bnlmhr + d_c_blue_bnlmhr + d_c_diag_bnlmhr;
517    let d_scale_bnlh = d_scale_lower_bhnl.permute([0, 2, 3, 1]) + d_scale_k3_bnlh;
518
519    san(&d_v_bnlmhp);
520    san(&d_da_bnlh);
521    san(&d_b_bnlmhr);
522    san(&d_c_bnlmhr);
523    san(&d_gamma_bnlh);
524    san(&d_scale_bnlh);
525    san(&d_initial_state_bhpr);
526
527    CombinedSingleSsdGrads {
528        d_v_bnlmhp,
529        d_da_bnlh,
530        d_b_bnlmhr,
531        d_c_bnlmhr,
532        d_gamma_bnlh,
533        d_scale_bnlh,
534        d_initial_state_bhpr,
535    }
536}