From efc851a37b27b420dc82bd74bdfe05bfc71c94c8 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Tue, 10 Feb 2026 06:17:47 -0800 Subject: [PATCH] Fix warmup_ratio deprecation for transformers >= 5.0 (#4019) * Fix warmup_ratio deprecation warning for transformers >= 5.0 In transformers 5.0, warmup_ratio is deprecated in favor of warmup_steps which now accepts float values (< 1 = ratio, >= 1 = absolute steps). The compiler now conditionally sets warmup_steps=0.1 on transformers >= 5.0 (same semantics as warmup_ratio=0.1) and keeps warmup_ratio=0.1 on older versions where warmup_steps only accepts int. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: Daniel Hanchen Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- unsloth/models/rl.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 32edcebaf8..181e9479df 100755 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -76,6 +76,14 @@ try: except Exception: torch_version = Version("0.0.0") +# Get transformers version for feature detection +try: + from transformers import __version__ as _transformers_version_raw + + transformers_version = Version(_transformers_version_raw) +except Exception: + transformers_version = Version("0.0.0") + def vLLMSamplingParams(**kwargs): from vllm import SamplingParams @@ -959,7 +967,6 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): "per_device_train_batch_size": 4, "gradient_accumulation_steps": 2, "weight_decay": 0.01, - "warmup_ratio": 0.1, "seed": 3407, "optim": "adamw_8bit", "learning_rate": 5e-05, @@ -986,6 +993,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # "dataloader_prefetch_factor" : 2, # "dataloader_num_workers" : 2, # Default is 0 means 1 } + # warmup_ratio deprecated in transformers >= 5.0; warmup_steps accepts float + if transformers_version >= Version("5.0.0"): + replacements["warmup_steps"] = 0.1 + else: + replacements["warmup_ratio"] = 0.1 + for k, v in replacements.items(): x = f"{k}( = [^,\n]{{1,}})?,\n" y = f"'{v}'" if type(v) is str else f"{v}"