// conv2d_backward.mbt 鈥?Conv2d backward pass.
//
// Conv2d forward:
// out[n, co, ho, wo] = bias[co] + sum_{ci, kh, kw} W[co, ci, kh, kw]
// * input[n, ci, ho*stride-pad+kh, wo*stride-pad+kw]
// (out-of-bounds input values are 0.)
//
// Conv2d backward (three gradients):
// d_input[n, ci, ih, iw] = sum_{co, kh, kw} W[co, ci, kh, kw]
// * d_output[n, co, ih-kh+pad, ...] / stride
// d_weight[co, ci, kh, kw] = sum_{n, ho, wo} d_output[n, co, ho, wo]
// * input[n, ci, ho*stride-pad+kh, wo*stride-pad+kw]
// d_bias[co] = sum_{n, ho, wo} d_output[n, co, ho, wo]
///|
/// Conv2d backward pass. Returns `(d_input, d_weight, d_bias)`.
pub fn conv2d_backward(
input : Array[Float],
d_output : Array[Float],
n : Int,
c_in : Int,
h : Int,
w : Int,
param : Conv2dParam,
) -> (Array[Float], Array[Float], Array[Float]) {
let c_out = param.c_out
let kh = param.kh
let kw = param.kw
let stride = param.stride
let pad = param.pad
let ho = (h + 2 * pad - kh) / stride + 1
let wo = (w + 2 * pad - kw) / stride + 1
let d_input : Array[Float] = Array::make(n * c_in * h * w, 0.0F)
let d_weight : Array[Float] = Array::make(c_out * c_in * kh * kw, 0.0F)
let d_bias : Array[Float] = Array::make(c_out, 0.0F)
let input_n_stride = c_in * h * w
let input_c_stride = h * w
let input_h_stride = w
let weight_co_stride = c_in * kh * kw
let weight_ci_stride = kh * kw
let d_out_n_stride = c_out * ho * wo
let d_out_co_stride = ho * wo
// ---- d_bias ----
for batch in 0..= 0 && ih < h {
let mut out_w = 0
while out_w < wo {
let iw = out_w * stride - pad + kw_i
if iw >= 0 && iw < w {
acc = acc +
d_output[d_out_co_off + out_h * wo + out_w] *
input[in_c_off + ih * input_h_stride + iw]
}
out_w = out_w + 1
}
}
out_h = out_h + 1
}
d_weight[w_ci_off + kh_i * kw + kw_i] = d_weight[w_ci_off + kh_i * kw + kw_i] + acc
kw_i = kw_i + 1
}
kh_i = kh_i + 1
}
ci = ci + 1
}
}
}
// ---- d_input ----
// d_input[n, ci, ih, iw] = sum_{co, kh, kw}
// W[co, ci, kh, kw] * d_output[n, co, ih+pad-kh)/stride, ...]
for batch in 0..= 0 && ho_idx % stride == 0 {
let ho_k = ho_idx / stride
if ho_k < ho {
let mut kw_i = 0
while kw_i < kw {
let wo_idx = iw + pad - kw_i
if wo_idx >= 0 && wo_idx % stride == 0 {
let wo_k = wo_idx / stride
if wo_k < wo {
acc = acc +
param.weight[w_ci_off + kh_i * kw + kw_i] *
d_output[d_out_co_off + ho_k * wo + wo_k]
}
}
kw_i = kw_i + 1
}
}
}
kh_i = kh_i + 1
}
co = co + 1
}
d_input[in_c_off + ih * input_h_stride + iw] = acc
iw = iw + 1
}
ih = ih + 1
}
ci = ci + 1
}
}
(d_input, d_weight, d_bias)
}