Fix: skip fp16/bf16 validation for full finetuning in RL trainers (#6813)
--------- Co-authored-by: Ayushman Paul <ayushman@HP>
This commit is contained in:
parent
22cd26f75d
commit
d33a7a7a1a
1 changed files with 3 additions and 0 deletions
|
|
@ -1015,6 +1015,9 @@ def _patch_trl_rl_trainers_impl(trainer_file = "grpo_trainer"):
|
|||
"dtype = _get_dtype(dtype)\n"
|
||||
"float16 = dtype == torch.float16\n"
|
||||
"bfloat16 = dtype == torch.bfloat16\n"
|
||||
"if full_finetuning:\n"
|
||||
" if bfloat16 and use_fp16: use_fp16 = False\n"
|
||||
" if float16 and use_bf16: use_bf16 = False\n"
|
||||
"if not force_float32 and (float16 and use_bf16): raise TypeError('Unsloth: Model is in float16 precision but you want to use bfloat16 precision. Set fp16 to `True` and bf16 to `False`')\n"
|
||||
"if not force_float32 and (bfloat16 and use_fp16): raise TypeError('Unsloth: Model is in bfloat16 precision but you want to use float16 precision. Set fp16 to `False` and bf16 to `True`')\n"
|
||||
"if force_float32:\n"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue