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
12 changes: 6 additions & 6 deletions specs/ml/activation/relu_activation.t27
Original file line number Diff line number Diff line change
Expand Up @@ -59,40 +59,40 @@ module Relu;
// TDD: Tests (from .tri behaviors)
// ═══════════════════════════════════════════════════════════

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

test forward_positive_identity
test "forward_positive_identity"
given input = [1.0, 2.0, 3.0]
when output = forward(input, ReLUConfig{ .negative_slope = 0.0, .inplace = false })
then output[0] == 1.0
and output[1] == 2.0
and output[2] == 3.0

test forward_negative_scaling
test "forward_negative_scaling"
given input = [-1.0, -2.0, -3.0]
when output = forward(input, ReLUConfig{ .negative_slope = 0.5, .inplace = false })
then output[0] == -0.5
and output[1] == -1.0
and output[2] == -1.5

test forward_zero_negative_slope
test "forward_zero_negative_slope"
given input = [-1.0, 0.0, 1.0]
when output = forward(input, ReLUConfig{ .negative_slope = 0.0, .inplace = false })
then output[0] == 0.0
and output[1] == 0.0
and output[2] == 1.0

test backward_input_gate
test "backward_input_gate"
given input = [-1.0, 0.0, 1.0]
when grad_input = backward([1.0, 1.0, 1.0], input, ReLUConfig{ .negative_slope = 0.5, .inplace = false })
then grad_input[0] == 0.5 // grad_output * negative_slope for input <= ZERO
and grad_input[1] == 0.5 // grad_output * negative_slope for input == ZERO
and grad_input[2] == 1.0 // grad_output for input > ZERO

test backward_basic_case
test "backward_basic_case"
given input = default_input()
when result = backward(input)
then result != undefined
Expand Down
Loading