Merge pull request #3834 from unslothai/rl-fixes
rl.py fixes: buffer reset, safer attribute access, typo fix
This commit is contained in:
commit
5b2ebe13c9
2 changed files with 22 additions and 13 deletions
|
|
@ -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,10 @@ from transformers.training_args import ParallelMode
|
|||
# Also patches W&B since multiple runs must use wandb.finish()
|
||||
import functools
|
||||
from types import MethodType
|
||||
try:
|
||||
from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers
|
||||
except:
|
||||
def reset_unsloth_gradient_checkpointing_buffers(): pass
|
||||
def prepare_for_training_mode(f):
|
||||
@functools.wraps(f)
|
||||
def wrapper(self, *args, **kwargs):
|
||||
|
|
@ -244,6 +248,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 +826,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 +1003,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
|
||||
|
|
|
|||
|
|
@ -879,12 +879,12 @@ def install_llama_cpp_make_non_blocking():
|
|||
IS_CMAKE = False
|
||||
if check == 0:
|
||||
# Uses old MAKE
|
||||
n_jobs = max(int(psutil.cpu_count() * 1.5), 1)
|
||||
n_jobs = max(int((psutil.cpu_count() or 1) * 1.5), 1)
|
||||
full_command = ["make", "all", "-j" + str(n_jobs), "-C", "llama.cpp"]
|
||||
IS_CMAKE = False
|
||||
else:
|
||||
# Uses new CMAKE
|
||||
n_jobs = max(int(psutil.cpu_count()), 1) # Use less CPUs since 1.5x faster
|
||||
n_jobs = max(int(psutil.cpu_count() or 1), 1) # Use less CPUs since 1.5x faster
|
||||
check = os.system(
|
||||
f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}"
|
||||
)
|
||||
|
|
@ -994,13 +994,13 @@ def install_llama_cpp_old(version = -10):
|
|||
# Try using MAKE
|
||||
commands = [
|
||||
"make clean -C llama.cpp",
|
||||
f"make all -j{psutil.cpu_count()*2} -C llama.cpp",
|
||||
f"make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp",
|
||||
]
|
||||
if try_execute(commands) == "CMAKE":
|
||||
# Instead use CMAKE
|
||||
commands = [
|
||||
f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}",
|
||||
f"cmake --build llama.cpp/build --config Release -j{psutil.cpu_count()*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}",
|
||||
f"cmake --build llama.cpp/build --config Release -j{(psutil.cpu_count() or 1)*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}",
|
||||
"cp llama.cpp/build/bin/llama-* llama.cpp",
|
||||
"rm -rf llama.cpp/build",
|
||||
]
|
||||
|
|
@ -1040,14 +1040,14 @@ def install_llama_cpp_blocking(use_cuda = False):
|
|||
"make clean -C llama.cpp",
|
||||
# https://github.com/ggerganov/llama.cpp/issues/7062
|
||||
# Weirdly GPU conversion for GGUF breaks??
|
||||
# f"{use_cuda} make all -j{psutil.cpu_count()*2} -C llama.cpp",
|
||||
f"make all -j{psutil.cpu_count()*2} -C llama.cpp",
|
||||
# f"{use_cuda} make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp",
|
||||
f"make all -j{(psutil.cpu_count() or 1)*2} -C llama.cpp",
|
||||
]
|
||||
if try_execute(commands) == "CMAKE":
|
||||
# Instead use CMAKE
|
||||
commands = [
|
||||
f"cmake llama.cpp -B llama.cpp/build -DBUILD_SHARED_LIBS=OFF -DGGML_CUDA=OFF {CURL_FLAG}",
|
||||
f"cmake --build llama.cpp/build --config Release -j{psutil.cpu_count()*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}",
|
||||
f"cmake --build llama.cpp/build --config Release -j{(psutil.cpu_count() or 1)*2} --clean-first --target {' '.join(LLAMA_CPP_TARGETS)}",
|
||||
"cp llama.cpp/build/bin/llama-* llama.cpp",
|
||||
"rm -rf llama.cpp/build",
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue