From 720adc5beb505aafe6ffa32f604a0bd42f6a27c4 Mon Sep 17 00:00:00 2001 From: leuitong Date: Fri, 7 Aug 2026 19:56:16 +0800 Subject: [PATCH] fix ppo callbacks This fix restores user callbacks on the PPO trainer that TRL's PPOTrainer.__init__ silently discards. --- swift/rlhf_trainers/ppo_trainer.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/swift/rlhf_trainers/ppo_trainer.py b/swift/rlhf_trainers/ppo_trainer.py index 3f2be52e23..1e122ab7a0 100644 --- a/swift/rlhf_trainers/ppo_trainer.py +++ b/swift/rlhf_trainers/ppo_trainer.py @@ -10,6 +10,7 @@ from swift.trainers import SwiftMixin from swift.utils import patch_getattr +from swift.callbacks import callbacks_map if version.parse(trl.__version__) >= version.parse('0.26.0'): from trl.experimental.ppo import PPOTrainer as HFPPOTrainer @@ -49,7 +50,6 @@ def __init__(self, model: PreTrainedModel, ref_model: PreTrainedModel, *_args, * 'reward_model', 'value_model', 'eval_dataset', - 'callbacks', ] } parameters = inspect.signature(ppo_trainer_init).parameters @@ -62,6 +62,12 @@ def __init__(self, model: PreTrainedModel, ref_model: PreTrainedModel, *_args, * else: new_kwargs['tokenizer'] = self.tokenizer ppo_trainer_init(self, model=model, ref_model=ref_model, **new_kwargs) + # TRL's `PPOTrainer.__init__` rebuilds `callback_handler`, discarding the user callbacks that + # `SwiftMixin` registered on the HF Trainer. Re-register them onto the new handler here. We + # call `add_callback` directly (rather than `SwiftMixin._add_callbacks`/`_get_callbacks`, + # whose names/signatures differ between swift branches) so this works across versions. + for callback in self.args.callbacks: + self.add_callback(callbacks_map[callback](self.args, self)) unwrap_model = self.accelerator.unwrap_model(self.model) patch_getattr(unwrap_model.__class__, 'policy')