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
70 changes: 54 additions & 16 deletions specs/ml/recurrent/rnn_cell.t27
Original file line number Diff line number Diff line change
Expand Up @@ -37,19 +37,53 @@ module RnnCell;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// forward_step(input: []f32) → void
fn forward_step(input: []f32) -> void {
// TODO: Implement from .tri spec
// forward_step(input: []f32, state: RNNState, params: RNNParams) -> RNNState
fn forward_step(input: []f32, state: RNNState, params: RNNParams) -> RNNState {
// Basic RNN forward step implementation
// Apply linear transformation to input
let input_transformed: []f32 = matmul(params.w_x, input);

// Apply linear transformation to hidden state
let hidden_transformed: []f32 = matmul(params.w_h, state.hidden);

// Add bias and apply activation function (tanh for cell state)
let cell_input: []f32 = add(add(input_transformed, hidden_transformed), params.b);
let new_cell: []f32 = tanh(cell_input);

// Apply activation function (tanh) for hidden state
let new_hidden: []f32 = tanh(new_cell);

return RNNState {
hidden: new_hidden,
cell: new_cell,
};
}

// forward_sequence(inputs: [][]f32) → void
fn forward_sequence(inputs: [][]f32) -> void {
// TODO: Implement from .tri spec
// forward_sequence(inputs: [][]f32, initial_state: []f32, params: RNNParams) -> [][]f32
fn forward_sequence(inputs: [][]f32, initial_state: []f32, params: RNNParams) -> [][]f32 {
// Basic RNN sequence processing implementation
let sequence_length = len(inputs);
let mut output_sequence: [][]f32 = [];

// Initialize initial state
let mut current_state = RNNState {
hidden: initial_state,
cell: zeros(len(initial_state)),
};

// Process each input in sequence
for i in 0..sequence_length {
current_state = forward_step(inputs[i], current_state, params);
output_sequence = append(output_sequence, current_state.hidden);
}

return output_sequence;
}

// init_state(hidden_size: u32) → void
fn init_state(hidden_size: u32) -> void {
// TODO: Implement from .tri spec
// init_state(hidden_size: u32) -> []f32
fn init_state(hidden_size: u32) -> []f32 {
// Initialize RNN state with zeros
return zeros(hidden_size);
}

// ═══════════════════════════════════════════════════════════
Expand All @@ -58,18 +92,23 @@ module RnnCell;

test forward_step_basic_case
given input = default_input()
when result = forward_step(input)
state = default_state()
params = default_params()
when result = forward_step(input, state, params)
then result != undefined

test forward_sequence_basic_case
given input = default_input()
when result = forward_sequence(input)
given inputs = default_sequence_inputs()
initial_state = default_input()
params = default_params()
when result = forward_sequence(inputs, initial_state, params)
then result != undefined

test init_state_basic_case
given input = default_input()
when result = init_state(input)
given hidden_size = 64
when result = init_state(hidden_size)
then result != undefined
then len(result) == hidden_size

// ═══════════════════════════════════════════════════════════
// TDD: Invariants (from .tri constraints)
Expand All @@ -93,5 +132,4 @@ module RnnCell;

invariant rnn_cell_constraint_4
given input = valid_input()
then true // "Goodfellow et al. (2016) - Sequence Modeling with RNNs and LSTMs"

then true // "Goodfellow et al. (2016) - Sequence Modeling with RNNs and LSTMs"
Loading