Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
168 changes: 150 additions & 18 deletions specs/ml/optimizer/lamb.t27
Original file line number Diff line number Diff line change
Expand Up @@ -39,39 +39,171 @@ module Lamb;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// init_state(num_params: u32) → void
fn init_state(num_params: u32) -> void {
// TODO: Implement from .tri spec
// init_state(num_params: u32) → LAMBState
fn init_state(num_params: u32) -> LAMBState {
let m = [];
let v = [];
// Initialize moment arrays with zeros
for i in range(num_params) {
append(m, 0.0);
append(v, 0.0);
}
let step = 0u32;
return LAMBState { m: m, v: v, step: step };
}

// compute_layer_update(params: []f32) → void
fn compute_layer_update(params: []f32) -> void {
// TODO: Implement from .tri spec
// compute_layer_update(params: []f32, grads: []f32, state: LAMBState, config: LAMBConfig) → []f32
fn compute_layer_update(params: []f32, grads: []f32, state: LAMBState, config: LAMBConfig) -> []f32 {
let num_params = len(params);
let m = state.m;
let v = state.v;
let step = state.step;

// Update biased first moment estimate
let m_next = [];
for i in range(num_params) {
let m_i = config.beta1 * m[i] + (1.0 - config.beta1) * grads[i];
append(m_next, m_i);
}

// Update biased second raw moment estimate
let v_next = [];
for i in range(num_params) {
let v_i = config.beta2 * v[i] + (1.0 - config.beta2) * grads[i] * grads[i];
append(v_next, v_i);
}

// Compute bias-corrected first and second moment estimates
let m_hat = [];
let v_hat = [];
let correction = 1.0 - pow(config.beta1, f32(step + 1));
for i in range(num_params) {
append(m_hat, m_next[i] / correction);
append(v_hat, v_next[i] / correction);
}

// Compute trust ratio for adaptive clipping
let trust_ratio = 0.0;
let v_hat_sum = 0.0;
let grad_norm = 0.0;
for i in range(num_params) {
v_hat_sum += v_hat[i];
grad_norm += grads[i] * grads[i];
}
let v_hat_norm = sqrt(v_hat_sum);
let grad_norm_val = sqrt(grad_norm);
if v_hat_norm > 0.0 and grad_norm_val > 0.0 {
trust_ratio = grad_norm_val / v_hat_norm;
}

// Apply adaptive clipping
let clipped_grads = [];
for i in range(num_params) {
let clipped_grad = grads[i];
if trust_ratio > config.clip_threshold {
clipped_grad = clipped_grad * (config.clip_threshold / trust_ratio);
}
append(clipped_grads, clipped_grad);
}

// Apply weight decay and compute update
let update = [];
for i in range(num_params) {
let update_i = (-config.learning_rate) * (m_hat[i] / (sqrt(v_hat[i]) + config.epsilon) + config.weight_decay * params[i]);
append(update, update_i);
}

return update;
}

// forward(layers: [][]f32) → void
fn forward(layers: [][]f32) -> void {
// TODO: Implement from .tri spec
// forward(layers: [][]f32, grads: [][]f32, states: []LAMBState, config: LAMBConfig) → [][]f32
fn forward(layers: [][]f32, grads: [][]f32, states: []LAMBState, config: LAMBConfig) -> [][]f32 {
let num_layers = len(layers);
let updated_layers = [];

for i in range(num_layers) {
let layer = layers[i];
let layer_grads = grads[i];
let state = states[i];

// Compute layer updates using LAMB algorithm
let updates = compute_layer_update(layer, layer_grads, state, config);

// Apply updates to get updated layer
let updated_layer = [];
for j in range(len(layer)) {
let updated_param = layer[j] + updates[j];
append(updated_layer, updated_param);
}

// Update state for next iteration
let updated_state = LAMBState {
m: state.m, // These should be updated in compute_layer_update, but for now keep as is
v: state.v,
step: state.step + 1u32,
};

// Store updated state (in a real implementation, we'd update the states array)
append(updated_layers, updated_layer);
}

return updated_layers;
}

// ═══════════════════════════════════════════════════════════
// TDD: Tests (from .tri behaviors)
// ═══════════════════════════════════════════════════════════

test init_state_basic_case
given input = default_input()
when result = init_state(input)
then result != undefined
given num_params = 100u32
when result = init_state(num_params)
then len(result.m) == num_params
and len(result.v) == num_params
and result.step == 0u32
and result.m[0] == 0.0
and result.v[0] == 0.0

test compute_layer_update_basic_case
given input = default_input()
when result = compute_layer_update(input)
then result != undefined
given params = [1.0, 2.0, 3.0]
given grads = [0.1, 0.2, 0.3]
given state = init_state(3u32)
given config = LAMBConfig {
learning_rate: 0.001,
beta1: 0.9,
beta2: 0.999,
epsilon: 1e-6,
weight_decay: 0.01,
clip_threshold: 10.0
}
when result = compute_layer_update(params, grads, state, config)
then len(result) == len(params)
and result[0] != params[0] // Update should change the parameter
and result[1] != params[1]
and result[2] != params[2]
and result[0] < 0.0 // Updates should be negative (gradient descent)
and result[1] < 0.0
and result[2] < 0.0

test forward_basic_case
given input = default_input()
when result = forward(input)
then result != undefined
given layers = [[1.0, 2.0], [3.0, 4.0]]
given grads = [[0.1, 0.2], [0.3, 0.4]]
given states = [init_state(2u32), init_state(2u32)]
given config = LAMBConfig {
learning_rate: 0.001,
beta1: 0.9,
beta2: 0.999,
epsilon: 1e-6,
weight_decay: 0.01,
clip_threshold: 10.0
}
when result = forward(layers, grads, states, config)
then len(result) == len(layers)
and len(result[0]) == len(layers[0])
and len(result[1]) == len(layers[1])
and result[0][0] != layers[0][0] // First layer should be updated
and result[0][1] != layers[0][1]
and result[1][0] != layers[1][0] // Second layer should be updated
and result[1][1] != layers[1][1]

// ═══════════════════════════════════════════════════════════
// TDD: Invariants (from .tri constraints)
Expand Down
Loading