rl.py fixes: buffer reset, safer attribute access, typo fix

1. Auto-reset gradient checkpointing buffers after trainer.train()
   - Import and call reset_unsloth_gradient_checkpointing_buffers() in
     prepare_for_training_mode wrapper to free memory after training
     while keeping buffers ready for subsequent runs

2. Replace eval/exec with safer getattr/setattr
   - eval(f"trl.trainer.{trainer}") -> getattr(trl.trainer, trainer)
   - exec(f"...{unwrap} = ...") -> setattr(current_trainer, unwrap, ...)
   - exec(f"Trainer.prediction_step=...") -> direct assignment

3. Fix psutil.cpu_count() potentially returning None
   - Change psutil.cpu_count()+4 to (psutil.cpu_count() or 1)+4
   - Prevents TypeError on systems where cpu_count() returns None

4. Fix typo: oriignal_is_vlm_text -> original_is_vlm_text
This commit is contained in:
danielhanchen 2026-01-04 12:21:39 +00:00
commit 3d15865bbc

View file

@ -199,15 +199,15 @@ def PatchRL(FastLanguageModel):
unwrap = "unwrap_model_for_generation"
for trainer in trainers:
try:
current_trainer = eval(f"trl.trainer.{trainer}")
current_trainer = getattr(trl.trainer, trainer)
except:
continue
if hasattr(current_trainer, unwrap):
try:
exec(f"trl.trainer.{trainer}.{unwrap} = unsloth_{unwrap}")
setattr(current_trainer, unwrap, unsloth_unwrap_model_for_generation)
except:
continue
exec(f"Trainer.prediction_step=unsloth_prediction_step")
Trainer.prediction_step = unsloth_prediction_step
selective_log_softmax = RL_REPLACEMENTS["selective_log_softmax"]
@ -234,6 +234,7 @@ from transformers.training_args import ParallelMode
# Also patches W&B since multiple runs must use wandb.finish()
import functools
from types import MethodType
from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers
def prepare_for_training_mode(f):
@functools.wraps(f)
def wrapper(self, *args, **kwargs):
@ -244,6 +245,11 @@ def prepare_for_training_mode(f):
# Return inference mode
if hasattr(self, 'model') and hasattr(self.model, "for_inference"):
self.model.for_inference()
# Reset gradient checkpointing buffers to free memory while staying ready for next run
try:
reset_unsloth_gradient_checkpointing_buffers()
except:
pass
# Patch W&B to enable logging on future runs, otherwise it'll overwrite the first run
try:
import wandb
@ -817,7 +823,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
num_proc_check = (
"if dataset_num_proc is None:\n"
" import psutil\n"
" dataset_num_proc = min(max(psutil.cpu_count()+4, 2), 64)\n"
" dataset_num_proc = min(max((psutil.cpu_count() or 1)+4, 2), 64)\n"
" memory_gb_left = psutil.virtual_memory().available / (1024**3)\n"
" if memory_gb_left <= 4: dataset_num_proc = 1 # Too risky, so set to 1\n"
" elif memory_gb_left <= 6: dataset_num_proc = min(2, dataset_num_proc)\n"
@ -994,10 +1000,10 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
# Temporary patch _is_vlm to False
# as of 0.22 it only exists in sfttrainer
oriignal_is_vlm_text = "self._is_vlm = True"
original_is_vlm_text = "self._is_vlm = True"
new_is_vlm_text = "self._is_vlm = False"
RLTrainer_source = RLTrainer_source.replace(
oriignal_is_vlm_text, new_is_vlm_text
original_is_vlm_text, new_is_vlm_text
)
# Remove multiple doc strings