diff --git a/specs/ml/rl/ppo_clip_loss.t27 b/specs/ml/rl/ppo_clip_loss.t27 index b6d21fc15b..637f851d91 100644 --- a/specs/ml/rl/ppo_clip_loss.t27 +++ b/specs/ml/rl/ppo_clip_loss.t27 @@ -36,34 +36,74 @@ module PpoClipLoss; // 3. Core Functions // ═══════════════════════════════════════════════════════════ - // compute_ratio(new_log_prob: f32) → void - fn compute_ratio(new_log_prob: f32) -> void { - // TODO: Implement from .tri spec + // compute_ratio(new_log_prob: f32, old_log_prob: f32) → f32 + fn compute_ratio(new_log_prob: f32, old_log_prob: f32) -> f32 { + // Compute probability ratio between new and old policies + return new_log_prob - old_log_prob; } - // clipped_surrogate(ratio: f32) → void - fn clipped_surrogate(ratio: f32) -> void { - // TODO: Implement from .tri spec + // clipped_surrogate(ratio: f32, advantage: f32, clip_epsilon: f32) → f32 + fn clipped_surrogate(ratio: f32, advantage: f32, clip_epsilon: f32) -> f32 { + // Apply clipping to the surrogate objective + let clipped_ratio = if ratio > 1.0 + clip_epsilon { + 1.0 + clip_epsilon + } else if ratio < 1.0 - clip_epsilon { + 1.0 - clip_epsilon + } else { + ratio + }; + return clipped_ratio * advantage; } - // policy_loss(ratios: []f32) → void - fn policy_loss(ratios: []f32) -> void { - // TODO: Implement from .tri spec + // policy_loss(ratios: []f32, advantages: []f32, clip_epsilon: f32) → f32 + fn policy_loss(ratios: []f32, advantages: []f32, clip_epsilon: f32) -> f32 { + // Compute policy loss using clipped surrogate objective + let total_loss = 0.0; + for i in 0..ratios.len { + let ratio = ratios[i]; + let advantage = advantages[i]; + let clipped_loss = clipped_surrogate(ratio, advantage, clip_epsilon); + total_loss += clipped_loss; + } + return total_loss / f32(ratios.len); } - // entropy_loss(entropies: []f32) → void - fn entropy_loss(entropies: []f32) -> void { - // TODO: Implement from .tri spec + // entropy_loss(entropies: []f32, entropy_coef: f32) → f32 + fn entropy_loss(entropies: []f32, entropy_coef: f32) -> f32 { + // Compute entropy loss to encourage exploration + let total_entropy = 0.0; + for entropy in entropies { + total_entropy += entropy; + } + return -entropy_coef * (total_entropy / f32(entropies.len)); } - // value_loss(predicted_values: []f32) → void - fn value_loss(predicted_values: []f32) -> void { - // TODO: Implement from .tri spec + // value_loss(predicted_values: []f32, target_values: []f32, old_values: []f32, clip_coef: f32) → f32 + fn value_loss(predicted_values: []f32, target_values: []f32, old_values: []f32, clip_coef: f32) -> f32 { + // Compute value loss using clipped value function + let total_loss = 0.0; + for i in 0..predicted_values.len { + let predicted = predicted_values[i]; + let target = target_values[i]; + let old = old_values[i]; + let value_diff = predicted - old; + let clipped_value_diff = if value_diff > clip_coef { + clip_coef + } else if value_diff < -clip_coef { + -clip_coef + } else { + value_diff + }; + let value_loss = 0.5 * (target - (old + clipped_value_diff)) * (target - (old + clipped_value_diff)); + total_loss += value_loss; + } + return total_loss / f32(predicted_values.len); } - // total_loss(policy_loss_value: f32) → void - fn total_loss(policy_loss_value: f32) -> void { - // TODO: Implement from .tri spec + // total_loss(policy_loss_value: f32, value_loss_value: f32, entropy_loss_value: f32, value_loss_coef: f32) → f32 + fn total_loss(policy_loss_value: f32, value_loss_value: f32, entropy_loss_value: f32, value_loss_coef: f32) -> f32 { + // Combine all loss components + return policy_loss_value + value_loss_coef * value_loss_value + entropy_loss_value; } // ═══════════════════════════════════════════════════════════ @@ -71,34 +111,60 @@ module PpoClipLoss; // ═══════════════════════════════════════════════════════════ test compute_ratio_basic_case - given input = default_input() - when result = compute_ratio(input) - then result != undefined + given + new_log_prob = 0.5 + old_log_prob = 0.3 + when result = compute_ratio(new_log_prob, old_log_prob) + then result == 0.2 test clipped_surrogate_basic_case - given input = default_input() - when result = clipped_surrogate(input) - then result != undefined + given + ratio = 1.1 + advantage = 2.0 + clip_epsilon = 0.2 + when result = clipped_surrogate(ratio, advantage, clip_epsilon) + then result == 1.2 * 2.0 + + test clipped_surrogate_clipping_case + given + ratio = 1.5 // Above clip threshold of 1.2 + advantage = 2.0 + clip_epsilon = 0.2 + when result = clipped_surrogate(ratio, advantage, clip_epsilon) + then result == 1.2 * 2.0 test policy_loss_basic_case - given input = default_input() - when result = policy_loss(input) - then result != undefined + given + ratios = [1.0, 1.1, 0.9] + advantages = [2.0, 1.5, 3.0] + clip_epsilon = 0.2 + when result = policy_loss(ratios, advantages, clip_epsilon) + then result > 0.0 test entropy_loss_basic_case - given input = default_input() - when result = entropy_loss(input) - then result != undefined + given + entropies = [0.5, 0.6, 0.4] + entropy_coef = 0.01 + when result = entropy_loss(entropies, entropy_coef) + then result < 0.0 test value_loss_basic_case - given input = default_input() - when result = value_loss(input) - then result != undefined + given + predicted_values = [1.0, 2.0, 1.5] + target_values = [1.1, 2.1, 1.6] + old_values = [0.9, 1.9, 1.4] + clip_coef = 0.2 + when result = value_loss(predicted_values, target_values, old_values, clip_coef) + then result > 0.0 test total_loss_basic_case - given input = default_input() - when result = total_loss(input) - then result != undefined + given + policy_loss_value = 1.0 + value_loss_value = 0.5 + entropy_loss_value = -0.01 + value_loss_coef = 0.5 + when result = total_loss(policy_loss_value, value_loss_value, entropy_loss_value, value_loss_coef) + then result == 1.0 + 0.5 * 0.5 + (-0.01) // ═══════════════════════════════════════════════════════════ // TDD: Invariants (from .tri constraints)