burn_mamba/mamba3/single_ssd/ssd/serial_recalculated/
diag.rs1#![allow(non_snake_case)]
19
20use burn_stack::utils::fprim::F;
21use burn::backend::Backend;
22
23pub struct DiagGrads<B: Backend> {
27 pub d_v_bnlmhp: F<B, 6>,
29 pub d_c_bnlmhr: F<B, 6>,
31 pub d_b_bnlmhr: F<B, 6>,
33 pub d_gamma_bnlh: F<B, 4>,
35}
36
37pub fn y_diag_correction<B: Backend>(
42 v_bnlmhp: F<B, 6>,
43 b_bnlmhr: F<B, 6>,
44 c_bnlmhr: F<B, 6>,
45 gamma_bnlh: F<B, 4>,
46 siso_specialization: bool,
47) -> F<B, 6> {
48 let [.., mimo_rank, _nheads, _per_head_dim] = v_bnlmhp.dims();
49 if mimo_rank == 1 && siso_specialization {
50 y_diag_correction_siso(v_bnlmhp, b_bnlmhr, c_bnlmhr, gamma_bnlh)
51 } else {
52 y_diag_correction_mimo(v_bnlmhp, b_bnlmhr, c_bnlmhr, gamma_bnlh)
53 }
54}
55
56pub(crate) fn y_diag_correction_siso<B: Backend>(
59 v_bnlmhp: F<B, 6>,
60 b_bnlmhr: F<B, 6>,
61 c_bnlmhr: F<B, 6>,
62 gamma_bnlh: F<B, 4>,
63) -> F<B, 6> {
64 let qk_dot_bnlmh1 = (c_bnlmhr * b_bnlmhr).sum_dim(5);
65 let gamma_bnl1h1 = gamma_bnlh.unsqueeze_dims::<6>(&[3, 5]);
66 v_bnlmhp * qk_dot_bnlmh1 * gamma_bnl1h1
67}
68
69pub(crate) fn y_diag_correction_mimo<B: Backend>(
71 v_bnlmhp: F<B, 6>,
72 b_bnlmhr: F<B, 6>,
73 c_bnlmhr: F<B, 6>,
74 gamma_bnlh: F<B, 4>,
75) -> F<B, 6> {
76 let c_bnlhmr = c_bnlmhr.swap_dims(3, 4);
77 let b_bnlhrm = b_bnlmhr.permute([0, 1, 2, 4, 5, 3]);
78 let qk_dot_bnlhmM = c_bnlhmr.matmul(b_bnlhrm); let v_bnlhmp = v_bnlmhp.swap_dims(3, 4);
80 let y_d_bnlhmp = qk_dot_bnlhmM.matmul(v_bnlhmp); let gamma_bnlh11 = gamma_bnlh.unsqueeze_dims::<6>(&[4, 5]);
82 (y_d_bnlhmp * gamma_bnlh11).swap_dims(3, 4)
83}
84
85pub fn y_diag_correction_backward<B: Backend>(
90 d_y_bnlmhp: F<B, 6>,
91 v_bnlmhp: F<B, 6>,
92 b_bnlmhr: F<B, 6>,
93 c_bnlmhr: F<B, 6>,
94 gamma_bnlh: F<B, 4>,
95 siso_specialization: bool,
96) -> DiagGrads<B> {
97 let [.., mimo_rank, _nheads, _per_head_dim] = v_bnlmhp.dims();
98 if mimo_rank == 1 && siso_specialization {
99 y_diag_correction_backward_siso(d_y_bnlmhp, v_bnlmhp, b_bnlmhr, c_bnlmhr, gamma_bnlh)
100 } else {
101 y_diag_correction_backward_mimo(d_y_bnlmhp, v_bnlmhp, b_bnlmhr, c_bnlmhr, gamma_bnlh)
102 }
103}
104
105pub(crate) fn y_diag_correction_backward_siso<B: Backend>(
118 d_y_bnlmhp: F<B, 6>,
119 v_bnlmhp: F<B, 6>,
120 b_bnlmhr: F<B, 6>,
121 c_bnlmhr: F<B, 6>,
122 gamma_bnlh: F<B, 4>,
123) -> DiagGrads<B> {
124 let qk_dot_bnlmh1 = (c_bnlmhr.clone() * b_bnlmhr.clone()).sum_dim(5);
125 let dyv_bnlmh1 = (d_y_bnlmhp.clone() * v_bnlmhp).sum_dim(5);
126 let gamma_bnl1h1 = gamma_bnlh.unsqueeze_dims::<6>(&[3, 5]);
127
128 let d_gamma_bnlh: F<B, 4> = (qk_dot_bnlmh1.clone() * dyv_bnlmh1.clone())
130 .squeeze_dim::<5>(5) .squeeze_dim::<4>(3); let d_qk_dot_bnlmh1 = dyv_bnlmh1 * gamma_bnl1h1.clone();
134 let d_v_bnlmhp = d_y_bnlmhp * (qk_dot_bnlmh1 * gamma_bnl1h1);
135 let d_c_bnlmhr = b_bnlmhr * d_qk_dot_bnlmh1.clone();
136 let d_b_bnlmhr = c_bnlmhr * d_qk_dot_bnlmh1;
137
138 DiagGrads {
139 d_v_bnlmhp,
140 d_c_bnlmhr,
141 d_b_bnlmhr,
142 d_gamma_bnlh,
143 }
144}
145
146pub(crate) fn y_diag_correction_backward_mimo<B: Backend>(
148 d_y_bnlmhp: F<B, 6>,
149 v_bnlmhp: F<B, 6>,
150 b_bnlmhr: F<B, 6>,
151 c_bnlmhr: F<B, 6>,
152 gamma_bnlh: F<B, 4>,
153) -> DiagGrads<B> {
154 let c_bnlhmr = c_bnlmhr.swap_dims(3, 4); let b_bnlhmr = b_bnlmhr.swap_dims(3, 4); let v_bnlhmp = v_bnlmhp.swap_dims(3, 4); let d_y_bnlhmp = d_y_bnlmhp.swap_dims(3, 4); let qk_dot_bnlhmM = c_bnlhmr.clone().matmul(b_bnlhmr.clone().transpose());
161 let y_d_unw_bnlhmp = qk_dot_bnlhmM.clone().matmul(v_bnlhmp.clone());
163
164 let d_gamma_bnlh: F<B, 4> = (d_y_bnlhmp.clone() * y_d_unw_bnlhmp)
166 .sum_dim(5) .squeeze_dim::<5>(5) .sum_dim(4) .squeeze_dim::<4>(4); let gamma_bnlh11 = gamma_bnlh.unsqueeze_dims::<6>(&[4, 5]);
173 let d_y_d_unw_bnlhmp = d_y_bnlhmp * gamma_bnlh11;
174
175 let d_qk_dot_bnlhmM = d_y_d_unw_bnlhmp.clone().matmul(v_bnlhmp.transpose()); let d_v_bnlhmp = qk_dot_bnlhmM
180 .transpose() .matmul(d_y_d_unw_bnlhmp); let d_c_bnlhmr = d_qk_dot_bnlhmM.clone().matmul(b_bnlhmr); let d_b_bnlhmr = d_qk_dot_bnlhmM
187 .transpose() .matmul(c_bnlhmr); DiagGrads {
191 d_v_bnlmhp: d_v_bnlhmp.swap_dims(3, 4),
192 d_c_bnlmhr: d_c_bnlhmr.swap_dims(3, 4),
193 d_b_bnlmhr: d_b_bnlhmr.swap_dims(3, 4),
194 d_gamma_bnlh,
195 }
196}
197
198#[cfg(all(test, feature = "_dev-test"))]
203mod tests;