Latest TRL, GRPO + Bug fixes (#2645)

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* model_type_arch

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update loader.py

* check

* Update _utils.py

* Update loader.py

* Update loader.py

* Remove prints

* Update README.md

typo

* Update _utils.py

* Update _utils.py

* versioning

* Update _utils.py

* Update _utils.py

* Update _utils.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update llama.py

* Update vision.py

* HF Transfer

* fix(utils): add missing importlib import to fix NameError (#2134)

This commit fixes a NameError that occurs when `importlib` is referenced in _utils.py
without being imported, especially when UNSLOTH_USE_MODELSCOPE=1 is enabled.
By adding the missing import statement, the code will no longer throw a NameError.

* Add QLoRA Train and Merge16bit Test (#2130)

* add reference and unsloth lora merging tests

* add test / dataset printing to test scripts

* allow running tests from repo root

* add qlora test readme

* more readme edits

* ruff formatting

* additional readme comments

* forgot to add actual tests

* add apache license

* Update pyproject.toml

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update loader.py

* Update loader.py

* Revert

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Bug fix

* Update mapper.py

* check SDPA for Mistral 3, Pixtral

* Update vision.py

* Versioning

* Update rl_replacements.py

* Update README.md

* add model registry

* move hf hub utils to unsloth/utils

* refactor global model info dicts to dataclasses

* fix dataclass init

* fix llama registration

* remove deprecated key function

* start registry reog

* add llama vision

* quant types -> Enum

* remap literal quant types to QuantType Enum

* add llama model registration

* fix quant tag mapping

* add qwen2.5 models to registry

* add option to include original model in registry

* handle quant types per model size

* separate registration of base and instruct llama3.2

* add QwenQVQ to registry

* add gemma3 to registry

* add phi

* add deepseek v3

* add deepseek r1 base

* add deepseek r1 zero

* add deepseek distill llama

* add deepseek distill models

* remove redundant code when constructing model names

* add mistral small to registry

* rename model registration methods

* rename deepseek registration methods

* refactor naming for mistral and phi

* add global register models

* refactor model registration tests for new registry apis

* add model search method

* remove deprecated registration api

* add quant type test

* add registry readme

* make llama registration more specific

* clear registry when executing individual model registration file

* more registry readme updates

* Update _auto_install.py

* Llama4

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Synthetic data

* Update mapper.py

* Xet and Synthetic

* Update synthetic.py

* Update loader.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update pyproject.toml

* Delete .gitignore

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update _utils.py

* Update pyproject.toml

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update synthetic.py

* Update chat_templates.py

* Seasame force float16 / float32

* Fix Seasame

* Update loader.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update loader.py

* is_multimodal

* Update loader.py

* Update loader.py

* Update loader.py

* Update loader.py

* Update vision.py

* Update vision.py

* Update vision.py

* UNSLOTH_DISABLE_STATIC_GENERATION

* Update vision.py

* Auto vision detection

* Sesame

* Whisper

* Update loader.py

* Update loader.py

* Update loader.py

* Update mapper.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update vision.py

* Update loader.py

* Update loader.py

* Update loader.py

* Update loader.py

* Update _utils.py

* Update rl.py

* versioning

* Update rl.py

* Update rl.py

* Update rl.py

* Update rl.py

* Update rl.py

* logging

* Update pyproject.toml

* Update rl.py

---------

Co-authored-by: Jack Shi Wei Lun <87535974+jackswl@users.noreply.github.com>
Co-authored-by: naliazheli <nalia0316@gmail.com>
Co-authored-by: jeromeku <jerome.ku@gmail.com>
Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com>
This commit is contained in:
Daniel Han 2025-05-28 06:15:12 -07:00 committed by GitHub
commit 61dff3514a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 56 additions and 9 deletions

View file

@ -37,10 +37,10 @@ triton = [
]
huggingface = [
"unsloth_zoo>=2025.5.8",
"unsloth_zoo>=2025.5.10",
"packaging",
"tyro",
"transformers==4.51.3,!=4.47.0",
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2",
"datasets>=3.4.1",
"sentencepiece>=0.2.0",
"tqdm",
@ -48,7 +48,7 @@ huggingface = [
"wheel>=0.42.0",
"numpy",
"accelerate>=0.34.1",
"trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0,<=0.15.2",
"trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0",
"peft>=0.7.1,!=0.11.0",
"protobuf<4.0.0",
"huggingface_hub",
@ -381,10 +381,10 @@ colab-ampere-torch220 = [
"flash-attn>=2.6.3",
]
colab-new = [
"unsloth_zoo>=2025.5.8",
"unsloth_zoo>=2025.5.9",
"packaging",
"tyro",
"transformers==4.51.3,!=4.47.0",
"transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2",
"datasets>=3.4.1",
"sentencepiece>=0.2.0",
"tqdm",
@ -399,7 +399,7 @@ colab-new = [
]
colab-no-deps = [
"accelerate>=0.34.1",
"trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0,<=0.15.2",
"trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0",
"peft>=0.7.1",
"xformers",
"bitsandbytes>=0.45.5",

View file

@ -12,7 +12,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
__version__ = "2025.5.7"
__version__ = "2025.5.8"
__all__ = [
"SUPPORTS_BFLOAT16",

View file

@ -395,7 +395,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
if trainer_file in RL_METRICS_CHANGES:
process_extra_args = RL_METRICS_CHANGES[trainer_file]
for process_extra_arg in process_extra_args:
other_metrics_processor += process_extra_arg(call_args, extra_args)
other_metrics_processor += process_extra_arg(old_RLTrainer_source, old_RLConfig_source)
pass
# Add statistics as well!
@ -481,6 +481,39 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
extra_args += num_proc_check
pass
# Check for loss_type = dr_grpo and scale_rewards for GRPO
if "loss_type" in call_args and "scale_rewards" in call_args:
check_dr_grpo = \
"if loss_type.lower() == 'dr_grpo':\n"\
" loss_type = 'dr_grpo'\n"\
"elif loss_type.lower() == 'dapo':\n"\
" loss_type = 'dapo'\n"\
"if loss_type.lower() == 'dr_grpo':\n"\
" if scale_rewards == None:\n"\
" scale_rewards = True\n"\
" elif scale_rewards == True:\n"\
" print('The Dr GRPO paper recommends setting `scale_rewards` to False! Will override. Set it to `None` to force False.')\n"\
" scale_rewards = False\n"\
"elif loss_type.lower() == 'dapo':\n"\
" print('The DAPO paper recommends `mask_truncated_completions = True`')\n"\
" print('The DAPO paper recommends `epsilon_high = 0.28`')\n"\
" mask_truncated_completions = True\n"\
" epsilon_high = 0.28\n"\
"\n"
extra_args += check_dr_grpo
pass
# Check GRPO num_generations mismatch
if "per_device_train_batch_size" in call_args and "num_generations" in call_args:
check_num_generations = \
"if (per_device_train_batch_size // num_generations) * num_generations != per_device_train_batch_size:\n"\
" print('Unsloth: We now expect `per_device_train_batch_size` to be a multiple of `num_generations`.\\n"\
"We will change the batch size of ' + str(per_device_train_batch_size) + ' to the `num_generations` of ' + str(num_generations))\n"\
" per_device_train_batch_size = num_generations\n"\
"\n"
extra_args += check_num_generations
pass
# Edit config with anything extra
if trainer_file in RL_CONFIG_CHANGES:
process_extra_args = RL_CONFIG_CHANGES[trainer_file]

View file

@ -363,13 +363,27 @@ RL_CONFIG_CHANGES["grpo_trainer"].append(grpo_trainer_fix_batch_size)
def grpo_trainer_metrics(RLTrainer_source, RLConfig_source):
if "reward_funcs" not in RLTrainer_source: return ""
# For new TRL we have /mean and /std
use_mean = "rewards/{reward_func_name}/mean" in RLTrainer_source
use_std = "rewards/{reward_func_name}/std" in RLTrainer_source
if not use_mean:
use_normal = "rewards/{reward_func_name}" in RLTrainer_source
else:
use_normal = False
pass
log_metrics = \
"if not isinstance(reward_funcs, list): _reward_funcs = [reward_funcs]\n"\
"else: _reward_funcs = reward_funcs\n"\
"for reward_func in _reward_funcs:\n"\
" try:\n"\
" reward_func_name = reward_func.__name__\n"\
" other_metrics.append(f'rewards/{reward_func_name}')\n"\
f" if {use_mean}:\n"\
" other_metrics.append(f'rewards/{reward_func_name}/mean')\n"\
f" if {use_std}:\n"\
" other_metrics.append(f'rewards/{reward_func_name}/std')\n"\
f" if {use_normal}:\n"\
" other_metrics.append(f'rewards/{reward_func_name}')\n"\
" except: pass\n"
return log_metrics
pass