diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index 08fef52..2b69fd1 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -422,13 +422,23 @@ class DPOStrategy(BaseStrategy): def compute_loss_output(self, batch: Dict[str, Tensor]) -> LossOutput: batch = move_to_device(batch, self.device) chosen_ids, rejected_ids = batch["chosen"], batch["rejected"] - chosen_mask, rejected_mask = batch["chosen_mask"], batch["rejected_mask"] + chosen_loss_mask = batch["chosen_mask"] + rejected_loss_mask = batch["rejected_mask"] + chosen_attention_mask = batch.get("chosen_attention_mask") + rejected_attention_mask = batch.get("rejected_attention_mask") + if chosen_attention_mask is None: + chosen_attention_mask = chosen_ids.ne(0) + if rejected_attention_mask is None: + rejected_attention_mask = rejected_ids.ne(0) concat_ids = torch.cat([chosen_ids, rejected_ids], dim=0) - concat_loss_mask = torch.cat([chosen_mask, rejected_mask], dim=0) + concat_loss_mask = torch.cat([chosen_loss_mask, rejected_loss_mask], dim=0) + concat_attention_mask = torch.cat( + [chosen_attention_mask, rejected_attention_mask], dim=0 + ) # Build full attention mask: key-padding + causal - key_pad = concat_ids.bool()[:, None, None, :] # [B*2, 1, 1, S] + key_pad = concat_attention_mask.bool()[:, None, None, :] S = key_pad.shape[-1] causal = torch.tril( torch.ones(S, S, dtype=torch.bool, device=concat_ids.device) @@ -477,23 +487,42 @@ class DPOStrategy(BaseStrategy): return True def prepare_from_rollout(self, result: RolloutResult) -> Dict[str, Tensor]: - """Pick best/worst response per prompt by reward as chosen/rejected.""" + """Build prompt-conditioned chosen/rejected sequences from rollout. + + DPO scores each response conditioned on its original prompt. The + prompt remains visible to attention while the loss mask covers only + valid response tokens. + """ rewards = result.rewards + prompts = result.prompts + prompt_mask = result.prompt_mask.bool() responses = result.responses - masks = result.response_mask + response_masks = result.response_mask.bool() best = rewards.argmax(dim=-1) worst = rewards.argmin(dim=-1) B = responses.shape[0] idx = torch.arange(B, device=responses.device) - chosen = responses[idx, best] - chosen_mask = masks[idx, best].float() - rejected = responses[idx, worst] - rejected_mask = masks[idx, worst].float() + chosen_response = responses[idx, best] + chosen_response_mask = response_masks[idx, best] + rejected_response = responses[idx, worst] + rejected_response_mask = response_masks[idx, worst] + + chosen = torch.cat([prompts, chosen_response], dim=-1) + rejected = torch.cat([prompts, rejected_response], dim=-1) + prompt_loss_mask = torch.zeros_like(prompt_mask) + chosen_mask = torch.cat([prompt_loss_mask, chosen_response_mask], dim=-1) + rejected_mask = torch.cat([prompt_loss_mask, rejected_response_mask], dim=-1) + chosen_attention_mask = torch.cat([prompt_mask, chosen_response_mask], dim=-1) + rejected_attention_mask = torch.cat( + [prompt_mask, rejected_response_mask], dim=-1 + ) return { "chosen": chosen, "chosen_mask": chosen_mask, + "chosen_attention_mask": chosen_attention_mask, "rejected": rejected, "rejected_mask": rejected_mask, + "rejected_attention_mask": rejected_attention_mask, } diff --git a/tests/trainer/test_online_strategy.py b/tests/trainer/test_online_strategy.py index 7960f70..6ac2984 100644 --- a/tests/trainer/test_online_strategy.py +++ b/tests/trainer/test_online_strategy.py @@ -136,19 +136,58 @@ def test_grpo_prepare_from_rollout_mapping(device): assert batch["rewards"] is r.rewards -def test_dpo_prepare_from_rollout_picks_best_worst(device): +def test_dpo_prepare_from_rollout_conditions_responses_on_prompt(device): strat = _make_dpo(device) - r = _make_rollout_result(B=3, G=4, R=5, device=device) + r = _make_rollout_result(B=3, G=4, P=6, R=5, device=device) + r.prompt_mask[0, :2] = False + r.prompts[0, :2] = 0 + r.response_mask[1, :, -2:] = False + r.responses[1, :, -2:] = 0 batch = strat.prepare_from_rollout(r) - assert batch["chosen"].shape == (3, 5) - assert batch["rejected"].shape == (3, 5) - assert batch["chosen_mask"].shape == (3, 5) - assert batch["rejected_mask"].shape == (3, 5) + + assert batch["chosen"].shape == (3, 11) + assert batch["rejected"].shape == (3, 11) + assert batch["chosen_mask"].shape == (3, 11) + assert batch["rejected_mask"].shape == (3, 11) idx = torch.arange(3, device=device) expected_best = r.responses[idx, r.rewards.argmax(dim=-1)] expected_worst = r.responses[idx, r.rewards.argmin(dim=-1)] - assert torch.equal(batch["chosen"], expected_best) - assert torch.equal(batch["rejected"], expected_worst) + expected_best_mask = r.response_mask[idx, r.rewards.argmax(dim=-1)] + expected_worst_mask = r.response_mask[idx, r.rewards.argmin(dim=-1)] + + assert torch.equal(batch["chosen"][:, :6], r.prompts) + assert torch.equal(batch["rejected"][:, :6], r.prompts) + assert torch.equal(batch["chosen"][:, 6:], expected_best) + assert torch.equal(batch["rejected"][:, 6:], expected_worst) + assert not batch["chosen_mask"][:, :6].any() + assert not batch["rejected_mask"][:, :6].any() + assert torch.equal(batch["chosen_mask"][:, 6:], expected_best_mask) + assert torch.equal(batch["rejected_mask"][:, 6:], expected_worst_mask) + assert torch.equal(batch["chosen_attention_mask"][:, :6], r.prompt_mask) + assert torch.equal(batch["rejected_attention_mask"][:, :6], r.prompt_mask) + assert torch.equal(batch["chosen_attention_mask"][:, 6:], expected_best_mask) + assert torch.equal(batch["rejected_attention_mask"][:, 6:], expected_worst_mask) + + +def test_dpo_prepare_from_rollout_same_response_keeps_distinct_prompts(): + strat = _make_dpo("cpu") + r = _make_rollout_result(B=2, G=2, P=3, R=2, device="cpu") + r.prompts = torch.tensor([[0, 11, 12], [21, 22, 23]]) + r.prompt_mask = torch.tensor([[False, True, True], [True, True, True]]) + shared_response = torch.tensor([101, 102]) + r.responses[:] = shared_response + r.response_mask[:] = True + r.rewards = torch.tensor([[1.0, 0.0], [1.0, 0.0]]) + + batch = strat.prepare_from_rollout(r) + + assert torch.equal(batch["chosen"][:, 3:], shared_response.expand(2, -1)) + assert torch.equal(batch["rejected"][:, 3:], shared_response.expand(2, -1)) + assert torch.equal(batch["chosen"][:, :3], r.prompts) + assert torch.equal(batch["rejected"][:, :3], r.prompts) + assert not torch.equal(batch["chosen"][0], batch["chosen"][1]) + assert not batch["chosen_mask"][:, :3].any() + assert not batch["rejected_mask"][:, :3].any() def test_call_without_runner_falls_back_to_compute_loss_grpo(device):