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
59 changes: 50 additions & 9 deletions specs/ml/layers/dropout_layer.t27
Original file line number Diff line number Diff line change
Expand Up @@ -26,14 +26,49 @@ module Dropout;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// forward(input: []const f32) → void
fn forward(input: []const f32) -> void {
// TODO: Implement from .tri spec
// forward(input: []const f32, output: []f32, mask: []bool, training: bool, config: DropoutConfig) → void
fn forward(input: []const f32, output: []f32, mask: []bool, training: bool, config: DropoutConfig) -> void {
const n = input.len;
var i : usize = 0;
while (i < n) {
if (training) {
if (mask[i]) {
// Neuron is kept, scale during training if configured
if (config.scale_during_training) {
output[i] = input[i] * (1.0 / (1.0 - config.p));
} else {
output[i] = input[i];
}
} else {
// Neuron is dropped
output[i] = 0.0;
}
} else {
// During inference, always scale output
output[i] = input[i] * (1.0 - config.p);
}
i = i + 1;
}
}

// backward(grad_output: []const f32) → void
fn backward(grad_output: []const f32) -> void {
// TODO: Implement from .tri spec
// backward(grad_output: []const f32, mask: []const bool, grad_input: []f32, config: DropoutConfig) → void
fn backward(grad_output: []const f32, mask: []const bool, grad_input: []f32, config: DropoutConfig) -> void {
const n = grad_output.len;
var i : usize = 0;
while (i < n) {
if (mask[i]) {
// Gradient flows through kept neurons
if (config.scale_during_training) {
grad_input[i] = grad_output[i] * (1.0 / (1.0 - config.p));
} else {
grad_input[i] = grad_output[i];
}
} else {
// Gradient is zero for dropped neurons
grad_input[i] = 0.0;
}
i = i + 1;
}
}

// ═══════════════════════════════════════════════════════════
Expand All @@ -42,11 +77,17 @@ module Dropout;

test forward_basic_case
given input = default_input()
when result = forward(input)
output = allocf32(input.len)
mask = allocbool(input.len)
config = DropoutConfig { p = 0.5, inplace = false, scale_during_training = true }
when result = forward(input, output, mask, true, config)
then result != undefined

test backward_basic_case
given input = default_input()
when result = backward(input)
given grad_output = default_input()
mask = allocbool(grad_output.len)
grad_input = allocf32(grad_output.len)
config = DropoutConfig { p = 0.5, inplace = false, scale_during_training = true }
when result = backward(grad_output, mask, grad_input, config)
then result != undefined

Loading