diff --git a/trl/experimental/async_distillation/async_distillation_trainer.py b/trl/experimental/async_distillation/async_distillation_trainer.py index 67da9f959f2..4b522b0a9d3 100644 --- a/trl/experimental/async_distillation/async_distillation_trainer.py +++ b/trl/experimental/async_distillation/async_distillation_trainer.py @@ -1189,6 +1189,8 @@ class AsyncDistillationTrainer(_BaseTrainer): `rollout_worker` updates the policy itself). """ + loss_is_scaled_for_ga = True + _tag_names = ["trl", "async-distillation"] _name = "AsyncDistillation" _paper = { @@ -1317,12 +1319,7 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - compute_loss_func="non-None value to disable scaling", ) - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether - # the model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False precision = self.accelerator.mixed_precision dtype = { diff --git a/trl/experimental/async_grpo/async_grpo_trainer.py b/trl/experimental/async_grpo/async_grpo_trainer.py index 9feaf11856c..0edc018e5e0 100644 --- a/trl/experimental/async_grpo/async_grpo_trainer.py +++ b/trl/experimental/async_grpo/async_grpo_trainer.py @@ -1024,6 +1024,8 @@ class AsyncGRPOTrainer(_BaseTrainer): implementation to disable trainer-side weight sync. """ + loss_is_scaled_for_ga = True + _tag_names = ["trl", "async-grpo"] _name = "AsyncGRPO" _paper = { @@ -1210,12 +1212,7 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - compute_loss_func="non-None value to disable scaling", ) - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False precision = self.accelerator.mixed_precision dtype = { diff --git a/trl/experimental/cpo/cpo_trainer.py b/trl/experimental/cpo/cpo_trainer.py index ec0b5e7cba8..447b4b9c96d 100644 --- a/trl/experimental/cpo/cpo_trainer.py +++ b/trl/experimental/cpo/cpo_trainer.py @@ -117,6 +117,8 @@ class CPOTrainer(_BaseTrainer): metric values. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "cpo"] _name = "CPO" _paper = { @@ -413,11 +415,6 @@ def make_inputs_require_grad(module, input, output): preprocess_logits_for_metrics=preprocess_logits_for_metrics, ) - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - # Add tags for models that have been loaded with the correct transformers version if hasattr(self.model, "add_model_tags"): self.model.add_model_tags(self._tag_names) diff --git a/trl/experimental/orpo/orpo_trainer.py b/trl/experimental/orpo/orpo_trainer.py index b1555bc4266..e0e3c9cbb78 100644 --- a/trl/experimental/orpo/orpo_trainer.py +++ b/trl/experimental/orpo/orpo_trainer.py @@ -129,6 +129,8 @@ class ORPOTrainer(_BaseTrainer): metric values. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "orpo"] _name = "ORPO" _paper = { @@ -395,11 +397,6 @@ def make_inputs_require_grad(module, input, output): preprocess_logits_for_metrics=preprocess_logits_for_metrics, ) - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - # Add tags for models that have been loaded with the correct transformers version if hasattr(self.model, "add_model_tags"): self.model.add_model_tags(self._tag_names) diff --git a/trl/experimental/sdft/sdft_trainer.py b/trl/experimental/sdft/sdft_trainer.py index fd904438d87..c3650bf88ad 100644 --- a/trl/experimental/sdft/sdft_trainer.py +++ b/trl/experimental/sdft/sdft_trainer.py @@ -206,6 +206,8 @@ def build( class SDFTTrainer(_BaseTrainer): """Trainer for SDFT-style on-policy self-distillation with explicit teacher prompts.""" + loss_is_scaled_for_ga = True + _tag_names = ["trl", "sdft"] _name = "SDFT" config_cls = SDFTConfig @@ -393,7 +395,6 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - compute_loss_func="non-None value to disable scaling", ) self._last_loaded_step = -1 if self.use_vllm else 0 @@ -446,7 +447,6 @@ def __init__( self.model.add_model_tags(self._tag_names) self._setup_teacher_model() - self.model_accepts_loss_kwargs = False def _set_signature_columns_if_needed(self): if self._signature_columns is None: diff --git a/trl/experimental/sdpo/sdpo_trainer.py b/trl/experimental/sdpo/sdpo_trainer.py index 98f71ffb0b4..66ea0aef4dd 100644 --- a/trl/experimental/sdpo/sdpo_trainer.py +++ b/trl/experimental/sdpo/sdpo_trainer.py @@ -329,6 +329,8 @@ class SDPOTrainer(_BaseTrainer): next-token predictions back into the policy. """ + loss_is_scaled_for_ga = True + config_cls = SDPOConfig _tag_names = ["trl", "sdpo"] _name = "SDPO" @@ -521,7 +523,6 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - compute_loss_func="non-None value to disable scaling", ) self._last_loaded_step = -1 if self.use_vllm else 0 @@ -574,7 +575,6 @@ def __init__( self.model.add_model_tags(self._tag_names) self._setup_teacher_model() - self.model_accepts_loss_kwargs = False self.importance_sampling_level = args.importance_sampling_level self.scale_rewards = args.scale_rewards diff --git a/trl/experimental/server_distillation/server_distillation_trainer.py b/trl/experimental/server_distillation/server_distillation_trainer.py index 645efe39713..262539565de 100644 --- a/trl/experimental/server_distillation/server_distillation_trainer.py +++ b/trl/experimental/server_distillation/server_distillation_trainer.py @@ -336,9 +336,9 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N labels=trimmed_labels, ) - # The base trainer disables Trainer's built-in grad-accum loss scaling (via `compute_loss_func`) because it - # normalizes by the global completion-token count. The server path normalizes locally with `batchmean` and does - # not consume `num_items_in_batch`, so it must re-apply that scaling itself. + # The base trainer sets `loss_is_scaled_for_ga = True` because it normalizes by the global completion-token + # count. The server path normalizes locally with `batchmean` and does not consume `num_items_in_batch`, so it + # must re-apply that scaling itself. if self.model.training: loss = loss / self.current_gradient_accumulation_steps diff --git a/trl/experimental/ssd/ssd_trainer.py b/trl/experimental/ssd/ssd_trainer.py index ac8c6c9dcf3..d3b32938da6 100644 --- a/trl/experimental/ssd/ssd_trainer.py +++ b/trl/experimental/ssd/ssd_trainer.py @@ -79,6 +79,8 @@ class SSDTrainer(_BaseTrainer): ``prompt`` column. """ + loss_is_scaled_for_ga = True + _tag_names = ["trl", "ssd"] _name = "SSD" config_cls = SSDConfig @@ -237,7 +239,6 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - compute_loss_func="non-None value to disable scaling", ) if args.disable_dropout: @@ -245,8 +246,6 @@ def __init__( self.model.add_model_tags(self._tag_names) - self.model_accepts_loss_kwargs = False - if self.use_vllm: from ...generation.vllm_generation import VLLMGeneration diff --git a/trl/trainer/base_trainer.py b/trl/trainer/base_trainer.py index f75a2e8a684..5ca9ff0fa6f 100644 --- a/trl/trainer/base_trainer.py +++ b/trl/trainer/base_trainer.py @@ -16,9 +16,11 @@ from pathlib import Path import torch +import transformers from accelerate.utils import is_peft_model from datasets import Dataset from huggingface_hub.utils import send_telemetry +from packaging.version import Version from transformers import CONFIG_MAPPING, Trainer, is_wandb_available from .. import __version__ @@ -61,6 +63,9 @@ class _BaseTrainer(Trainer): + # Whether `compute_loss` already scales the loss for gradient accumulation, see `Trainer.loss_is_scaled_for_ga` + loss_is_scaled_for_ga = None + _tag_names = [] _name = "Base" _paper = {} @@ -68,6 +73,11 @@ class _BaseTrainer(Trainer): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) + # `Trainer.loss_is_scaled_for_ga` requires transformers 5.19; older versions only read these two attributes + if Version(transformers.__version__) < Version("5.19.0.dev0") and self.loss_is_scaled_for_ga is not None: + self.model_accepts_loss_kwargs = False + if self.loss_is_scaled_for_ga: + self.compute_loss_func = "non-None value to disable scaling" self._send_telemetry() def _send_telemetry(self): diff --git a/trl/trainer/distillation_trainer.py b/trl/trainer/distillation_trainer.py index a922755f130..6693a025d04 100644 --- a/trl/trainer/distillation_trainer.py +++ b/trl/trainer/distillation_trainer.py @@ -376,6 +376,8 @@ class DistillationTrainer(_BaseTrainer): use and that it has been fine-tuned for tool calling. """ + loss_is_scaled_for_ga = True + _tag_names = ["trl", "distillation"] _name = "Distillation" _paper = { @@ -715,12 +717,6 @@ def __init__( processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - # In Trainer, `training_step` scales the loss by `gradient_accumulation_steps` only if `compute_loss_func` - # is None. Here, loss scaling instead depends on the total number of completion tokens across the global - # accumulated batch. To control scaling ourselves, we must disable Trainer's built-in scaling. The simplest - # (though a bit hacky) way is to set `compute_loss_func` to any non-None value, which bypasses that behavior - # without rewriting `training_step`. - compute_loss_func="non-None value to disable scaling", ) # With several GPUs visible and no distributed launcher, `Trainer` wraps the model in `nn.DataParallel`, whose @@ -733,11 +729,6 @@ def __init__( "single GPU visible with `CUDA_VISIBLE_DEVICES`." ) - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - self._dist = DistributedBackend(self.accelerator) # Add tags to the model diff --git a/trl/trainer/dpo_trainer.py b/trl/trainer/dpo_trainer.py index 09473806e92..dfa8c498e3b 100644 --- a/trl/trainer/dpo_trainer.py +++ b/trl/trainer/dpo_trainer.py @@ -492,6 +492,8 @@ class DPOTrainer(_BaseTrainer): PEFT configuration used to wrap the model. If `None`, the model is not wrapped. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "dpo"] _name = "DPO" _paper = { @@ -947,11 +949,6 @@ def __init__( else: self._tp_size = 1 - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - # Add tags to the model self.model.add_model_tags(self._tag_names) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index a9aa3423182..4a9391bc4a3 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -282,6 +282,8 @@ class GRPOTrainer(_BaseTrainer): any time without prior notice. """ + loss_is_scaled_for_ga = True + _tag_names = ["trl", "grpo"] _name = "GRPO" _paper = { @@ -942,12 +944,6 @@ def get_reward(environments, _env_type=env_type, **kwargs): processing_class=processing_class, callbacks=callbacks, optimizers=optimizers, - # In Trainer, `training_step` scales the loss by `gradient_accumulation_steps` only if `compute_loss_func` - # is None. For DAPO, loss scaling instead depends on the total number of completions tokens across the - # global accumulated batch. To control scaling ourselves, we must disable Trainer's built-in scaling. The - # simplest (though a bit hacky) way is to set `compute_loss_func` to any non-None value, which bypasses - # that behavior without rewriting `training_step`. - compute_loss_func="non-None value to disable scaling", ) # With several GPUs visible and no distributed launcher, `Trainer` wraps the model in `nn.DataParallel`, whose @@ -1136,10 +1132,6 @@ def cast_outputs_to_original_dtype(module, args, output): # Keep training-specific generation kwargs to overwrite model's original generation config self.generation_kwargs = generation_kwargs - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False self._dist = DistributedBackend(self.accelerator) # Add tags to the model diff --git a/trl/trainer/kto_trainer.py b/trl/trainer/kto_trainer.py index 8e08aae6064..386b549eaa7 100644 --- a/trl/trainer/kto_trainer.py +++ b/trl/trainer/kto_trainer.py @@ -548,6 +548,8 @@ class KTOTrainer(_BaseTrainer): PEFT configuration used to wrap the model. If `None`, the model is not wrapped. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "kto"] _name = "KTO" _paper = { @@ -957,11 +959,6 @@ def __init__( else: self._tp_size = 1 - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - # Add tags to the model self.model.add_model_tags(self._tag_names) diff --git a/trl/trainer/reward_trainer.py b/trl/trainer/reward_trainer.py index 316411e7434..e55627e02d0 100644 --- a/trl/trainer/reward_trainer.py +++ b/trl/trainer/reward_trainer.py @@ -327,6 +327,8 @@ class RewardTrainer(_BaseTrainer): to ensure that the reward head is properly trained. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "reward-trainer"] _name = "Reward" _template_file = "rm_model_card.md" @@ -619,11 +621,6 @@ def __init__( else: self._tp_size = 1 - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False - # Add tags to the model self.model.add_model_tags(self._tag_names) diff --git a/trl/trainer/rloo_trainer.py b/trl/trainer/rloo_trainer.py index 4751fb6c40b..3d7603e851d 100644 --- a/trl/trainer/rloo_trainer.py +++ b/trl/trainer/rloo_trainer.py @@ -216,6 +216,8 @@ class RLOOTrainer(_BaseTrainer): PEFT configuration used to wrap the model. If `None`, the model is not wrapped. """ + loss_is_scaled_for_ga = False + _tag_names = ["trl", "rloo"] _name = "RLOO" _paper = { @@ -773,10 +775,6 @@ def __init__( # Keep training-specific generation kwargs to overwrite model's original generation config self.generation_kwargs = generation_kwargs - # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the - # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set - # self.model_accepts_loss_kwargs to False to enable scaling. - self.model_accepts_loss_kwargs = False self._dist = DistributedBackend(self.accelerator) # Add tags to the model