Skip to content
Open
Show file tree
Hide file tree
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
13 changes: 13 additions & 0 deletions docs/source/reference/llms.rst
Original file line number Diff line number Diff line change
Expand Up @@ -544,6 +544,19 @@ SFT
SFTLoss
SFTLossOutput

Reward Model Training
~~~~~~~~~~~~~~~~~~~~~

.. currentmodule:: torchrl.objectives.llm

.. autosummary::
:toctree: generated/
:template: rl_template.rst

RewardModelLoss
RewardModelLossOutput
reward_model_loss

.. currentmodule:: torchrl.data.llm

.. autosummary::
Expand Down
11 changes: 11 additions & 0 deletions docs/source/reference/llms_objectives.rst
Original file line number Diff line number Diff line change
Expand Up @@ -29,3 +29,14 @@ SFT

SFTLoss
SFTLossOutput

Reward Model Training
---------------------

.. autosummary::
:toctree: generated/
:template: rl_template.rst

RewardModelLoss
RewardModelLossOutput
reward_model_loss
158 changes: 158 additions & 0 deletions test/llm/test_llm_objectives.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
import torch

from tensordict import lazy_stack, MetaData, TensorDict
from tensordict.nn import TensorDictModule
from torchrl._utils import logger
from torchrl.data import History, LazyStackStorage, ReplayBuffer
from torchrl.data.llm.history import _CHAT_TEMPLATES
Expand All @@ -29,6 +30,11 @@
MCAdvantageSelector,
RayMCAdvantage,
)
from torchrl.objectives.llm.reward import (
reward_model_loss,
RewardModelLoss,
RewardModelLossOutput,
)
from torchrl.objectives.llm.sft import SFTLoss

