From bf9ee21454a06d84eb00a3d742572ea57ae6ed95 Mon Sep 17 00:00:00 2001 From: Trinity Bee Date: Tue, 15 Sep 2026 18:58:29 +0000 Subject: [PATCH 1/2] Implement LSTM forward and init functions with proper signatures and tests - Update forward function signature to match quoted specification - Update init function signature to match quoted specification - Implement LSTM forward pass with proper gate computations - Implement LSTM initialization with Xavier initialization and bias setup - Update test blocks to match new function signatures - Add proper validation and assertions Closes #3713 --- specs/ml/recurrent/lstm_cell.t27 | 76 +++++++++++++++++++++++++++----- 1 file changed, 66 insertions(+), 10 deletions(-) diff --git a/specs/ml/recurrent/lstm_cell.t27 b/specs/ml/recurrent/lstm_cell.t27 index 036d0ede12..1a154bda5d 100644 --- a/specs/ml/recurrent/lstm_cell.t27 +++ b/specs/ml/recurrent/lstm_cell.t27 @@ -41,14 +41,65 @@ 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); } // ═══════════════════════════════════════════════════════════ @@ -56,13 +107,18 @@ module Lstm; // ═══════════════════════════════════════════════════════════ test forward_basic_case - given input = default_input() - when result = forward(input) + 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) + given weights = default_lstm_weights(), + config = default_lstm_config() + when result = init(weights, config) then result != undefined // ═══════════════════════════════════════════════════════════ From 95647da5ee15cf3e7bd2be0c338eb8be22f119f6 Mon Sep 17 00:00:00 2001 From: Trinity Bee Date: Tue, 15 Sep 2026 23:20:37 +0000 Subject: [PATCH 2/2] Update LSTM cell test format to use quoted test names Closes #3713 --- specs/ml/recurrent/lstm_cell.t27 | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/specs/ml/recurrent/lstm_cell.t27 b/specs/ml/recurrent/lstm_cell.t27 index 1a154bda5d..7549b7b41c 100644 --- a/specs/ml/recurrent/lstm_cell.t27 +++ b/specs/ml/recurrent/lstm_cell.t27 @@ -106,7 +106,7 @@ module Lstm; // TDD: Tests (from .tri behaviors) // ═══════════════════════════════════════════════════════════ - test forward_basic_case + test "forward_basic_case" given input = default_input(), state_prev = default_lstm_state(), weights = default_lstm_weights(), @@ -115,7 +115,7 @@ module Lstm; when result = forward(input, state_prev, weights, state_next, config) then result != undefined - test init_basic_case + test "init_basic_case" given weights = default_lstm_weights(), config = default_lstm_config() when result = init(weights, config)