From caf4877aedb86928022a14614b83a92153b713e7 Mon Sep 17 00:00:00 2001 From: Trinity Bee Date: Tue, 15 Sep 2026 18:41:01 +0000 Subject: [PATCH] Implement DQN target network functions with proper signatures and tests - Add missing target parameter to hard_update function - Add missing target and tau parameters to soft_update function - Add missing config parameter to should_update function and change return type to bool - Implement hard_update to copy source weights to target - Implement soft_update with Polyak averaging: target = tau * source + (1 - tau) * target - Implement should_update to check if step matches update frequency - Update all test blocks to match new function signatures - Add proper test inputs using constants and default values Closes #3719 --- specs/ml/rl/dqn_target_network.t27 | 37 ++++++++++++++++++------------ 1 file changed, 22 insertions(+), 15 deletions(-) diff --git a/specs/ml/rl/dqn_target_network.t27 b/specs/ml/rl/dqn_target_network.t27 index c85f2cd2f8..82fa11747d 100644 --- a/specs/ml/rl/dqn_target_network.t27 +++ b/specs/ml/rl/dqn_target_network.t27 @@ -31,19 +31,26 @@ 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; } // ═══════════════════════════════════════════════════════════ @@ -51,18 +58,18 @@ module DqnTarget; // ═══════════════════════════════════════════════════════════ 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 // ═══════════════════════════════════════════════════════════