// residual_block.mbt — Residual block (v0.20.2).
//
// Implements the basic residual block from He et al. 2015:
//
// y = F(x, W) + W_s * x // F = conv-bn-relu-conv-bn, W_s = projection (optional)
// out = relu(y)
//
// `shortcut_conv` and `shortcut_bn` are `None` for identity shortcut
// (input shape must match the block output). For downsample blocks
// (stride > 1 or c changes), pass `Some(1x1 conv)` + `Some(BN)`.
//
// Cache holds everything backward needs:
// x, sum (pre-ReLU), bn1_out (pre-ReLU of branch),
// bn1_cache, bn2_cache, shortcut_bn_cache? (if any),
// conv2_input (ReLU of bn1_out).
//
// Conventions:
// - Stride is applied to the first conv (stride=2 in downsample
// blocks halves the spatial size).
// - The second conv uses stride=1.
// - Both main-path convs use 3x3 kernels with pad=1 (preserves
// spatial size at stride=1).
// - BN follows each conv (pre-activation is NOT used; this is the
// original ResNet v1 design).
///|
/// A basic residual block.
pub struct ResidualBlock {
conv1 : Conv2dParam
bn1 : BatchNorm2d
conv2 : Conv2dParam
bn2 : BatchNorm2d
shortcut_conv : Conv2dParam?
shortcut_bn : BatchNorm2d?
stride : Int
}
///|
/// Cache returned by `residual_block_forward` and consumed by
/// `residual_block_backward`.
pub struct ResidualCache {
// Saved input tensor x (for shortcut path).
x : Array[Float]
// Saved post-add sum: needed by ReLU backward.
sum : Array[Float]
// Saved post-BN1, pre-ReLU activation: needed by ReLU backward.
bn1_out : Array[Float]
// BN caches for the main path's two BatchNorm2d calls.
bn1_cache : BatchNormCache
bn2_cache : BatchNormCache
// BN cache for the shortcut BN (if any).
shortcut_bn_cache : BatchNormCache?
// Saved shortcut output (BN output of shortcut conv, or identity x).
shortcut_output : Array[Float]
// Input / intermediate shapes (for conv2 backward).
conv2_input : Array[Float]
conv2_input_n : Int
conv2_input_c : Int
conv2_input_h : Int
conv2_input_w : Int
// Input shape.
n : Int
c : Int
h : Int
w : Int
}
///|
/// Parameter gradients returned by `residual_block_backward`.
pub struct ResidualGrads {
conv1_d_weight : Array[Float]
conv1_d_bias : Array[Float]
bn1_d_gamma : Array[Float]
bn1_d_beta : Array[Float]
conv2_d_weight : Array[Float]
conv2_d_bias : Array[Float]
bn2_d_gamma : Array[Float]
bn2_d_beta : Array[Float]
shortcut_conv_d_weight : Array[Float]?
shortcut_conv_d_bias : Array[Float]?
shortcut_bn_d_gamma : Array[Float]?
shortcut_bn_d_beta : Array[Float]?
}
///|
/// Build an identity-shortcut basic block (input shape must match
/// block output: c_in == c_out, stride == 1).
pub fn ResidualBlock::identity(c : Int, seed : UInt64) -> ResidualBlock {
let rng = Xoshiro::new(seed)
let w1 = he_init(rng, c * c * 3 * 3, c * 3 * 3)
let b1 : Array[Float] = Array::make(c, 0.0F)
let conv1 = Conv2dParam::new(w1, b1, c, c, 3, 3, stride=1, pad=1)
let bn1 = BatchNorm2d::new(c)
let w2 = he_init(rng, c * c * 3 * 3, c * 3 * 3)
let b2 : Array[Float] = Array::make(c, 0.0F)
let conv2 = Conv2dParam::new(w2, b2, c, c, 3, 3, stride=1, pad=1)
let bn2 = BatchNorm2d::new(c)
{ conv1, bn1, conv2, bn2, shortcut_conv: None, shortcut_bn: None, stride: 1 }
}
///|
/// Build a downsample basic block with a 1x1 conv projection shortcut.
pub fn ResidualBlock::projection(
c_in : Int,
c_out : Int,
stride : Int,
seed : UInt64,
) -> ResidualBlock {
let rng = Xoshiro::new(seed)
let w1 = he_init(rng, c_out * c_in * 3 * 3, c_in * 3 * 3)
let b1 : Array[Float] = Array::make(c_out, 0.0F)
let conv1 = Conv2dParam::new(w1, b1, c_out, c_in, 3, 3, stride=stride, pad=1)
let bn1 = BatchNorm2d::new(c_out)
let w2 = he_init(rng, c_out * c_out * 3 * 3, c_out * 3 * 3)
let b2 : Array[Float] = Array::make(c_out, 0.0F)
let conv2 = Conv2dParam::new(w2, b2, c_out, c_out, 3, 3, stride=1, pad=1)
let bn2 = BatchNorm2d::new(c_out)
let ws = he_init(rng, c_out * c_in, c_in)
let bs : Array[Float] = Array::make(c_out, 0.0F)
let shortcut_conv = Conv2dParam::new(ws, bs, c_out, c_in, 1, 1, stride=stride, pad=0)
let shortcut_bn = BatchNorm2d::new(c_out)
{
conv1, bn1, conv2, bn2,
shortcut_conv: Some(shortcut_conv), shortcut_bn: Some(shortcut_bn),
stride,
}
}
///|
/// Forward pass for a residual block.
pub fn residual_block_forward(
b : ResidualBlock,
x : Array[Float],
n : Int,
c : Int,
h : Int,
w : Int,
) -> (Array[Float], ResidualCache) {
// ---- Main path ----
let conv1_out = conv2d_forward(x, n, c, h, w, b.conv1)
let h1 = (h + 2 * b.conv1.pad - b.conv1.kh) / b.conv1.stride + 1
let w1 = (w + 2 * b.conv1.pad - b.conv1.kw) / b.conv1.stride + 1
let c1 = b.conv1.c_out
let (bn1_out, bn1_cache) = batch_norm2d_forward(conv1_out, n, c1, h1, w1, b.bn1)
let relu1 = relu_forward(bn1_out)
let conv2_out = conv2d_forward(relu1, n, c1, h1, w1, b.conv2)
let h2 = (h1 + 2 * b.conv2.pad - b.conv2.kh) / b.conv2.stride + 1
let w2 = (w1 + 2 * b.conv2.pad - b.conv2.kw) / b.conv2.stride + 1
let c2 = b.conv2.c_out
let (bn2_out, bn2_cache) = batch_norm2d_forward(conv2_out, n, c2, h2, w2, b.bn2)
// ---- Shortcut path ----
let (shortcut_out, sbn_cache_opt) = match (b.shortcut_conv, b.shortcut_bn) {
(Some(scv), Some(sbn)) => {
let sout_pre = conv2d_forward(x, n, c, h, w, scv)
let sh_pre = (h + 2 * scv.pad - scv.kh) / scv.stride + 1
let sw_pre = (w + 2 * scv.pad - scv.kw) / scv.stride + 1
let (sout, sbn_cache) = batch_norm2d_forward(
sout_pre, n, scv.c_out, sh_pre, sw_pre, sbn,
)
(sout, Some(sbn_cache))
}
_ => (x, None)
}
// ---- Sum + ReLU ----
let sum = add_forward(bn2_out, shortcut_out)
let out = relu_forward(sum)
let cache : ResidualCache = {
x,
sum,
bn1_out,
bn1_cache,
bn2_cache,
shortcut_bn_cache: sbn_cache_opt,
shortcut_output: shortcut_out,
conv2_input: relu1,
conv2_input_n: n,
conv2_input_c: c1,
conv2_input_h: h1,
conv2_input_w: w1,
n, c, h, w,
}
(out, cache)
}
///|
/// Backward pass for a residual block. Returns (d_input, grads).
pub fn residual_block_backward(
b : ResidualBlock,
cache : ResidualCache,
d_output : Array[Float],
) -> (Array[Float], ResidualGrads) {
let n = cache.n
let c = cache.c
let h = cache.h
let w = cache.w
// ReLU backward at the block output.
let d_sum = relu_backward(cache.sum, d_output)
let (d_bn2, d_shortcut) = add_backward(d_sum)
let (d_conv2_in, d_bn2_gamma, d_bn2_beta) = batch_norm2d_backward(
d_bn2, cache.bn2_cache, b.bn2,
)
let (d_relu1, d_conv2_w, d_conv2_b) = conv2d_backward(
cache.conv2_input, d_conv2_in, cache.conv2_input_n,
cache.conv2_input_c, cache.conv2_input_h, cache.conv2_input_w,
b.conv2,
)
let d_bn1 = relu_backward(cache.bn1_out, d_relu1)
let (d_conv1_in, d_bn1_gamma, d_bn1_beta) = batch_norm2d_backward(
d_bn1, cache.bn1_cache, b.bn1,
)
let (d_x_main, d_conv1_w, d_conv1_b) = conv2d_backward(
cache.x, d_conv1_in, n, c, h, w, b.conv1,
)
match (b.shortcut_conv, b.shortcut_bn, cache.shortcut_bn_cache) {
(Some(scv), Some(sbn), Some(sbn_cache)) => {
// BN backward first (operating on d_shortcut of shape
// [n, scv.c_out, h_out, w_out]) -> returns d_shortcut_bn_in
// of the same shape.
let (d_shortcut_bn_in, d_sbn_g, d_sbn_b) = batch_norm2d_backward(
d_shortcut, sbn_cache, sbn,
)
// Then conv backward (operating on d_shortcut_bn_in of shape
// [n, scv.c_out, h_out, w_out]) -> returns d_x_shortcut of
// shape [n, c, h, w] matching d_x_main.
let (d_x_shortcut, d_sc_w, d_sc_b) = conv2d_backward(
cache.x, d_shortcut_bn_in, n, c, h, w, scv,
)
let merged : Array[Float] = Array::make(d_x_main.length(), 0.0F)
for i in 0.. {
let merged : Array[Float] = Array::make(d_x_main.length(), 0.0F)
for i in 0..