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
138 changes: 102 additions & 36 deletions specs/ml/rl/ppo_clip_loss.t27
Original file line number Diff line number Diff line change
Expand Up @@ -36,69 +36,135 @@ 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;
}

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

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)
Expand Down
Loading