pub struct Conv1dDims {
pub d_inner: usize,
pub d_conv: usize,
}
pub fn backward_conv1d_step(
d_x_branch: &mut [f32], d_conv_weight: &mut [f32], d_conv_bias: &mut [f32], d_conv_out: &[f32], conv_state: &[f32], conv_weight: &[f32], dims: Conv1dDims,
) {
let di = dims.d_inner;
let dc = dims.d_conv;
for d in 0..di {
let base = d * dc;
d_conv_bias[d] += d_conv_out[d];
for k in 0..dc {
d_conv_weight[base + k] += d_conv_out[d] * conv_state[base + k];
}
d_x_branch[d] = d_conv_out[d] * conv_weight[base + dc - 1];
}
}
pub fn backward_rms_norm(
dx: &mut [f32], d_scale: &mut [f32], dy: &[f32], x: &[f32], scale_and_rms: (&[f32], &[f32]), batch: usize,
dim: usize,
) {
let (scale, rms_vals) = scale_and_rms;
let dim_f = dim as f32;
for (b, &rms_b) in rms_vals.iter().enumerate().take(batch) {
let off = b * dim;
let inv_rms = 1.0 / rms_b;
let inv_rms2 = inv_rms * inv_rms;
let mut dot = 0.0_f32;
for i in 0..dim {
dot += dy[off + i] * scale[i] * x[off + i];
}
dot *= inv_rms2;
for i in 0..dim {
d_scale[i] += dy[off + i] * x[off + i] * inv_rms;
dx[off + i] = inv_rms * (scale[i] * dy[off + i] - x[off + i] * dot / dim_f);
}
}
}