Files
modelscope--ms-swift/swift/rlhf_trainers/kto_trainer.py
T
wehub-resource-sync a203934033
Lint test / lint (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:34:58 +08:00

133 lines
5.8 KiB
Python

# Copyright (c) ModelScope Contributors. All rights reserved.
import torch
import torch.nn as nn
import trl
from packaging import version
from peft import PeftModel
from transformers import PreTrainedModel
from typing import Dict, Optional, Union
from swift.trainers import SwiftMixin, disable_gradient_checkpointing
from swift.utils import get_logger
from .rlhf_mixin import RLHFTrainerMixin
logger = get_logger()
if version.parse(trl.__version__) >= version.parse('0.26.0'):
from trl.experimental.kto import KTOTrainer as HFKTOTrainer
else:
from trl import KTOTrainer as HFKTOTrainer
del HFKTOTrainer.__init__
class KTOTrainer(RLHFTrainerMixin, SwiftMixin, HFKTOTrainer):
def __init__(self,
model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,
ref_model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,
*_args,
**kwargs):
args = kwargs['args']
args.disable_dropout = True
self.desirable_weight = args.desirable_weight
self.undesirable_weight = args.undesirable_weight
self.precompute_ref_log_probs = args.precompute_ref_log_probs
if hasattr(args, 'loss_type'):
self.loss_type = args.loss_type
else:
self.loss_type = 'kto'
self.ref_adapter_name = getattr(args, 'ref_adapter_name', None)
self.model_adapter_name = None
# Not all losses require a KL calculation
self.calculate_KL = True
if self.loss_type in ['apo_zero_unpaired']:
self.calculate_KL = False
super().__init__(model, ref_model, *_args, **kwargs)
# Code borrowed from huggingface/trl
def forward(
self, model: nn.Module, batch: Dict[str, Union[list, torch.LongTensor]]
) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:
KL_logps = self._compute_kl_logps(model, batch)
model_kwargs, labels = self._get_model_kwargs(batch, 'completion_')
if self.aux_loss_enabled:
model_kwargs['output_router_logits'] = True
outputs = model(**model_kwargs)
completion_logits = outputs.logits
completion_logps, completion_logits = self.get_batch_logps(model_kwargs, completion_logits, labels)
if completion_logps.shape[0] != len(batch['label']):
raise ValueError('There is a mismatch between the number of examples in this batch and the number of '
'examples for which an output sequence was predicted.')
chosen_idx = [i for i in range(completion_logps.shape[0]) if batch['label'][i] is True]
rejected_idx = [i for i in range(completion_logps.shape[0]) if batch['label'][i] is False]
chosen_logps = completion_logps[chosen_idx, ...]
rejected_logps = completion_logps[rejected_idx, ...]
chosen_logits = completion_logits[chosen_idx]
rejected_logits = completion_logits[rejected_idx]
if self.aux_loss_enabled:
return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, KL_logps, outputs.aux_loss)
else:
return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, KL_logps)
def _get_model_kwargs(self, inputs, prefix: str):
model_kwargs = {k[len(prefix):]: v for k, v in inputs.items() if k.startswith(prefix)}
use_logits_to_keep = self.get_use_logits_to_keep(self.template.sequence_parallel_size == 1)
if use_logits_to_keep:
self.prepare_logits_to_keep(model_kwargs)
labels = model_kwargs['labels']
if not self.is_encoder_decoder:
model_kwargs.pop('labels')
return model_kwargs, labels
def get_batch_logps(
self,
inputs,
logits: torch.FloatTensor,
labels: torch.LongTensor,
) -> torch.FloatTensor:
text_position_ids = inputs.pop('text_position_ids', None)
if text_position_ids is None:
text_position_ids = inputs.get('position_ids')
if logits.shape[1] != labels.shape[1]:
# for llava, the model returns logits for the entire sequence, including the image tokens
# (placed before the text tokens)
logits = logits[:, -labels.shape[1]:]
if not self.is_encoder_decoder and self.template.sequence_parallel_size == 1:
# Shift so that tokens < n predict n
labels = torch.roll(labels, shifts=-1, dims=1)
per_token_logps, sum_logits, loss_mask = self.get_per_token_logps(
logits, labels, label_pad_token_id=self.label_pad_token_id, reduction='sum')
if self.template.padding_free:
cu_seqlens = self.get_cu_seqlens(text_position_ids, inputs.get('logits_to_keep'))
all_logps = per_token_logps.new_zeros(cu_seqlens.shape[0] - 1)
all_logits = per_token_logps.new_zeros(cu_seqlens.shape[0] - 1)
for i in range(cu_seqlens.shape[0] - 1):
start, end = cu_seqlens[i], cu_seqlens[i + 1]
all_logps[i] = per_token_logps[:, start:end].sum()
all_logits[i] = sum_logits[:, start:end].sum()
else:
all_logps = per_token_logps.sum(-1)
all_logits = sum_logits.sum(-1)
return all_logps, all_logits
# Code borrowed from huggingface/trl (compat trl<0.17)
def _compute_kl_logps(self, model, batch):
"""Compute KL log probabilities for a given batch."""
KL_logps = None
if self.calculate_KL:
KL_model_kwargs, labels = self._get_model_kwargs(batch, 'KL_completion_')
with torch.no_grad(), disable_gradient_checkpointing(model, self.args.gradient_checkpointing_kwargs):
KL_logits = model(**KL_model_kwargs).logits
KL_logps, _ = self.get_batch_logps(KL_model_kwargs, KL_logits, labels)
return KL_logps