1#![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#[non_exhaustive]
34pub struct CombinedSingleSsdGrads<B: Backend> {
35 pub d_v_bnlmhp: F<B, 6>,
37 pub d_da_bnlh: F<B, 4>,
39 pub d_b_bnlmhr: F<B, 6>,
41 pub d_c_bnlmhr: F<B, 6>,
43 pub d_gamma_bnlh: F<B, 4>,
45 pub d_scale_bnlh: F<B, 4>,
47 pub d_initial_state_bhpr: F<B, 4>,
49}
50
51#[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 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 let (da_cumsum_bhnl, da_chunk_end_bhn) = k1_ssd_chunk_cumsum(da_bnlh.clone());
108 san(&da_cumsum_bhnl);
109
110 let cb_bnhLMLM = k2_ssd_bmm(c_bnlmhr.clone(), b_bnlmhr.clone());
112 san(&cb_bnhLMLM);
113
114 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 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 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 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 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 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 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 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 let scale_bhLM: F<B, 3> = scale_bnlh
214 .clone()
215 .slice(s![.., i_chunk, .., ..]) .squeeze_dim::<3>(1)
217 .swap_dims(1, 2) .unsqueeze_dim::<4>(3) .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 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() .matmul(d_ch_bhLMp.clone()) .transpose(); 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 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 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) .expand([batch, nheads, chunk_len * mimo_rank, chunk_len * mimo_rank]);
297
298 let prod_bhLMLM = cb_bhLMLM.clone() * decay_strict_bhLMLM.clone();
300 let w_bhLMLM = prod_bhLMLM.clone() * scale_col_bhLMLM.clone();
301
302 let d_w_bhLMLM: F<B, 4> = d_y_bhLMp.clone().matmul(v_bhLMp.clone().transpose());
304 san(&d_w_bhLMLM);
305
306 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 let d_prod_bhLMLM = d_w_bhLMLM.clone() * scale_col_bhLMLM;
313 let d_scale_at_bhLMLM = d_w_bhLMLM * prod_bhLMLM;
314
315 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 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 let d_scale_lower_bhl: F<B, 3> = d_scale_at_bhLMLM
334 .sum_dim(2) .squeeze_dim::<3>(2) .reshape([batch, nheads, chunk_len, mimo_rank])
337 .sum_dim(3) .squeeze_dim::<3>(3); vec_lower_d_scale_bhl.push(d_scale_lower_bhl);
340
341 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 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 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 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()); let d_k_scaled_bnhLMr: F<B, 5> = decayed_v_bnhpLM
409 .transpose() .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 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 let d_k_scaled_bnlmhr: F<B, 6> = d_k_scaled_bnhLMr
447 .swap_dims(2, 3) .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) .squeeze_dim::<5>(5) .sum_dim(3) .squeeze_dim::<4>(3); 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 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 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 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 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}