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 <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
7df8654dc4
commit
efc851a37b
1 changed files with 14 additions and 1 deletions
|
|
@ -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}"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue