diff --git a/swift/megatron/trainers/reward_trainer.py b/swift/megatron/trainers/reward_trainer.py index 1f7d47d438..913ad5fff9 100644 --- a/swift/megatron/trainers/reward_trainer.py +++ b/swift/megatron/trainers/reward_trainer.py @@ -16,8 +16,13 @@ def loss_func(self, output_tensor, *, data): margin = data.pop('margin', None) num_samples = output_tensor.shape[0] if packed_seq_params is None else packed_seq_params.seq_lens.shape[0] rewards = self.get_last_tokens(output_tensor, packed_seq_params, data.get('attention_mask')) - rewards_chosen, rewards_rejected = torch.split(rewards, num_samples // 2, dim=0) + batch_size = num_samples // 2 + rewards_chosen, rewards_rejected = torch.split(rewards, batch_size, dim=0) if margin is not None: + margin = margin.to(device=rewards_chosen.device, dtype=rewards_chosen.dtype) + if margin.numel() != batch_size: + raise ValueError(f'Expected {batch_size} margins, got {margin.numel()}.') + margin = margin.reshape_as(rewards_chosen) loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - margin).mean() else: loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected).mean() diff --git a/swift/rlhf_trainers/reward_trainer.py b/swift/rlhf_trainers/reward_trainer.py index b69d3ad2b9..31fd7fab0e 100644 --- a/swift/rlhf_trainers/reward_trainer.py +++ b/swift/rlhf_trainers/reward_trainer.py @@ -46,6 +46,10 @@ def compute_loss(self, rewards = model(**inputs).logits rewards_chosen, rewards_rejected = torch.split(rewards, batch_size, dim=0) if margin is not None: + margin = margin.to(device=rewards_chosen.device, dtype=rewards_chosen.dtype) + if margin.numel() != batch_size: + raise ValueError(f'Expected {batch_size} margins, got {margin.numel()}.') + margin = margin.reshape_as(rewards_chosen) loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected - margin).mean() else: loss = -nn.functional.logsigmoid(rewards_chosen - rewards_rejected).mean() diff --git a/swift/template/base.py b/swift/template/base.py index 3df77b97be..aeb4c82b1e 100644 --- a/swift/template/base.py +++ b/swift/template/base.py @@ -1788,8 +1788,9 @@ def _rlhf_data_collator(self, res = self._data_collator(new_batch, padding_to=padding_to) # reward modeling - margin = [b['margin'] for b in batch if b.get('margin') is not None] - if margin: + has_margin = any(b.get('margin') is not None for b in batch) + if has_margin: + margin = [0.0 if b.get('margin') is None else b['margin'] for b in batch] res['margin'] = torch.tensor(margin, dtype=torch.float) return res