Skip to main content

burn_mamba/mamba2/ssd/
serial.rs

1//! # Serial-over-chunks SSD (Mamba-2)
2//!
3//! The chunkwise SSD scan expressed as a serial loop over chunks, mirroring the
4//! five Triton kernels of the reference `ssd_combined.py` (`ssd_chunk_state.py`,
5//! `ssd_bmm.py`, `ssd_state_passing.py`, `ssd_chunk_scan.py`):
6//!
7//! - **K1** [`k1_ssd_chunk_cumsum`] — per-chunk cumulative `Δ·A` decays.
8//! - **K2** [`k2_ssd_bmm`] — the intra-chunk `C·Bᵀ` block matmul.
9//! - **K3** [`k3_ssd_chunk_state`] — each chunk's contribution to its end state
10//!   (assuming a zero state at the chunk's start).
11//! - **K4** `k4_ssd_state_passing` — the serial inter-chunk scan that carries the
12//!   running state across chunk boundaries.
13//! - **K5** [`k5_ssd_chunk_scan`] — combines the intra-chunk (attention-like) and
14//!   inter-chunk (state-carried) contributions into the output `y`.
15//!
16//! This produces identical values and gradients to [`super::minimal`]; the
17//! serial form keeps per-chunk tensors small (lower peak memory) and is the
18//! basis of the recompute backward in [`super::serial_recalculated`].  Gradients
19//! here still flow through plain autodiff.
20
21#![allow(unused_variables)]
22
23use crate::mamba2::prelude::*;
24use crate::modules::sanity as san;
25use burn::prelude::*;
26
27impl Mamba2SsdInput {
28    /// Forward pass for the Mamba-2 SSD module (serial-over-chunks form).
29    ///
30    /// Returns:
31    /// - `y_bnlhp`.
32    /// - `final_state_bhpr`.
33    #[allow(non_snake_case)]
34    pub fn ssd_serial(self) -> (Tensor<5>, Tensor<4>) {
35        let input = self;
36        let [batch, nchunks, chunk_len, nheads, per_head_dim] = input.x_bnlhp.dims();
37        let [.., state_rank] = input.b_bnlhr.dims();
38        assert!(nchunks > 0, "sequence length must be at least 1");
39
40        san(&input.x_bnlhp);
41        san(&input.dt_bnlh);
42        san(&input.a_decay_h);
43        san(&input.b_bnlhr);
44        san(&input.c_bnlhr);
45        san(&input.d_h);
46        san(&input.initial_state_bhpr);
47
48        assert!(
49            input.init_state_hpr.is_none(),
50            "init_state_hpr not yet implemented"
51        );
52
53        // ── Permutes ──────────────────────────────────────────────────────────────────
54        // Note: dt_bnlh calculation (originally in Kernel 1) moved to Step 4 (before padding).
55        let dt_discretized_bhnl = input.dt_bnlh.permute([0, 3, 1, 2]);
56        assert_eq!(
57            [batch, nheads, nchunks, chunk_len],
58            dt_discretized_bhnl.dims()
59        );
60        san(&dt_discretized_bhnl);
61
62        // ── Kernel 1 ──────────────────────────────────────────────────────────────────
63        // IO: (..) -> (da_cumsum_bhnl [used in K3+K5][*], da_chunk_end_bhn [used in K4][omitted][*])
64        let (da_cumsum_bhnl, da_chunk_end_bhn): (Tensor<4>, Tensor<3>) =
65            k1_ssd_chunk_cumsum(dt_discretized_bhnl.clone(), input.a_decay_h.clone());
66        assert_eq!([batch, nheads, nchunks, chunk_len], da_cumsum_bhnl.dims());
67        assert_eq!([batch, nheads, nchunks], da_chunk_end_bhn.dims());
68        san(&da_cumsum_bhnl);
69        san(&da_chunk_end_bhn);
70
71        // ── Kernel 2 ──────────────────────────────────────────────────────────────────
72        // IO: (..) -> (cb_bnhll [used in K5][!])
73        let cb_bnhll: Tensor<5> = k2_ssd_bmm(input.c_bnlhr.clone(), input.b_bnlhr.clone());
74        assert_eq!(
75            [batch, nchunks, nheads, chunk_len, chunk_len],
76            cb_bnhll.dims()
77        );
78        san(&cb_bnhll);
79
80        // ── Kernel 3 ──────────────────────────────────────────────────────────────────
81        // IO: (..) -> (intra_chunk_state_bnhpr [used in K4][!])
82        let intra_chunk_state_bnhpr: Tensor<5> = k3_ssd_chunk_state(
83            input.x_bnlhp.clone(),
84            input.b_bnlhr.clone(),
85            da_cumsum_bhnl.clone(),
86            dt_discretized_bhnl.clone(),
87        );
88        assert_eq!(
89            [batch, nchunks, nheads, per_head_dim, state_rank],
90            intra_chunk_state_bnhpr.dims()
91        );
92        san(&intra_chunk_state_bnhpr);
93
94        // ── Kernel 4 ──────────────────────────────────────────────────────────────────
95        // IO: (..) -> (chunk_input_state_bnhpr [used in K5][!], final_state_bhpr [final output])
96        let (chunk_input_state_bnhpr, final_state_bhpr): (Tensor<5>, Tensor<4>) =
97            k4_ssd_state_passing(
98                intra_chunk_state_bnhpr.clone(),
99                da_chunk_end_bhn.clone(),
100                input.initial_state_bhpr,
101            );
102        assert_eq!(
103            [batch, nchunks, nheads, per_head_dim, state_rank],
104            chunk_input_state_bnhpr.dims()
105        );
106        assert_eq!(
107            [batch, nheads, per_head_dim, state_rank],
108            final_state_bhpr.dims()
109        );
110        san(&chunk_input_state_bnhpr);
111        san(&final_state_bhpr);
112
113        // ── Kernel 5 ──────────────────────────────────────────────────────────────────
114        let y_bnlhp: Tensor<5> = k5_ssd_chunk_scan(
115            da_cumsum_bhnl,
116            dt_discretized_bhnl,
117            input.x_bnlhp,
118            input.c_bnlhr,
119            cb_bnhll,
120            chunk_input_state_bnhpr,
121            input.d_h,
122        );
123        assert_eq!(
124            [batch, nchunks, chunk_len, nheads, per_head_dim],
125            y_bnlhp.dims()
126        );
127        san(&y_bnlhp);
128
129        (y_bnlhp, final_state_bhpr)
130    }
131}
132
133/// Based on the Kernel 1 Triton reference `_chunk_cumsum_fwd_kernel` (`ssd_chunk_state.py`).
134///
135/// Returns:
136/// - da_cumsum_bhnl `[used in K3+K5][*]` - intra-chunk cumsum.
137/// - da_chunk_end_bhn `[used in K4][omitted][*]` - last da_cumsum per chunk.
138pub fn k1_ssd_chunk_cumsum(
139    dt_discretized_bhnl: Tensor<4>,
140    a_decay_h: Tensor<1>,
141) -> (Tensor<4>, Tensor<3>) {
142    let [batch, nheads, nchunks, chunk_len] = dt_discretized_bhnl.dims();
143    let da_cumsum_bhnl: Tensor<4> = {
144        let a_decay_bhnl = a_decay_h
145            // - 1/6: unsqueeze-dims: (a_decay_h [*]) -> (a_decay_1h11)
146            .unsqueeze_dims::<4>(&[0, 2, 3]) // a_decay_1h11
147            // - 2: expand: (a_decay_1h11) -> (a_decay_bhnl)
148            .expand([batch, nheads, nchunks, chunk_len]);
149        // - 3: mul: (dt_discretized_bhnl [*], a_decay_bhnl) -> (da_bhnl)
150        // - 4: cumsum: (da_bhnl) -> (da_cumsum_bhnl [out][*])
151        (dt_discretized_bhnl * a_decay_bhnl).cumsum(3)
152    };
153    assert_eq!([batch, nheads, nchunks, chunk_len], da_cumsum_bhnl.dims());
154
155    let da_chunk_end_bhn = da_cumsum_bhnl
156        .clone()
157        // - 5: slice: (da_cumsum_bhnl [*]) -> (da_cumsum_bhn1)
158        .slice(s![.., .., .., -1]) // da_cumsum_bhn1
159        // - 6/6: squeeze: (da_cumsum_bhn1) -> (da_chunk_end_bhn [out])
160        .squeeze_dim::<3>(3);
161    assert_eq!([batch, nheads, nchunks], da_chunk_end_bhn.dims());
162
163    (da_cumsum_bhnl, da_chunk_end_bhn)
164}
165
166/// Based on the Kernel 2 Triton reference `_bmm_chunk_fwd_kernel` (`ssd_bmm.py`).
167///
168/// Returns:
169/// - cb_bnhll `[used in K5][!]`.
170pub fn k2_ssd_bmm(c_bnlhr: Tensor<5>, b_bnlhr: Tensor<5>) -> Tensor<5> {
171    let [batch, nchunks, chunk_len, nheads, _state_rank] = c_bnlhr.dims();
172
173    // - 1/3: permute: (c_bnlhr [in][*]) -> (c_bnhlr)
174    let c_bnhlr = c_bnlhr.permute([0, 1, 3, 2, 4]);
175    // - 2: permute: (b_bnlhr [in][*]) -> (b_bnhrl)
176    let b_bnhrl = b_bnlhr.permute([0, 1, 3, 4, 2]);
177    // - 3/3: matmul: (c_bnhlr, b_bnhrl) -> (cb_bnhll [out][!])
178    let cb_bnhll: Tensor<5> = c_bnhlr.matmul(b_bnhrl);
179    assert_eq!(
180        [batch, nchunks, nheads, chunk_len, chunk_len],
181        cb_bnhll.dims()
182    );
183    // Note: cb_bnhll is then only used by Kernel 5.
184    cb_bnhll
185}
186
187/// Based on the Kernel 3 Triton reference `_chunk_state_fwd_kernel` (`ssd_chunk_state.py`).
188///
189/// Returns:
190/// - cb_bngll `[used in K5][!]` - state assuming zero initial state at each chunk boundary.
191/// - b_bar_scale_bhnl `[*]` - intermediary
192pub fn k3_ssd_chunk_state(
193    x_bnlhp: Tensor<5>,
194    b_bnlhr: Tensor<5>,
195    da_cumsum_bhnl: Tensor<4>,
196    dt_discretized_bhnl: Tensor<4>,
197) -> Tensor<5> {
198    use burn::tensor::s;
199
200    let [batch, nchunks, chunk_len, nheads, per_head_dim] = x_bnlhp.dims();
201    let [.., state_rank] = b_bnlhr.dims();
202
203    // permute b and x to prepare them for the matmul
204    // - 1/15: permute: (x_bnlhp [in][*]) -> (x_bnhpl)
205    let x_bnhpl = x_bnlhp.clone().permute([0, 1, 3, 4, 2]);
206    assert_eq!(
207        [batch, nchunks, nheads, per_head_dim, chunk_len],
208        x_bnhpl.dims()
209    );
210    // - 2: permute: (b_bnlhr [in][*]) -> (b_bnhlr)
211    let b_bnhlr = b_bnlhr.permute([0, 1, 3, 2, 4]);
212    assert_eq!(
213        [batch, nchunks, nheads, chunk_len, state_rank],
214        b_bnhlr.dims()
215    );
216
217    // scale b
218    let b_scaled_bnhlr = {
219        let b_bar_scale_bhnl = {
220            let da_cumsum_last_in_chunk_bhn1 =
221                // - 6: slice: (da_cumsum_bhnl [in][*]) -> (da_cumsum_last_in_chunk_bhn1)
222                da_cumsum_bhnl.clone().slice(s![.., .., .., -1]);
223            assert_eq!(
224                [batch, nheads, nchunks, 1],
225                da_cumsum_last_in_chunk_bhn1.dims()
226            );
227
228            // - 7: expand: (da_cumsum_last_in_chunk_bhn1) -> (da_cumsum_last_bhnl)
229            let da_cumsum_last_bhnl =
230                da_cumsum_last_in_chunk_bhn1.expand([batch, nheads, nchunks, chunk_len]);
231            // - 8: sub: (da_cumsum_last_bhnl, da_cumsum_bhnl [from K1][*]) -> (da_delta_bhnl)
232            let da_delta_bhnl = da_cumsum_last_bhnl - da_cumsum_bhnl.clone();
233            // - 9: exp: (da_delta_bhnl) -> (forward_decay_to_chunk_end_bhnl [+])
234            let forward_decay_to_chunk_end_bhnl = da_delta_bhnl.exp();
235            assert_eq!(
236                [batch, nheads, nchunks, chunk_len],
237                forward_decay_to_chunk_end_bhnl.dims()
238            );
239
240            // - 10: mul: (forward_decay_to_chunk_end_bhnl [+], dt_discretized_bhnl [in][*]) -> (b_bar_scale_bhnl [+])
241            forward_decay_to_chunk_end_bhnl * dt_discretized_bhnl.clone()
242        };
243        assert_eq!([batch, nheads, nchunks, chunk_len], b_bar_scale_bhnl.dims());
244
245        // - 11: permute: (b_bar_scale_bhnl [+]) -> (b_bar_scale_bnhl)
246        let b_bar_scale_bnhl = b_bar_scale_bhnl.permute([0, 2, 1, 3]);
247        assert_eq!([batch, nchunks, nheads, chunk_len], b_bar_scale_bnhl.dims());
248        let b_bar_scale_bnhlr = b_bar_scale_bnhl
249            // - 12: unsqueeze: (b_bar_scale_bnhl) -> (b_bar_scale_bnhl1)
250            .unsqueeze_dim::<5>(4) // b_bar_scale_bnhl1
251            // - 13: expand: (b_bar_scale_bnhl1) -> (b_bar_scale_bnhlr)
252            .expand([batch, nchunks, nheads, chunk_len, state_rank]);
253        // - 14: mul: (b_bnhlr, b_bar_scale_bnhlr) -> (b_scaled_bnhlr [+])
254        b_bnhlr * b_bar_scale_bnhlr
255    };
256    assert_eq!(
257        [batch, nchunks, nheads, chunk_len, state_rank],
258        b_scaled_bnhlr.dims()
259    );
260
261    // - 15/15: matmul: (x_bnhpl, b_scaled_bnhlr [+]) -> (intra_chunk_state_bnhpr [out][!])
262    let intra_chunk_state_bnhpr: Tensor<5> = x_bnhpl.matmul(b_scaled_bnhlr);
263    assert_eq!(
264        [batch, nchunks, nheads, per_head_dim, state_rank],
265        intra_chunk_state_bnhpr.dims()
266    );
267    intra_chunk_state_bnhpr
268}
269
270/// Based on the Kernel 4 Triton reference `_state_passing_fwd_kernel` (`ssd_state_passing.py`).
271///
272/// Returns:
273/// - chunk_input_state_bnhpr `[used in K5][!]`.
274/// - final_state_bhpr `[final output]`.
275pub fn k4_ssd_state_passing(
276    intra_chunk_state_bnhpr: Tensor<5>,
277    da_chunk_end_bhn: Tensor<3>,
278    initial_state_bhpr: Tensor<4>,
279) -> (Tensor<5>, Tensor<4>) {
280    let [batch, nchunks, nheads, per_head_dim, state_rank] = intra_chunk_state_bnhpr.dims();
281    let flat_state_dim = per_head_dim * state_rank;
282
283    // - 1/5: init-mut: (initial_state_bhpr [in][*]) -> (running_state_bhpr)
284    let mut running_state_bhpr = initial_state_bhpr;
285    assert_eq!(
286        [batch, nheads, per_head_dim, state_rank],
287        running_state_bhpr.dims()
288    );
289
290    let mut chunk_input_state_vec_bhpr = Vec::with_capacity(nchunks + 1);
291    // - 2: vec-push: (running_state_bhpr [elem]) -> (chunk_input_state_vec_bhpr [vec][!])
292    chunk_input_state_vec_bhpr.push(running_state_bhpr.clone());
293
294    // - 3: serial-loop: (0..nchunks)
295    for i_chunk in 0..nchunks {
296        let intra_state_bhpr = intra_chunk_state_bnhpr
297            .clone()
298            //   - 3.1/3.9: slice: (intra_chunk_state_bnhpr [in][!]) -> (intra_chunk_state_b1hpr)
299            .slice(s![.., i_chunk, .., .., ..]) // intra_chunk_state_b1hpr
300            //   - 3.2: squeeze: (intra_chunk_state_b1hpr) -> (intra_state_bhpr)
301            .squeeze_dim::<4>(1);
302        assert_eq!(
303            [batch, nheads, per_head_dim, state_rank],
304            intra_state_bhpr.dims()
305        );
306
307        let decay_bhpr = da_chunk_end_bhn
308            .clone()
309            //   - 3.3: slice: (da_chunk_end_bhn [in][*]) -> (da_chunk_end_bh1)
310            .slice(s![.., .., i_chunk]) // da_chunk_end_bh1
311            //   - 3.4: exp: (da_chunk_end_bh1) -> (exp_da_chunk_end_bh1)
312            .exp() // exp_da_chunk_end_bh1
313            //   - 3.5: unsqueeze: (exp_da_chunk_end_bh1) -> (exp_da_chunk_end_bh11)
314            .unsqueeze_dim::<4>(3) // exp_da_chunk_end_bh11
315            //   - 3.6: expand: (exp_da_chunk_end_bh11) -> (decay_bhpr)
316            .expand([batch, nheads, per_head_dim, state_rank]);
317
318        // SSM recurrence: running_state = decay * running_state + intra_state
319        running_state_bhpr =
320        //   - 3.7: mul: (decay_bhpr, running_state_bhpr) -> (running_state_bhpr)
321            (decay_bhpr * running_state_bhpr) // running_state_bhpr
322        //   - 3.8: add: (running_state_bhpr, intra_state_bhpr) -> (running_state_bhpr)
323            + intra_state_bhpr;
324        //   - 3.9/3.9: vec-push: (running_state_bhpr [elem]) -> (chunk_input_state_vec_bhpr [vec][!])
325        chunk_input_state_vec_bhpr.push(running_state_bhpr.clone());
326    }
327
328    // - 4: vec-pop: (chunk_input_state_vec_bhpr [vec][!]) -> (final_state_bhpr [elem][out][!])
329    let final_state_bhpr = chunk_input_state_vec_bhpr.pop().unwrap();
330    assert_eq!(
331        [batch, nheads, per_head_dim, state_rank],
332        final_state_bhpr.dims()
333    );
334
335    // - 5/5: stack: (chunk_input_state_vec_bhpr [!]) -> (chunk_input_state_bnhpr [out][!])
336    let chunk_input_state_bnhpr = Tensor::stack(chunk_input_state_vec_bhpr, 1);
337    assert_eq!(
338        [batch, nchunks, nheads, per_head_dim, state_rank],
339        chunk_input_state_bnhpr.dims()
340    );
341
342    (chunk_input_state_bnhpr, final_state_bhpr)
343}
344
345/// Based on the Kernel 5 Triton reference `_chunk_scan_fwd_kernel` (`ssd_chunk_scan.py`).
346///
347/// Returns:
348/// - y_bnlhp `[final output]`
349pub fn k5_ssd_chunk_scan(
350    da_cumsum_bhnl: Tensor<4>,
351    dt_discretized_bhnl: Tensor<4>,
352    x_bnlhp: Tensor<5>,
353    c_bnlhr: Tensor<5>,
354    cb_bnhll: Tensor<5>,
355    chunk_input_state_bnhpr: Tensor<5>,
356    d_h: Tensor<1>,
357) -> Tensor<5> {
358    let [batch, nchunks, chunk_len, nheads, per_head_dim] = x_bnlhp.dims();
359    let device = x_bnlhp.device();
360
361    // Rearrange inputs to the common [batch, nchunks, nheads, ...] ordering used below.
362    // - 1/36: permute: (da_cumsum_bhnl [*]) -> (da_cumsum_bnhl)
363    let da_cumsum_bnhl = da_cumsum_bhnl.permute([0, 2, 1, 3]);
364    san(&da_cumsum_bnhl);
365    // - 2: permute: (dt_discretized_bhnl [*]) -> (dt_bnhl)
366    let dt_bnhl = dt_discretized_bhnl.permute([0, 2, 1, 3]);
367    san(&dt_bnhl);
368    // - 3: permute: (x_bnlhp [*]) -> (x_bnhlp)
369    let x_bnhlp = x_bnlhp.clone().permute([0, 1, 3, 2, 4]);
370    san(&x_bnhlp);
371
372    // B/C are already per-head — only a permute is needed.
373    let c_bnhlr = c_bnlhr.permute([0, 1, 3, 2, 4]);
374    san(&c_bnhlr);
375    san(&cb_bnhll);
376
377    // ── BLUE: exp(dA[l]) · C[l,:] @ state_in^T ─────────────────────────────
378    //
379    //   blue[b,n,h,l,p] = exp(da[b,n,h,l]) · Σ_r  c[b,n,h,l,r] · state[b,n,h,p,r]
380    //
381    //   [b,n,h,l,r] @ [b,n,h,r,p]  →  [b,n,h,l,p]
382    let exp_da_cumsum_bnhlp = da_cumsum_bnhl
383        .clone()
384        // - 11: exp: (da_cumsum_bnhl) -> (exp_da_cumsum_bnhl)
385        .exp()
386        // - 12: unsqueeze: (exp_da_cumsum_bnhl) -> (exp_da_cumsum_bnhl1)
387        .unsqueeze_dim::<5>(4) // exp_da_cumsum_bnhl1
388        // - 13: expand: (exp_da_cumsum_bnhl1) -> (exp_da_cumsum_bnhlp)
389        .expand([batch, nchunks, nheads, chunk_len, per_head_dim]);
390    san(&exp_da_cumsum_bnhlp);
391    // - 14: permute: (chunk_input_state_bnhpr [!]) -> (chunk_input_state_bnhrp)
392    let chunk_input_state_bnhrp = chunk_input_state_bnhpr.permute([0, 1, 2, 4, 3]);
393    // - 15: matmul: (c_bnhlr, chunk_input_state_bnhrp) -> (blue_bnhlp)
394    let blue_scaled_bnhlp = c_bnhlr
395        .matmul(chunk_input_state_bnhrp)  // blue_bnhlp
396        // - 16: mul: (blue_bnhlp, exp_da_cumsum_bnhlp) -> (blue_scaled_bnhlp)
397        * exp_da_cumsum_bnhlp;
398    san(&blue_scaled_bnhlp);
399
400    // ── ORANGE: causal CB_weighted @ X ──────────────────────────────────────
401    //
402    //   orange[b,n,h,l,p] = Σ_{s≤l} CB[l,s] · exp(da[l]-da[s]) · dt[s] · x[s,p]
403    //
404    // Precompute the full lower-triangular weight matrix, then do a single matmul.
405    //
406    let da_cumsum_target_bnhll = da_cumsum_bnhl
407        .clone()
408        // - 17: unsqueeze: (da_cumsum_bnhl) -> (da_cumsum_bnhl1)
409        .unsqueeze_dim::<5>(4) // da_cumsum_bnhl1
410        // - 18: expand: (da_cumsum_bnhl1) -> (da_cumsum_target_bnhll)
411        .expand([batch, nchunks, nheads, chunk_len, chunk_len]);
412    // println!("{}", da_cumsum_target_bnhll);
413    san(&da_cumsum_target_bnhll);
414    let da_cumsum_source_bnhll = da_cumsum_bnhl
415        // - 19: unsqueeze: (da_cumsum_bnhl) -> (da_cumsum_bnh1l)
416        .unsqueeze_dim::<5>(3) // da_cumsum_bnh1l
417        // - 20: expand: (da_cumsum_bnh1l) -> (da_cumsum_source_bnhll)
418        .expand([batch, nchunks, nheads, chunk_len, chunk_len]);
419    // println!("{}", da_cumsum_source_bnhll);
420    san(&da_cumsum_source_bnhll);
421    // - 21: sub: (da_cumsum_target_bnhll, da_cumsum_source_bnhll) -> (da_cumsum_diff_bnhll)
422    let da_cumsum_diff_bnhll = da_cumsum_target_bnhll - da_cumsum_source_bnhll;
423    san(&da_cumsum_diff_bnhll);
424
425    // note: overflow instability at step 22, a `minimal::segsum`-like upper triangle protection is necessary.
426    // - 21.1: tril-mask: (0) -> (causal_mask_ll), expanded as a view to causal_mask_bnhll.
427    // true above the main diagonal, false at diagonal and below.
428    // Built at [L,L] and broadcast — the mask values do not depend on (b,n,h).
429    let causal_mask_bnhll: Tensor<5, burn::prelude::Bool> =
430        Tensor::<2, burn::prelude::Bool>::tril_mask([chunk_len, chunk_len], 0, &device)
431            .reshape([1, 1, 1, chunk_len, chunk_len])
432            .expand([batch, nchunks, nheads, chunk_len, chunk_len]);
433    // - 21.2: mask-fill: (da_cumsum_diff_bnhll, causal_mask_bnhll) -> (da_cumsum_diff_masked_bnhll)
434    // Causal mask and exp stabilizer: above upper diagonal set to -inf.
435    let da_cumsum_diff_masked_bnhll =
436        da_cumsum_diff_bnhll.mask_fill(causal_mask_bnhll, f32::NEG_INFINITY);
437
438    // - 22: exp: (da_cumsum_diff_masked_bnhll) -> (da_cumsum_diff_exp_bnhll)
439    let da_cumsum_diff_exp_bnhll = da_cumsum_diff_masked_bnhll.exp();
440    san(&da_cumsum_diff_exp_bnhll);
441    let dt_source_bnhll = dt_bnhl
442        // - 23: unsqueeze: (dt_bnhl) -> (dt_bnh1l)
443        .unsqueeze_dim::<5>(3) // dt_bnh1l
444        // - 24: expand: (dt_bnh1l) -> (dt_source_bnhll)
445        .expand([batch, nchunks, nheads, chunk_len, chunk_len]);
446    san(&dt_source_bnhll);
447
448    // note: steps 25, 26 and 29 are no longer necessary.
449    // // Causal mask (0 above the main diagonal, 1 elsewhere).
450    // let causal_mask_bnhll =
451    //     // - 25: ones: (1) -> (ones_bnhll)
452    //     Tensor::ones([batch, nchunks, nheads, chunk_len, chunk_len], &device)
453    //     // - 26: tril: (ones_bnhll, 0) -> (causal_mask_bnhll)
454    //     .tril(0);
455
456    //   [b,n,h,l,l] @ [b,n,h,l,p]  →  [b,n,h,l,p]
457    // - 27: mul: (cb_bnhll, da_cumsum_diff_exp_bnhll) -> (orange_lhs_partial1_bnhll)
458    let orange_lhs_partial1_bnhll = cb_bnhll * da_cumsum_diff_exp_bnhll;
459    san(&orange_lhs_partial1_bnhll);
460    // - 28: mul: (orange_lhs_partial1_bnhll, dt_source_bnhll) -> (orange_lhs_partial2_bnhll)
461    let orange_lhs_partial2_bnhll = orange_lhs_partial1_bnhll * dt_source_bnhll;
462    san(&orange_lhs_partial2_bnhll);
463    // // - 29: mul: (orange_lhs_partial2_bnhll, causal_mask_bnhll) -> (orange_lhs_partial3_bnhll)
464    // let orange_lhs_partial3_bnhll = orange_lhs_partial2_bnhll * causal_mask_bnhll;
465    // san(&orange_lhs_partial3_bnhll);
466    // - 30: matmul: (orange_lhs_partial3_bnhll, x_bnhlp) -> (orange_bnhlp)
467    // - 30: matmul: (orange_lhs_partial2_bnhll, x_bnhlp) -> (orange_bnhlp)
468    let orange_bnhlp = orange_lhs_partial2_bnhll.matmul(x_bnhlp);
469    san(&orange_bnhlp);
470
471    // ── SKIP: D[h] · x[l,p] ─────────────────────────────────────────────────
472    //
473    //   D_HAS_HDIM = False: D is a scalar per head, shape [nheads].
474    //   Triton: `acc += x_residual * D`
475    let skip_bnlhp = d_h
476        // - 31: unsqueeze-dims: (d_h [*]) -> (d_111h1)
477        .unsqueeze_dims::<5>(&[0, 1, 2, 4]) // d_111h1
478        // - 32: expand: (d_111h1) -> (d_bnlhp)
479        .expand([
480            batch,
481            nchunks,
482            chunk_len,
483            nheads,
484            per_head_dim,
485        ]) // d_bnlhp
486    // - 33: mul: (d_bnlhp, x_bnlhp[*]) -> (skip_bnlhp)
487    * x_bnlhp;
488    san(&skip_bnlhp);
489
490    // Permute BLUE + ORANGE from [b,n,h,l,p] back to [b,n,l,h,p], then add SKIP.
491    // - 34: add: (blue_scaled_bnhlp, orange_bnhlp) -> (y_partial_bnhlp)
492    let y_partial_bnhlp = blue_scaled_bnhlp + orange_bnhlp;
493    san(&y_partial_bnhlp);
494    // - 35: permute: (y_partial_bnhlp) -> (y_partial_bnlhp)
495    let y_partial_bnlhp = y_partial_bnhlp.permute([0, 1, 3, 2, 4]);
496    san(&y_partial_bnlhp);
497    // - 36/36: add: (y_partial_bnlhp, skip_bnlhp) -> (y_bnlhp [out])
498    let y_bnlhp: Tensor<5> = y_partial_bnlhp + skip_bnlhp;
499    san(&y_bnlhp);
500
501    assert_eq!(
502        [batch, nchunks, chunk_len, nheads, per_head_dim],
503        y_bnlhp.dims()
504    );
505    y_bnlhp
506}