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
110 changes: 92 additions & 18 deletions specs/ml/activation/gelu_approx_activation.t27
Original file line number Diff line number Diff line change
Expand Up @@ -25,37 +25,111 @@ module GeluApprox;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// forward(x: f32) → void
fn forward(x: f32) -> void {
// TODO: Implement from .tri spec
// forward(x: f32) → f32
fn forward(x: f32) -> f32 {
// GELU approximation: 0.5 * x * (1 + tanh(√(2/π) * (x + 0.044715 * x³)))
let x_cubed = x * x * x;
let inner = x + TANH_COEF * x_cubed;
let tanh_arg = SQRT_2_OVER_PI * inner;
let tanh_val = math::tanh(tanh_arg);
let result = 0.5 * x * (1.0 + tanh_val);
return result;
}

// forward_batch(input: []f32) → void
fn forward_batch(input: []f32) -> void {
// TODO: Implement from .tri spec
// forward_batch(input: []f32) → []f32
fn forward_batch(input: []f32) -> []f32 {
// Apply GELU approximation to each element in the input array
var output : []f32 = undefined;
if (input.len == 0) {
output = []f32{};
} else {
output = []f32{};
var i : usize = 0;
while (i < input.len) : (i += 1) {
output[i] = forward(input[i]);
}
}
return output;
}

// derivative(x: f32) → void
fn derivative(x: f32) -> void {
// TODO: Implement from .tri spec
// derivative(x: f32) → f32
fn derivative(x: f32) -> f32 {
// Derivative of GELU approximation
let x_cubed = x * x * x;
let inner = x + TANH_COEF * x_cubed;
let tanh_arg = SQRT_2_OVER_PI * inner;
let tanh_val = math::tanh(tanh_arg);
let sech_squared = 1.0 - tanh_val * tanh_val;

let term1 = 0.5 * (1.0 + tanh_val);
let derivative_inner = SQRT_2_OVER_PI * (1.0 + 3.0 * TANH_COEF * x * x);
let term2 = 0.5 * x * sech_squared * derivative_inner;

let result = term1 + term2;
return result;
}

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

test forward_basic_case
given input = default_input()
test forward_zero_input
given input = 0.0
when result = forward(input)
then result != undefined
then result == 0.0

test forward_batch_basic_case
given input = default_input()
test forward_positive_input
given input = 1.0
when result = forward(input)
then result > 0.0 and result < 1.0

test forward_negative_input
given input = -1.0
when result = forward(input)
then result < 0.0 and result > -1.0

test forward_batch_empty
given input = []f32{len = 0}
when result = forward_batch(input)
then result.len == 0

test forward_batch_single_element
given input = []f32{len = 1; [0.0]}
when result = forward_batch(input)
then result.len == 1 and result[0] == 0.0

test forward_batch_multiple_elements
given input = []f32{len = 3; [1.0, -1.0, 0.5]}
when result = forward_batch(input)
then result.len == 3 and result[0] > 0.0 and result[1] < 0.0 and result[2] > 0.0

test derivative_zero_input
given input = 0.0
when result = derivative(input)
then result > 0.0 and result < 1.0

test derivative_positive_input
given input = 1.0
when result = derivative(input)
then result > 0.0 and result < 1.5

test derivative_negative_input
given input = -1.0
when result = derivative(input)
then result > 0.0 and result < 1.0

test "forward_basic"
given input = 0.5
when result = forward(input)
then result >= 0.0

test "forward_batch_basic"
given input = []f32{len = 2; [1.0, -0.5]}
when result = forward_batch(input)
then result != undefined
then result.len == 2

test derivative_basic_case
given input = default_input()
test "derivative_basic"
given input = 2.0
when result = derivative(input)
then result != undefined
then result >= 0.0

Loading