// optimizer_sgd.mbt — SGD + SGD-with-momentum (v0.15.0).
//
// All optimisers operate on raw `(Array[Float], Array[Float])` tuples for
// (weight, bias). Layer-specific wrappers (`sgd_update_linear` /
// `sgd_update_conv`) delegate to the array-level helpers, copying the
// resulting arrays into a fresh parameter struct (params are immutable).
//
// Conventions:
// - `weight[i] -= lr * d_weight[i]` (vanilla SGD)
// - SGD-with-momentum (Polyak / Sutskever):
// v <- momentum * v + d_weight
// weight <- weight - lr * v
// Equivalently `v <- momentum * v + d_weight` then `weight -= lr * v`,
// which is the PyTorch convention (NOT the original Sutskever-2013
// `weight -= lr * (momentum * v + grad)` form).
// - Bias updates use the same rule, kept independent from weight
// velocity (no shared buffer).
// - All state is Float32; we follow the project's bit-exact contract.
///|
/// Vanilla SGD: returns a fresh `(weight, bias)` pair after applying
/// `weight -= lr * d_weight` and `bias -= lr * d_bias` element-wise.
pub fn sgd_update_arrays(
weight : Array[Float],
bias : Array[Float],
d_weight : Array[Float],
d_bias : Array[Float],
lr : Float,
) -> (Array[Float], Array[Float]) {
let w2 : Array[Float] = Array::make(weight.length(), 0.0F)
for i in 0.. SGDMomentumState {
{ v_w: Array::make(weight_len, 0.0F), v_b: Array::make(bias_len, 0.0F) }
}
///|
/// One SGD-with-momentum step. Mutates and returns the velocity state
/// alongside the updated `(weight, bias)` arrays.
///
/// v_w <- momentum * v_w + d_weight
/// v_b <- momentum * v_b + d_bias
/// weight <- weight - lr * v_w
/// bias <- bias - lr * v_b
pub fn sgd_momentum_update_arrays(
weight : Array[Float],
bias : Array[Float],
d_weight : Array[Float],
d_bias : Array[Float],
state : SGDMomentumState,
lr : Float,
momentum : Float,
) -> (Array[Float], Array[Float], SGDMomentumState) {
// Update velocity.
let new_v_w : Array[Float] = Array::make(weight.length(), 0.0F)
for i in 0.. LinearParam {
let (w, b) = sgd_update_arrays(param.weight, param.bias, d_weight, d_bias, lr)
{ weight: w, bias: b, in_features: param.in_features, out_features: param.out_features }
}
///|
/// SGD wrapper for `Conv2dParam`. Returns a fresh `Conv2dParam` with
/// updated weights / biases.
pub fn sgd_update_conv(
param : Conv2dParam,
d_weight : Array[Float],
d_bias : Array[Float],
lr : Float,
) -> Conv2dParam {
let (w, b) = sgd_update_arrays(param.weight, param.bias, d_weight, d_bias, lr)
{ weight: w, bias: b, c_out: param.c_out, c_in: param.c_in,
kh: param.kh, kw: param.kw, stride: param.stride, pad: param.pad }
}
///|
/// SGD-with-momentum wrapper for `LinearParam`.
pub fn sgd_momentum_update_linear(
param : LinearParam,
d_weight : Array[Float],
d_bias : Array[Float],
state : SGDMomentumState,
lr : Float,
momentum : Float,
) -> (LinearParam, SGDMomentumState) {
let (w, b, s2) = sgd_momentum_update_arrays(
param.weight, param.bias, d_weight, d_bias, state, lr, momentum,
)
let p2 = { weight: w, bias: b, in_features: param.in_features,
out_features: param.out_features }
(p2, s2)
}
///|
/// SGD-with-momentum wrapper for `Conv2dParam`.
pub fn sgd_momentum_update_conv(
param : Conv2dParam,
d_weight : Array[Float],
d_bias : Array[Float],
state : SGDMomentumState,
lr : Float,
momentum : Float,
) -> (Conv2dParam, SGDMomentumState) {
let (w, b, s2) = sgd_momentum_update_arrays(
param.weight, param.bias, d_weight, d_bias, state, lr, momentum,
)
let p2 = { weight: w, bias: b, c_out: param.c_out, c_in: param.c_in,
kh: param.kh, kw: param.kw, stride: param.stride, pad: param.pad }
(p2, s2)
}