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
80 changes: 68 additions & 12 deletions specs/ml/recurrent/lstm_cell.t27
Original file line number Diff line number Diff line change
Expand Up @@ -41,28 +41,84 @@ module Lstm;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// forward(input: []const f32) → void
fn forward(input: []const f32) -> void {
// TODO: Implement from .tri spec
// forward(input: []const f32, state_prev: LSTMState, weights: LSTMWeights, state_next: LSTMState, config: LSTMConfig) → void
fn forward(input: []const f32, state_prev: LSTMState, weights: LSTMWeights, state_next: LSTMState, config: LSTMConfig) -> void {
// LSTM forward pass implementation
// Extract input size from input slice
let input_size = input.size();

// Validate input size matches configuration
assert(input_size == config.input_size);

// Concatenate input with previous hidden state
let concat_input = math::concatenate(input, state_prev.h);

// Compute gate activations
let forget_gate = math::sigmoid(math::dot(concat_input, weights.Wf) + weights.bf);
let input_gate = math::sigmoid(math::dot(concat_input, weights.Wi) + weights.bi);
let output_gate = math::sigmoid(math::dot(concat_input, weights.Wo) + weights.bo);
let cell_gate = math::tanh(math::dot(concat_input, weights.Wg) + weights.bg);

// Update cell state
state_next.c = math::add(math::multiply(forget_gate, state_prev.c),
math::multiply(input_gate, cell_gate));

// Compute new hidden state
let activated_cell = math::tanh(state_next.c);
state_next.h = math::multiply(output_gate, activated_cell);
}

// init(weights: LSTMWeights) → void
fn init(weights: LSTMWeights) -> void {
// TODO: Implement from .tri spec
// init(weights: LSTMWeights, config: LSTMConfig) → void
fn init(weights: LSTMWeights, config: LSTMConfig) -> void {
// LSTM initialization implementation
// Validate weight dimensions match configuration
let expected_weight_size = (config.input_size + config.hidden_size) * config.hidden_size;
let expected_bias_size = config.hidden_size;

assert(weights.Wf.size() == expected_weight_size);
assert(weights.Wi.size() == expected_weight_size);
assert(weights.Wo.size() == expected_weight_size);
assert(weights.Wg.size() == expected_weight_size);

assert(weights.bf.size() == expected_bias_size);
assert(weights.bi.size() == expected_bias_size);
assert(weights.bo.size() == expected_bias_size);
assert(weights.bg.size() == expected_bias_size);

// Initialize forget gate bias to 1.0 for better convergence
// Note: Using array fill operations instead of loops
weights.bf = math::fill(weights.bf.size(), 1.0);

// Initialize other biases to 0.0
weights.bi = math::fill(weights.bi.size(), 0.0);
weights.bo = math::fill(weights.bo.size(), 0.0);
weights.bg = math::fill(weights.bg.size(), 0.0);

// Initialize weights with small random values (Xavier initialization)
let scale = 1.0 / math::sqrt(f32(config.input_size + config.hidden_size));
weights.Wf = math::random_uniform_array(weights.Wf.size(), -scale, scale);
weights.Wi = math::random_uniform_array(weights.Wi.size(), -scale, scale);
weights.Wo = math::random_uniform_array(weights.Wo.size(), -scale, scale);
weights.Wg = math::random_uniform_array(weights.Wg.size(), -scale, scale);
}

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

test forward_basic_case
given input = default_input()
when result = forward(input)
test "forward_basic_case"
given input = default_input(),
state_prev = default_lstm_state(),
weights = default_lstm_weights(),
state_next = default_lstm_state(),
config = default_lstm_config()
when result = forward(input, state_prev, weights, state_next, config)
then result != undefined

test init_basic_case
given input = default_input()
when result = init(input)
test "init_basic_case"
given weights = default_lstm_weights(),
config = default_lstm_config()
when result = init(weights, config)
then result != undefined

// ═══════════════════════════════════════════════════════════
Expand Down
Loading