use super::backward_ops::{Conv1dDims, backward_conv1d_step, backward_rms_norm};
use super::flat::{MambaBackboneFlat, MambaLayerFlat};
use super::scratch::BackwardPhaseScratch;
use super::weights::{TrainMambaLayerWeights, TrainMambaWeights};
use crate::ops::blas::sgemm_backward;
use crate::ops::dims::MambaDims;
use crate::ops::fast_math::fast_exp_scalar;
pub fn backward_mamba_layer_batched(
d_temporal_flat: &mut [f32],
d_layer: &mut TrainMambaLayerWeights,
acts: &MambaLayerFlat,
w: &TrainMambaLayerWeights,
a_neg: &[f32],
scratch: &mut BackwardPhaseScratch,
dims: &MambaDims,
) {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let dc = dims.d_conv;
let dt_rank = dims.dt_rank;
let xdbl_dim = dims.xdbl_dim;
let seq_len = dims.seq_len;
acts.copy_gated_all(&mut scratch.gated_buf);
sgemm_backward(
&mut scratch.d_gated_flat,
&mut d_layer.out_proj_w,
None,
d_temporal_flat,
&scratch.gated_buf,
&w.out_proj_w,
(seq_len, di, dm),
);
for t in 0..seq_len {
let t_off = t * di;
let gate_post_silu = acts.gate_post_silu(t);
let gate_pre_silu = acts.gate_pre_silu(t);
let y_vals = acts.y(t);
for d in 0..di {
let dg = scratch.d_gated_flat[t_off + d];
scratch.d_y_flat[t_off + d] = dg * gate_post_silu[d];
let x = gate_pre_silu[d];
let sigma = 1.0 / (1.0 + fast_exp_scalar(-x));
let silu_grad = sigma * (1.0 + x * (1.0 - sigma));
scratch.d_gate_flat[t_off + d] = dg * y_vals[d] * silu_grad;
}
}
scratch.d_h.fill(0.0);
scratch.d_conv_carry.fill(0.0);
let b_offset = dt_rank;
let c_offset = dt_rank + ds;
for t in (0..seq_len).rev() {
let t_off_di = t * di;
let t_off_xdbl = t * xdbl_dim;
scratch.d_delta_flat[t_off_di..t_off_di + di].fill(0.0);
scratch.d_u_flat[t_off_di..t_off_di + di].fill(0.0);
scratch.d_xdbl_flat[t_off_xdbl..t_off_xdbl + xdbl_dim].fill(0.0);
{
let acts_delta = acts.delta(t);
let acts_u = acts.u(t);
let acts_da_exp = acts.da_exp(t);
let acts_h_prev = acts.h_prev(t);
let acts_h_curr = acts.h_curr(t);
let acts_xdbl = acts.xdbl(t);
for d in 0..di {
let delta_d = acts_delta[d];
let u_d = acts_u[d];
let dy_d = scratch.d_y_flat[t_off_di + d];
d_layer.d_param[d] += dy_d * u_d;
scratch.d_u_flat[t_off_di + d] += dy_d * w.d_param[d];
for n in 0..ds {
let idx = d * ds + n;
let a_dn = a_neg[idx];
let da = acts_da_exp[idx];
let h_prev = acts_h_prev[idx];
let b_n = acts_xdbl[b_offset + n];
let c_n = acts_xdbl[c_offset + n];
let h_curr = acts_h_curr[idx];
scratch.d_h[idx] += dy_d * c_n;
let dh = scratch.d_h[idx];
scratch.d_delta_flat[t_off_di + d] += dh * (a_dn * da * h_prev + u_d * b_n);
scratch.d_u_flat[t_off_di + d] += dh * delta_d * b_n;
scratch.d_xdbl_flat[t_off_xdbl + b_offset + n] += dh * delta_d * u_d;
scratch.d_xdbl_flat[t_off_xdbl + c_offset + n] += dy_d * h_curr;
d_layer.a_log[idx] += dh * da * delta_d * a_dn * h_prev;
scratch.d_h[idx] = da * dh;
}
}
}
{
let acts_delta_raw = acts.delta_raw(t);
for (d, &raw) in acts_delta_raw.iter().enumerate().take(di) {
scratch.d_delta_raw_flat[t_off_di + d] = if raw > 20.0 {
scratch.d_delta_flat[t_off_di + d]
} else {
let sig = 1.0 / (1.0 + fast_exp_scalar(-raw));
scratch.d_delta_flat[t_off_di + d] * sig
};
}
}
}
acts.copy_xdbl_dt_all(&mut scratch.xdbl_dt_buf);
sgemm_backward(
&mut scratch.d_dt_input_flat,
&mut d_layer.dt_proj_w,
Some(&mut d_layer.dt_proj_b),
&scratch.d_delta_raw_flat,
&scratch.xdbl_dt_buf,
&w.dt_proj_w,
(seq_len, dt_rank, di),
);
for t in 0..seq_len {
let xdbl_off = t * xdbl_dim;
let dt_off = t * dt_rank;
for i in 0..dt_rank {
scratch.d_xdbl_flat[xdbl_off + i] += scratch.d_dt_input_flat[dt_off + i];
}
}
acts.copy_u_all(&mut scratch.u_buf);
sgemm_backward(
&mut scratch.d_u_xproj_flat,
&mut d_layer.x_proj_w,
None,
&scratch.d_xdbl_flat,
&scratch.u_buf,
&w.x_proj_w,
(seq_len, di, xdbl_dim),
);
for (du, &du_xp) in scratch
.d_u_flat
.iter_mut()
.zip(scratch.d_u_xproj_flat.iter())
{
*du += du_xp;
}
for t in 0..seq_len {
let t_off = t * di;
let post_conv = acts.post_conv(t);
for (d, &x) in post_conv.iter().enumerate().take(di) {
let sig = 1.0 / (1.0 + fast_exp_scalar(-x));
scratch.d_conv_out_flat[t_off + d] =
scratch.d_u_flat[t_off + d] * sig * (1.0 + x * (1.0 - sig));
}
}
for t in (0..seq_len).rev() {
let t_off = t * di;
backward_conv1d_step(
&mut scratch.d_x_branch_flat[t_off..t_off + di],
&mut d_layer.conv1d_weight,
&mut d_layer.conv1d_bias,
&scratch.d_conv_out_flat[t_off..t_off + di],
acts.conv_state(t),
&w.conv1d_weight,
Conv1dDims {
d_inner: di,
d_conv: dc,
},
);
if dc > 1 {
let carry_stride = dc - 1;
for d in 0..di {
let carry_base = d * carry_stride;
let w_base = d * dc;
scratch.d_x_branch_flat[t_off + d] += scratch.d_conv_carry[carry_base];
for k in 0..carry_stride - 1 {
scratch.d_conv_carry[carry_base + k] = scratch.d_conv_carry[carry_base + k + 1]
+ scratch.d_conv_out_flat[t_off + d] * w.conv1d_weight[w_base + dc - 2 - k];
}
scratch.d_conv_carry[carry_base + carry_stride - 1] =
scratch.d_conv_out_flat[t_off + d] * w.conv1d_weight[w_base];
}
}
}
for t in 0..seq_len {
let t_off_di = t * di;
let proj_off = t * 2 * di;
scratch.d_proj_flat[proj_off..proj_off + di]
.copy_from_slice(&scratch.d_x_branch_flat[t_off_di..t_off_di + di]);
scratch.d_proj_flat[proj_off + di..proj_off + 2 * di]
.copy_from_slice(&scratch.d_gate_flat[t_off_di..t_off_di + di]);
}
acts.copy_post_norm_all(&mut scratch.post_norm_buf);
sgemm_backward(
&mut scratch.d_norm_flat,
&mut d_layer.in_proj_w,
None,
&scratch.d_proj_flat,
&scratch.post_norm_buf,
&w.in_proj_w,
(seq_len, dm, 2 * di),
);
for t in 0..seq_len {
let off = t * dm;
let rms_slice = [acts.rms_val(t)];
scratch.d_pre_norm_flat[off..off + dm].fill(0.0);
backward_rms_norm(
&mut scratch.d_pre_norm_flat[off..off + dm],
&mut d_layer.norm_weight,
&scratch.d_norm_flat[off..off + dm],
acts.residual(t),
(&w.norm_weight, &rms_slice),
1,
dm,
);
for d in 0..dm {
d_temporal_flat[off + d] += scratch.d_pre_norm_flat[off + d];
}
}
}
pub fn backward_mamba_backbone_batched(
d_temporal_flat: &mut [f32],
d_mamba: &mut TrainMambaWeights,
acts: &MambaBackboneFlat,
mamba_w: &TrainMambaWeights,
a_neg_all: &[f32],
scratch: &mut BackwardPhaseScratch,
dims: &MambaDims,
) {
let dm = dims.d_model;
let di = dims.d_inner;
let ds = dims.d_state;
let mid = dims.mamba_input_dim;
let seq_len = dims.seq_len;
let n_layers = dims.n_layers;
let a_neg_per_layer = di * ds;
{
let norm_f_dx = &mut scratch.d_input_proj_scratch[..seq_len * dm];
norm_f_dx.fill(0.0);
scratch.d_norm_f_weight_local.fill(0.0);
backward_rms_norm(
norm_f_dx,
&mut scratch.d_norm_f_weight_local,
&d_temporal_flat[..seq_len * dm],
&acts.norm_f_input[..seq_len * dm],
(&mamba_w.norm_f_weight, &acts.norm_f_rms[..seq_len]),
seq_len,
dm,
);
d_temporal_flat[..seq_len * dm].copy_from_slice(&norm_f_dx[..seq_len * dm]);
for (a, b) in d_mamba
.norm_f_weight
.iter_mut()
.zip(&scratch.d_norm_f_weight_local)
{
*a += b;
}
}
for layer_idx in (0..n_layers).rev() {
let a_neg_start = layer_idx * a_neg_per_layer;
backward_mamba_layer_batched(
d_temporal_flat,
&mut d_mamba.layers[layer_idx],
&acts.layers[layer_idx],
&mamba_w.layers[layer_idx],
&a_neg_all[a_neg_start..a_neg_start + a_neg_per_layer],
scratch,
dims,
);
}
sgemm_backward(
&mut scratch.d_input_proj_scratch,
&mut d_mamba.input_proj_w,
Some(&mut d_mamba.input_proj_b),
&d_temporal_flat[..seq_len * dm],
&acts.input_proj_inputs[..seq_len * mid],
&mamba_w.input_proj_w,
(seq_len, mid, dm),
);
}