pub(super) struct DepthwiseGeometry {
pub input: (usize, usize),
pub output: (usize, usize),
pub channels: usize,
pub depth_multiplier: usize,
pub kernel: (usize, usize),
pub strides: (usize, usize),
pub pad_before: (usize, usize),
}
impl DepthwiseGeometry {
pub fn out_channels(&self) -> usize {
self.channels * self.depth_multiplier
}
fn input_item(&self) -> usize {
self.input.0 * self.input.1 * self.channels
}
fn output_item(&self) -> usize {
self.output.0 * self.output.1 * self.out_channels()
}
#[inline]
fn tap_offset(&self, oh: usize, ow: usize, kh: usize, kw: usize) -> Option<usize> {
let ih = (oh * self.strides.0 + kh).checked_sub(self.pad_before.0)?;
let iw = (ow * self.strides.1 + kw).checked_sub(self.pad_before.1)?;
if ih >= self.input.0 || iw >= self.input.1 {
return None;
}
Some((ih * self.input.1 + iw) * self.channels)
}
}
pub(super) fn depthwise_forward_row(
g: &DepthwiseGeometry,
src: &[f32],
ker: &[f32],
bias: Option<&[f32]>,
b: usize,
oh: usize,
out_row: &mut [f32],
) {
let (kh_size, kw_size) = g.kernel;
let out_channels = g.out_channels();
let dm = g.depth_multiplier;
let in_item = b * g.input_item();
for ow in 0..g.output.1 {
let acc = &mut out_row[ow * out_channels..(ow + 1) * out_channels];
match bias {
Some(bias) => acc.copy_from_slice(bias),
None => acc.fill(0.0),
}
for kh in 0..kh_size {
for kw in 0..kw_size {
let Some(off) = g.tap_offset(oh, ow, kh, kw) else {
continue;
};
let x = &src[in_item + off..][..g.channels];
let k = &ker[(kh * kw_size + kw) * out_channels..][..out_channels];
if dm == 1 {
for ((a, &xc), &kc) in acc.iter_mut().zip(x).zip(k) {
*a += xc * kc;
}
} else {
for (c, &xc) in x.iter().enumerate() {
let base = c * dm;
for m in 0..dm {
acc[base + m] += xc * k[base + m];
}
}
}
}
}
}
}
pub(super) struct DepthwiseGradients {
pub weight: Vec<f32>,
pub bias: Vec<f32>,
pub input: Vec<f32>,
}
pub(super) fn depthwise_item_gradients(
g: &DepthwiseGeometry,
src: &[f32],
grad: &[f32],
ker: &[f32],
b: usize,
) -> DepthwiseGradients {
let (kh_size, kw_size) = g.kernel;
let out_channels = g.out_channels();
let dm = g.depth_multiplier;
let mut weight = vec![0.0f32; kh_size * kw_size * out_channels];
let mut bias = vec![0.0f32; out_channels];
let mut input = vec![0.0f32; g.input_item()];
let in_item = b * g.input_item();
let g_item = b * g.output_item();
for oh in 0..g.output.0 {
for ow in 0..g.output.1 {
let gr = &grad[g_item + (oh * g.output.1 + ow) * out_channels..][..out_channels];
for (j, &gj) in gr.iter().enumerate() {
bias[j] += gj;
}
for kh in 0..kh_size {
for kw in 0..kw_size {
let Some(off) = g.tap_offset(oh, ow, kh, kw) else {
continue;
};
let k_off = (kh * kw_size + kw) * out_channels;
if dm == 1 {
let x = &src[in_item + off..][..g.channels];
let wg = &mut weight[k_off..][..g.channels];
let kc = &ker[k_off..][..g.channels];
let dx = &mut input[off..][..g.channels];
for ((((w, d), &xc), &kv), &gj) in
wg.iter_mut().zip(dx.iter_mut()).zip(x).zip(kc).zip(gr)
{
*w += xc * gj;
*d += kv * gj;
}
} else {
for c in 0..g.channels {
let xc = src[in_item + off + c];
let base = c * dm;
let mut dxc = 0.0f32;
for m in 0..dm {
let gj = gr[base + m];
weight[k_off + base + m] += xc * gj;
dxc += ker[k_off + base + m] * gj;
}
input[off + c] += dxc;
}
}
}
}
}
}
DepthwiseGradients {
weight,
bias,
input,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tap_offset_skips_padding_on_both_edges() {
let g = DepthwiseGeometry {
input: (3, 3),
output: (3, 3),
channels: 2,
depth_multiplier: 1,
kernel: (3, 3),
strides: (1, 1),
pad_before: (1, 1),
};
assert_eq!(g.tap_offset(0, 0, 0, 0), None);
assert_eq!(g.tap_offset(0, 0, 1, 1), Some(0));
assert_eq!(g.tap_offset(2, 2, 2, 2), None);
assert_eq!(g.tap_offset(2, 2, 1, 1), Some(16));
}
}