// conv2d.mbt 鈥?2D convolution primitive (forward only).
//
// Layout: NCHW row-major flat `Array[Float]`.
//
// Conventions:
// - Input: [n, c_in, h, w] length = n * c_in * h * w
// - Weight: [c_out, c_in, kh, kw] length = c_out * c_in * kh * kw
// - Bias: [c_out] length = c_out
// - Output: [n, c_out, ho, wo] length = n * c_out * ho * wo
//
// ho = (h + 2*pad - kh) / stride + 1
// wo = (w + 2*pad - kw) / stride + 1
//
// No external dependencies; pure Float32 loop, matches the project's
// bit-exact Float32 contract. Forward only; no backward (autograd
// not in scope).
///|
/// Conv2d parameter container. `weight` is laid out as
/// `[c_out, c_in, kh, kw]` row-major; `bias` is `[c_out]`.
/// `stride` is shared across spatial dims (no asymmetric strides yet);
/// `pad` is the half-padding applied symmetrically on both sides
/// (e.g. `pad = kh // 2` gives "same" padding for odd kernel sizes).
pub struct Conv2dParam {
weight : Array[Float]
bias : Array[Float]
c_out : Int
c_in : Int
kh : Int
kw : Int
stride : Int
pad : Int
}
///|
/// Build a Conv2dParam from raw arrays. Does NOT copy the arrays; the
/// caller retains ownership.
pub fn Conv2dParam::new(
weight : Array[Float],
bias : Array[Float],
c_out : Int,
c_in : Int,
kh : Int,
kw : Int,
stride? : Int = 1,
pad? : Int = 0,
) -> Conv2dParam {
{ weight, bias, c_out, c_in, kh, kw, stride, pad }
}
///|
/// Forward pass for a 2D convolution.
///
/// `input` : length = n * c_in * h * w
/// returns : length = n * c_out * ho * wo
///
/// Pad-region values are zero (i.e. zero-padded). Out-of-bounds reads
/// contribute nothing.
pub fn conv2d_forward(
input : Array[Float],
n : Int,
c_in : Int,
h : Int,
w : Int,
param : Conv2dParam,
) -> Array[Float] {
let c_out = param.c_out
let kh = param.kh
let kw = param.kw
let stride = param.stride
let pad = param.pad
// Output spatial dims.
let ho = (h + 2 * pad - kh) / stride + 1
let wo = (w + 2 * pad - kw) / stride + 1
let out : Array[Float] = Array::make(n * c_out * ho * wo, 0.0F)
// Precompute row strides for input / weight / output indexing.
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 out_n_stride = c_out * ho * wo
let out_co_stride = ho * wo
for batch in 0..= 0 && ih < h {
let mut out_w = 0
while out_w < wo {
let iw_base = out_w * stride - pad
let mut kw_i = 0
while kw_i < kw {
let iw = iw_base + kw_i
if iw >= 0 && iw < w {
let in_v = input[in_c_off + ih * input_h_stride + iw]
if in_v != 0.0F {
let w_v = param.weight[w_ci_off + kh_i * kw + kw_i]
if w_v != 0.0F {
let acc = out[out_co_off + out_h * wo + out_w]
out[out_co_off + out_h * wo + out_w] = acc +
w_v * in_v
}
}
}
kw_i = kw_i + 1
}
out_w = out_w + 1
}
}
kh_i = kh_i + 1
}
out_h = out_h + 1
}
ci = ci + 1
}
}
}
out
}