fugu-lite / tests /test_training_objectives.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
1.18 kB
import torch
from fugu_lite.train_rl import contextual_bandit_loss
from fugu_lite.train_sft import soft_label_loss
def test_sft_loss_prefers_high_reward_worker():
rewards = torch.tensor([[1.0, 0.0]])
good_logits = torch.tensor([[4.0, -4.0]])
bad_logits = torch.tensor([[-4.0, 4.0]])
assert soft_label_loss(good_logits, rewards, 0.1) < soft_label_loss(
bad_logits, rewards, 0.1
)
def test_expected_reward_has_useful_gradient():
logits = torch.zeros((1, 2), requires_grad=True)
rewards = torch.tensor([[1.0, 0.0]])
loss, metrics = contextual_bandit_loss(
logits,
rewards,
estimator="expected_reward",
entropy_coefficient=0.0,
)
loss.backward()
assert logits.grad[0, 0] < 0
assert logits.grad[0, 1] > 0
assert metrics["zero_spread_fraction"] == 0.0
def test_reinforce_loss_is_finite_with_equal_rewards():
torch.manual_seed(0)
logits = torch.zeros((2, 3), requires_grad=True)
rewards = torch.ones((2, 3))
loss, metrics = contextual_bandit_loss(logits, rewards, estimator="reinforce")
assert torch.isfinite(loss)
assert metrics["zero_spread_fraction"] == 1.0