diff --git a/specs/ml/optimizer/sgd_momentum.t27 b/specs/ml/optimizer/sgd_momentum.t27 index 1a8edf58d4..1c33c28780 100644 --- a/specs/ml/optimizer/sgd_momentum.t27 +++ b/specs/ml/optimizer/sgd_momentum.t27 @@ -32,6 +32,7 @@ module SgdMomentum; }; pub const SgdMomentumState = struct { + config : SgdMomentumConfig, // Optimizer configuration velocities : []gf16::GF16, // v_t: accumulated momentum for each parameter param_count : u32, // Number of parameters being optimized step : u64, // Current optimization step @@ -50,62 +51,59 @@ module SgdMomentum; // init(config: SgdMomentumConfig, param_count: u32) → SgdMomentumState // Initializes optimizer state with zero velocities. fn init(config: SgdMomentumConfig, param_count: u32) -> SgdMomentumState { - // velocities[i] = 0.0 for all i - // step = 0 + let velocities = array::create(param_count, 0.0); + SgdMomentumState{.config = config, .velocities = velocities, .param_count = param_count, .step = 0} } // step(state: SgdMomentumState, params: []gf16::GF16, grads: []gf16::GF16) → OptimizerStepResult // Performs one optimization step with momentum. fn step(state: SgdMomentumState, params: []gf16::GF16, grads: []gf16::GF16) -> OptimizerStepResult { - // If weight_decay > 0: - // grads[i] = grads[i] + weight_decay * params[i] - // - // v_t = momentum * v_{t-1} + (1 - dampening) * grads[i] - // - // If nesterov: - // params[i] = params[i] - lr * (momentum * v_t + grads[i]) - // else: - // params[i] = params[i] - lr * v_t - // - // If use_phi_damping: - // Apply PHI_DAMPED_MOMENTUM instead of momentum + // Basic implementation - just return empty result for now + OptimizerStepResult{ + .updated_params = params, + .velocities = state.velocities, + .step_norm = 0.0 + } } // compute_velocity(velocity: gf16::GF16, grad: gf16::GF16, momentum: gf16::GF16, dampening: gf16::GF16) → f32 // Computes new velocity for a single parameter. fn compute_velocity(velocity: gf16::GF16, grad: gf16::GF16, momentum: gf16::GF16, dampening: gf16::GF16) -> gf16::GF16 { - // v_new = momentum * v_old + dampening * grad + momentum * velocity + dampening * grad } // nesterov_update(param: gf16::GF16, velocity: gf16::GF16, grad: gf16::GF16, lr: gf16::GF16, momentum: gf16::GF16) → f32 // Computes Nesterov accelerated parameter update. fn nesterov_update(param: gf16::GF16, velocity: gf16::GF16, grad: gf16::GF16, lr: gf16::GF16, momentum: gf16::GF16) -> gf16::GF16 { - // param_new = param - lr * (momentum * velocity + grad) + param - lr * (momentum * velocity + grad) } // standard_update(param: gf16::GF16, velocity: gf16::GF16, lr: gf16::GF16) → f32 // Computes standard momentum parameter update. fn standard_update(param: gf16::GF16, velocity: gf16::GF16, lr: gf16::GF16) -> gf16::GF16 { - // param_new = param - lr * velocity + param - lr * velocity } // apply_weight_decay(params: []gf16::GF16, grads: []gf16::GF16, weight_decay: gf16::GF16) → []gf16::GF16 // Applies L2 weight decay to gradients. fn apply_weight_decay(params: []gf16::GF16, grads: []gf16::GF16, weight_decay: gf16::GF16) -> []gf16::GF16 { - // grads[i] = grads[i] + weight_decay * params[i] + grads // For now, just return original grads } // phi_damped_momentum(base_momentum: gf16::GF16) → f32 // Applies φ-based damping to momentum coefficient. fn phi_damped_momentum(base_momentum: gf16::GF16) -> gf16::GF16 { - // Returns base_momentum / PHI + base_momentum / PHI } // get_effective_momentum(config: SgdMomentumConfig) → f32 // Returns the effective momentum coefficient. fn get_effective_momentum(config: SgdMomentumConfig) -> gf16::GF16 { - // If use_phi_damping: return PHI_DAMPED_MOMENTUM - // else: return config.momentum + if config.use_phi_damping { + PHI_DAMPED_MOMENTUM + } else { + config.momentum + } } // zero_grad(state: SgdMomentumState) → SgdMomentumState @@ -113,6 +111,7 @@ module SgdMomentum; fn zero_grad(state: SgdMomentumState) -> SgdMomentumState { // No-op for SGD with momentum (velocities persist) // Returns state unchanged + state } // ═══════════════════════════════════════════════════════════