Fix TRL 0.27.0 GRPO compatibility and PEFT model handling (#3969)
* Fix TRL 0.27.0 GRPO compatibility and PEFT model handling - Remove use_reentrant=False from gradient_checkpointing_kwargs for TRL 0.27.0+ TRL 0.27.0 auto-sets use_reentrant=False in GRPOConfig.__post_init__, but Unsloth gradient checkpointing requires use_reentrant=True. This adds a post-init cleanup that removes the setting when present. - Handle prepare_peft_model standalone function pattern for TRL 0.22.0+ TRL changed from self._prepare_peft_model() method to prepare_peft_model() standalone function. Both patterns are now bypassed to let Unsloth handle PEFT model preparation. Tested with TRL versions 0.22.2, 0.23.1, 0.24.0, 0.25.1, 0.26.2, and 0.27.1. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: danielhanchen <danielhanchen@users.noreply.github.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
7dd3ae8768
commit
949f1ce573
1 changed files with 16 additions and 0 deletions
|
|
@ -357,6 +357,7 @@ class Unsloth{RLConfig_name}({RLConfig_name}):
|
|||
)
|
||||
self.unsloth_logit_chunk_multiplier = unsloth_logit_chunk_multiplier
|
||||
{max_seq_length_post}
|
||||
{RLConfig_post}
|
||||
pass
|
||||
|
||||
{RLTrainer_extras}
|
||||
|
|
@ -1025,6 +1026,18 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
RLConfig_extra_args = extra_args
|
||||
RLConfig_call_args = call_args
|
||||
|
||||
# TRL 0.27.0+ forces use_reentrant=False in gradient_checkpointing_kwargs.
|
||||
# Unsloth gradient checkpointing requires use_reentrant=True, so we remove
|
||||
# the setting after super().__init__() when it gets auto-applied.
|
||||
RLConfig_post = ""
|
||||
if trl_version >= Version("0.27.0") and RLConfig_name == "GRPOConfig":
|
||||
RLConfig_post = (
|
||||
" # Unsloth: Remove use_reentrant=False forced by TRL 0.27.0+\n"
|
||||
" if getattr(self, 'gradient_checkpointing_kwargs', None) is not None:\n"
|
||||
" if 'use_reentrant' in self.gradient_checkpointing_kwargs:\n"
|
||||
" del self.gradient_checkpointing_kwargs['use_reentrant']\n"
|
||||
)
|
||||
|
||||
# Patch vLLM and other functions
|
||||
RLTrainer_extras = patch_functions(
|
||||
RLTrainer, trainer_file, RLTrainer_name, all_imports, imports
|
||||
|
|
@ -1077,6 +1090,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
RLConfig_extra_args = RLConfig_extra_args,
|
||||
RLConfig_call_args = RLConfig_call_args,
|
||||
RLConfig_kwargs = ",**kwargs"[1 if RLConfig_call_args.endswith(",") else 0 :],
|
||||
RLConfig_post = RLConfig_post,
|
||||
RLTrainer_extras = RLTrainer_extras,
|
||||
RLTrainer_post = RLTrainer_post,
|
||||
RL_pre = RL_pre,
|
||||
|
|
@ -1230,6 +1244,8 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import
|
|||
init = init.replace(
|
||||
"model = self._prepare_peft_model(model, peft_config, args)\n", "pass\n"
|
||||
)
|
||||
# TRL 0.22.0+ uses prepare_peft_model as a standalone function
|
||||
init = init.replace("model = prepare_peft_model(model, peft_config, args)", "pass")
|
||||
|
||||
# Skip add_adapter("ref") for reference model computation
|
||||
# Unsloth: We comment out the "ref" adapter creation because:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue