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
37 changes: 22 additions & 15 deletions specs/ml/rl/dqn_target_network.t27
Original file line number Diff line number Diff line change
Expand Up @@ -31,38 +31,45 @@ module DqnTarget;
// 3. Core Functions
// ═══════════════════════════════════════════════════════════

// hard_update(source: []f32) → void
fn hard_update(source: []f32) -> void {
// TODO: Implement from .tri spec
// hard_update(source: []f32, target: []f32) → void
fn hard_update(source: []f32, target: []f32) -> void {
// Copy source network weights to target network
for i in 0..source.length {
target[i] = source[i];
}
}

// soft_update(source: []f32) → void
fn soft_update(source: []f32) -> void {
// TODO: Implement from .tri spec
// soft_update(source: []f32, target: []f32, tau: f32) → void
fn soft_update(source: []f32, target: []f32, tau: f32) -> void {
// Polyak averaging: target = tau * source + (1 - tau) * target
for i in 0..source.length {
target[i] = tau * source[i] + (1.0 - tau) * target[i];
}
}

// should_update(step: u32) → void
fn should_update(step: u32) -> void {
// TODO: Implement from .tri spec
// should_update(step: u32, config: TargetUpdateConfig) → bool
fn should_update(step: u32, config: TargetUpdateConfig) -> bool {
// Update based on configured frequency and method
return step % config.update_freq == 0;
}

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

test hard_update_basic_case
given input = default_input()
when result = hard_update(input)
given source = default_input(), target = default_input()
when result = hard_update(source, target)
then result != undefined

test soft_update_basic_case
given input = default_input()
when result = soft_update(input)
given source = default_input(), target = default_input(), tau = DEFAULT_TAU
when result = soft_update(source, target, tau)
then result != undefined

test should_update_basic_case
given input = default_input()
when result = should_update(input)
given step = 1000u32, config = TargetUpdateConfig { method = UpdateMethod, tau = DEFAULT_TAU, update_freq = DEFAULT_UPDATE_FREQ }
when result = should_update(step, config)
then result != undefined

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