_has_transformers = importlib.util.find_spec("transformers") is not None
Expand Down Expand Up @@ -1123,6 +1129,158 @@ def test_grpo_loss_with_real_models(
assert torch.isfinite(result.loss_objective)


class _Scorer(torch.nn.Module):
"""A tiny, dependency-free score network mapping input_ids -> scalar score."""

def __init__(self, vocab_size: int = 128, embed_dim: int = 8):
super().__init__()
self.embed = torch.nn.Embedding(vocab_size, embed_dim)
self.head = torch.nn.Linear(embed_dim, 1)

def forward(self, input_ids):
return self.head(self.embed(input_ids).float().mean(-2))


class TestRewardModel:
"""Tests for the model-agnostic Bradley-Terry :class:`RewardModelLoss`."""

vocab_size = 128
seq_len = 16
batch_size = 4

def _score_network(self):
return TensorDictModule(
_Scorer(self.vocab_size), in_keys=["input_ids"], out_keys=["score"]
)

def _make_data(self, chosen_key="chosen", rejected_key="rejected"):
def _ids():
return torch.randint(0, self.vocab_size, (self.batch_size, self.seq_len))

return TensorDict(
{
chosen_key: TensorDict(input_ids=_ids(), batch_size=[self.batch_size]),
rejected_key: TensorDict(
input_ids=_ids(), batch_size=[self.batch_size]
),
},
batch_size=[self.batch_size],
)

def test_reward_model_loss_fn_matches_formula(self):
"""The bare loss helper matches -log_sigmoid(chosen - rejected)."""
chosen = torch.randn(8)
rejected = torch.randn(8)
expected = -torch.nn.functional.logsigmoid(chosen - rejected)
torch.testing.assert_close(
reward_model_loss(chosen, rejected, "none"), expected
)
torch.testing.assert_close(
reward_model_loss(chosen, rejected, "mean"), expected.mean()
)
torch.testing.assert_close(
reward_model_loss(chosen, rejected, "sum"), expected.sum()
)

def test_forward_output_type_and_backward(self):
loss_fn = RewardModelLoss(score_network=self._score_network())
out = loss_fn(self._make_data())
assert isinstance(out, RewardModelLossOutput)
assert out.loss_reward_model.shape == ()
assert torch.isfinite(out.loss_reward_model)
assert out.loss_center is None
# accuracy is a detached metric in [0, 1]
assert 0.0 <= out.accuracy.item() <= 1.0
assert not out.accuracy.requires_grad
# gradients flow to the score network
out.loss_reward_model.backward()
grads = [
p.grad for p in loss_fn.score_network.parameters() if p.grad is not None
]
assert grads, "expected gradients to flow into the score network"

@pytest.mark.parametrize("reduction", ["mean", "sum", "none"])
def test_reduction(self, reduction):
loss_fn = RewardModelLoss(
score_network=self._score_network(), reduction=reduction
)
out = loss_fn(self._make_data())
if reduction == "none":
assert out.loss_reward_model.shape == (self.batch_size,)
else:
assert out.loss_reward_model.shape == ()

def test_precomputed_scores_no_network(self):
"""With score_network=None the scores are read directly from the inputs."""
loss_fn = RewardModelLoss(score_network=None)
chosen = torch.randn(self.batch_size)
rejected = torch.randn(self.batch_size)
td = TensorDict(
{
"chosen": TensorDict(score=chosen, batch_size=[self.batch_size]),
"rejected": TensorDict(score=rejected, batch_size=[self.batch_size]),
},
batch_size=[self.batch_size],
)
out = loss_fn(td)
expected = -torch.nn.functional.logsigmoid(chosen - rejected).mean()
torch.testing.assert_close(out.loss_reward_model, expected)
expected_acc = (chosen > rejected).float().mean()
torch.testing.assert_close(out.accuracy, expected_acc)

def test_center_coeff(self):
loss_fn = RewardModelLoss(score_network=self._score_network(), center_coeff=0.1)
out = loss_fn(self._make_data())
assert out.loss_center is not None
assert out.loss_center.requires_grad
# total differentiable loss is the sum of the loss_ terms
(out.loss_reward_model + out.loss_center).backward()

def test_nested_keys(self):
"""Exercise NestedKey inputs via set_keys."""
loss_fn = RewardModelLoss(score_network=self._score_network())
loss_fn.set_keys(chosen=("data", "chosen"), rejected=("data", "rejected"))
inner = self._make_data()
td = TensorDict({"data": inner}, batch_size=[self.batch_size])
out = loss_fn(td)
assert torch.isfinite(out.loss_reward_model)
assert ("data", "chosen") in loss_fn.in_keys

def test_nested_score_key(self):
"""Exercise NestedKey score lookup via set_keys."""
loss_fn = RewardModelLoss(score_network=None)
loss_fn.set_keys(score=("metrics", "score"))
chosen = torch.randn(self.batch_size)
rejected = torch.randn(self.batch_size)
td = TensorDict(
{
"chosen": TensorDict(
{"metrics": TensorDict(score=chosen, batch_size=[self.batch_size])},
batch_size=[self.batch_size],
),
"rejected": TensorDict(
{
"metrics": TensorDict(
score=rejected, batch_size=[self.batch_size]
)
},
batch_size=[self.batch_size],
),
},
batch_size=[self.batch_size],
)
out = loss_fn(td)
expected = -torch.nn.functional.logsigmoid(chosen - rejected).mean()
torch.testing.assert_close(out.loss_reward_model, expected)

def test_missing_key_raises(self):
loss_fn = RewardModelLoss(score_network=self._score_network())
td = self._make_data()
del td["rejected"]
with pytest.raises(KeyError):
loss_fn(td)


if __name__ == "__main__":
args, unknown = argparse.ArgumentParser().parse_known_args()
pytest.main([__file__, "--capture", "no", "--exitfirst"] + unknown)
4 changes: 4 additions & 0 deletions torchrl/objectives/llm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
MCAdvantageSelector,
RayMCAdvantage,
)
from .reward import reward_model_loss, RewardModelLoss, RewardModelLossOutput
from .sft import SFTLoss, SFTLossOutput

__all__ = [
Expand All @@ -29,6 +30,9 @@
"MCAdvantage",
"MCAdvantageSelector",
"RayMCAdvantage",
"RewardModelLoss",
"RewardModelLossOutput",
"reward_model_loss",
"SFTLoss",
"SFTLossOutput",
]
Loading
Loading