From 720adc5beb505aafe6ffa32f604a0bd42f6a27c4 Mon Sep 17 00:00:00 2001 From: leuitong Date: Fri, 7 Aug 2026 19:56:16 +0800 Subject: [PATCH 1/3] 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') From 31e4fad85ecf559a5fa79a0c92feb8d4b86e3669 Mon Sep 17 00:00:00 2001 From: leuitong Date: Mon, 10 Aug 2026 10:39:38 +0800 Subject: [PATCH 2/3] fix isort --- swift/rlhf_trainers/ppo_trainer.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/swift/rlhf_trainers/ppo_trainer.py b/swift/rlhf_trainers/ppo_trainer.py index 1e122ab7a0..661e9f38d1 100644 --- a/swift/rlhf_trainers/ppo_trainer.py +++ b/swift/rlhf_trainers/ppo_trainer.py @@ -1,16 +1,18 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import inspect -import trl from contextlib import contextmanager +from typing import Optional + +import trl from packaging import version from torch.utils.data import DataLoader from transformers import PreTrainedModel from transformers import Trainer as HfTrainer -from typing import Optional +from swift.callbacks import callbacks_map 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 From 5b0097545beaa6f0021ae7be623398cf25a41ed1 Mon Sep 17 00:00:00 2001 From: leuitong Date: Mon, 10 Aug 2026 17:16:17 +0800 Subject: [PATCH 3/3] fix isort --- swift/rlhf_trainers/ppo_trainer.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/swift/rlhf_trainers/ppo_trainer.py b/swift/rlhf_trainers/ppo_trainer.py index 661e9f38d1..bf3f4d8577 100644 --- a/swift/rlhf_trainers/ppo_trainer.py +++ b/swift/rlhf_trainers/ppo_trainer.py @@ -1,13 +1,12 @@ # Copyright (c) ModelScope Contributors. All rights reserved. import inspect -from contextlib import contextmanager -from typing import Optional - import trl +from contextlib import contextmanager from packaging import version from torch.utils.data import DataLoader from transformers import PreTrainedModel from transformers import Trainer as HfTrainer +from typing import Optional from swift.callbacks import callbacks_map from swift.trainers import SwiftMixin