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
30 changes: 24 additions & 6 deletions specs/ml/layers/flatten_layer.t27
Original file line number Diff line number Diff line change
Expand Up @@ -29,14 +29,32 @@ module Flatten;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// calc_output_size(input_dims: []u32) → void
fn calc_output_size(input_dims: []u32) -> void {
// TODO: Implement from .tri spec
// calc_output_size(input_dims: []u32) → u32
fn calc_output_size(input_dims: []u32) -> u32 {
// Calculate the total number of elements in the input tensor
var total_size: u32 = 1;
for i in range(len(input_dims)) {
total_size = total_size * input_dims[i];
}
return total_size;
}

// forward(input: []f32) → void
fn forward(input: []f32) -> void {
// TODO: Implement from .tri spec
// forward(input: []f32, config: FlattenConfig) → []f32
fn forward(input: []f32, config: FlattenConfig) -> []f32 {
// Calculate output size using the existing function
var output_size: u32 = calc_output_size(config.input_shape);

// Create output array
var output: []f32 = make([]f32, output_size);

// Flatten the input tensor into the output array
var output_index: u32 = 0;
for input_val in input {
output[output_index] = input_val;
output_index = output_index + 1;
}

return output;
}

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