Fix bugs (#1701)
* Phi 4 * Update llama.py * Torch.Cuda Is Available Condition and Warning (#1545) * check for torch.cuda and triton if available on my machine(mac m3) the cuda were not available * Update pyproject.toml * Update __init__.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update mistral.py * Update mistral.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Fix * Bug fixes * Update mapper.py * Add dropout to granite to match HF's implementation (#1557) Signed-off-by: datta0 <venkatadattasainimmaturi@gmail.com> * Update llama.py * Update llama.py * Bug fixes * fix: flash_attn_detection_error (#1556) * fix: flash_attn_detection_error * Update _utils.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update mapper.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * dim fix * Update _utils.py * Torch 2.6 support * 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 * Faster inference? * Update llama.py * Update llama.py * Update utils.py * Update llama.py * Update llama.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update mapper.py * Fast Inference via vLLM * Update llama.py * Update llama.py * Update utils.py * Create rl.py * PatchRL * Update rl.py * Update rl.py * Update rl.py * PatchRLStatistics * Update rl.py * Update rl.py * Update rl.py * Update utils.py * Update utils.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * RL metrics * Update rl.py * RL metrics * Update __init__.py * Update rl.py * Update rl.py * Update rl.py * Update chat_templates.py * Update mapper.py * Fp8 cache * Update llama.py * Update llama.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update __init__.py * Update loader.py * Update rl.py * Update rl.py * Update _utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Better TRL handling * Update rl.py * Update tokenizer_utils.py * Auto patching * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update rl.py * Update tokenizer_utils.py * Update rl.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update tokenizer_utils.py * Update rl.py * Update rl.py * Update rl.py * max seq length * Update rl.py * Update rl.py * Patching * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * NEFTune * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Extra replacements * Update rl_replacements.py * Update rl.py * extra RL replacements * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update _utils.py * Update loader_utils.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * autocast * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update pyproject.toml * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update _utils.py --------- Signed-off-by: datta0 <venkatadattasainimmaturi@gmail.com> Co-authored-by: AminWhat <88392440+aminwhat@users.noreply.github.com> Co-authored-by: Datta Nimmaturi <datta.nimmaturi@nutanix.com> Co-authored-by: Zhe Zhang <2631992879@qq.com>
This commit is contained in:
parent
3232119b47
commit
47b81fdd23
7 changed files with 670 additions and 323 deletions
|
|
@ -187,9 +187,9 @@ cu124onlytorch260 = [
|
|||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp311-cp311-manylinux_2_28_x86_64.whl ; python_version=='3.11' and platform_system == 'Linux'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp312-cp312-manylinux_2_28_x86_64.whl ; python_version=='3.12' and platform_system == 'Linux'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp39-cp39-win_amd64.whl ; python_version=='3.9' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp310-cp310-win_amd64.whl ; python_version=='3.10' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp311-cp311-win_amd64.whl ; python_version=='3.11' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post2-cp312-cp312-win_amd64.whl ; python_version=='3.12' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post3-cp310-cp310-win_amd64.whl ; python_version=='3.10' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post3-cp311-cp311-win_amd64.whl ; python_version=='3.11' and platform_system == 'Windows'",
|
||||
"xformers @ https://download.pytorch.org/whl/cu124/xformers-0.0.29.post3-cp312-cp312-win_amd64.whl ; python_version=='3.12' and platform_system == 'Windows'",
|
||||
]
|
||||
cu126onlytorch260 = [
|
||||
"xformers @ https://download.pytorch.org/whl/cu126/xformers-0.0.29.post2-cp39-cp39-manylinux_2_28_x86_64.whl ; python_version=='3.9' and platform_system == 'Linux'",
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__version__ = "2025.2.4"
|
||||
__version__ = "2025.2.5"
|
||||
|
||||
__all__ = [
|
||||
"SUPPORTS_BFLOAT16",
|
||||
|
|
@ -131,6 +131,7 @@ logging.getLogger("transformers.tokenization_utils_base").setLevel(logging.CRITI
|
|||
|
||||
# Ignore logging messages
|
||||
class HideLoggingMessage(logging.Filter):
|
||||
__slots__ = "text",
|
||||
def __init__(self, text): self.text = text
|
||||
def filter(self, x): return not (self.text in x.getMessage())
|
||||
pass
|
||||
|
|
@ -138,6 +139,8 @@ pass
|
|||
# The speedups for torchdynamo mostly come wih GPU Ampere or higher and which is not detected here.
|
||||
from transformers.training_args import logger as transformers_training_args_logger
|
||||
transformers_training_args_logger.addFilter(HideLoggingMessage("The speedups"))
|
||||
# torch.distributed process group is initialized, but parallel_mode != ParallelMode.DISTRIBUTED.
|
||||
transformers_training_args_logger.addFilter(HideLoggingMessage("torch.distributed"))
|
||||
del transformers_training_args_logger
|
||||
|
||||
# Using the default loss: `ForCausalLMLoss`.
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@
|
|||
import torch
|
||||
import gc
|
||||
import math
|
||||
from functools import partial
|
||||
from typing import Optional, Tuple, List, Union
|
||||
from ._utils import *
|
||||
from ._utils import __version__
|
||||
|
|
@ -447,20 +448,28 @@ def LlamaAttention_fast_forward(
|
|||
A = flash_attn_func(Q, K, V, causal = True)
|
||||
else:
|
||||
# Grouped query attention
|
||||
if n_groups != 1:
|
||||
K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
|
||||
V = V[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
|
||||
K = K.reshape(bsz, n_heads, kv_seq_len, head_dim)
|
||||
V = V.reshape(bsz, n_heads, kv_seq_len, head_dim)
|
||||
if SDPA_HAS_GQA:
|
||||
# Needs (batch_size, n_heads, seq_len, head_dim)
|
||||
# is_casual and attention_mask must not be both set!
|
||||
A = scaled_dot_product_attention(Q, K, V, attn_mask = attention_mask, is_causal = False, enable_gqa = n_groups != 1)
|
||||
# Go back to (batch_size, seq_len, n_heads, head_dim)
|
||||
A = A.transpose(1, 2)#.contiguous()
|
||||
else:
|
||||
if n_groups != 1:
|
||||
K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
|
||||
V = V[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim)
|
||||
K = K.reshape(bsz, n_heads, kv_seq_len, head_dim)
|
||||
V = V.reshape(bsz, n_heads, kv_seq_len, head_dim)
|
||||
pass
|
||||
# Must be contiguous or else results are False!
|
||||
# https://github.com/pytorch/pytorch/issues/112577
|
||||
Q, K, V = Q.contiguous(), K.contiguous(), V.contiguous()
|
||||
# Needs (batch_size, n_heads, seq_len, head_dim)
|
||||
# is_casual and attention_mask must not be both set!
|
||||
A = scaled_dot_product_attention(Q, K, V, attn_mask = attention_mask, is_causal = False)
|
||||
# Go back to (batch_size, seq_len, n_heads, head_dim)
|
||||
A = A.transpose(1, 2).contiguous()
|
||||
pass
|
||||
# Must be contiguous or else results are False!
|
||||
# https://github.com/pytorch/pytorch/issues/112577
|
||||
Q, K, V = Q.contiguous(), K.contiguous(), V.contiguous()
|
||||
# Needs (batch_size, n_heads, seq_len, head_dim)
|
||||
# is_casual and attention_mask must not be both set!
|
||||
A = scaled_dot_product_attention(Q, K, V, attn_mask = attention_mask, is_causal = False)
|
||||
# Go back to (batch_size, seq_len, n_heads, head_dim)
|
||||
A = A.transpose(1, 2).contiguous()
|
||||
pass
|
||||
attn_output = A.reshape(bsz, q_len, n_heads*head_dim)
|
||||
attn_output = self.apply_o(self, attn_output)
|
||||
|
|
@ -699,6 +708,7 @@ def LlamaModel_fast_forward(
|
|||
if attention_mask is None:
|
||||
padding_mask = None
|
||||
elif self.training:
|
||||
# elif attention_mask is not None and self.training:
|
||||
attention_mask = None
|
||||
padding_mask = None
|
||||
else:
|
||||
|
|
@ -714,6 +724,7 @@ def LlamaModel_fast_forward(
|
|||
past_key_values_length,
|
||||
sliding_window = getattr(self.config, "sliding_window", None),
|
||||
)
|
||||
attention_mask = attention_mask.to(torch.bool)
|
||||
pass
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
|
@ -1802,8 +1813,6 @@ class FastLlamaModel:
|
|||
model = convert_vllm_to_huggingface(quant_state_dict, model_config, dtype)
|
||||
model.vllm_engine = llm
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
|
||||
from functools import partial
|
||||
model.fast_generate_batches = partial(generate_batches, model.vllm_engine)
|
||||
pass
|
||||
# Return old flag
|
||||
|
|
@ -1952,13 +1961,13 @@ class FastLlamaModel:
|
|||
Trainer._inner_training_loop = _fast_inner_training_loop
|
||||
|
||||
# Save max_seq_length
|
||||
model.max_seq_length = max_position_embeddings
|
||||
model.max_seq_length = max_seq_length
|
||||
internal_model = model
|
||||
while hasattr(internal_model, "model"):
|
||||
internal_model.max_seq_length = max_position_embeddings
|
||||
internal_model.max_seq_length = max_seq_length
|
||||
internal_model = internal_model.model
|
||||
pass
|
||||
internal_model.max_seq_length = max_position_embeddings
|
||||
internal_model.max_seq_length = max_seq_length
|
||||
|
||||
# We check the tokenizer first for errors
|
||||
if fix_tokenizer:
|
||||
|
|
@ -2146,8 +2155,6 @@ class FastLlamaModel:
|
|||
signature = str(inspect.signature(LoraConfig))
|
||||
SUPPORTS_LOFTQ = "loftq_config" in signature
|
||||
SUPPORTS_RSLORA = "use_rslora" in signature
|
||||
|
||||
assert(max_seq_length <= model.max_seq_length)
|
||||
|
||||
if lora_dropout != 0:
|
||||
logger.warning_once(
|
||||
|
|
@ -2632,6 +2639,10 @@ class FastLlamaModel:
|
|||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
pass
|
||||
|
||||
# Add for_inference and for_training
|
||||
model.for_training = partial(FastLlamaModel.for_training, model)
|
||||
model.for_inference = partial(FastLlamaModel.for_inference, model)
|
||||
return model
|
||||
pass
|
||||
|
||||
|
|
@ -2739,3 +2750,5 @@ class FastLlamaModel:
|
|||
pass
|
||||
pass
|
||||
|
||||
from .rl import PatchFastRL
|
||||
PatchFastRL(FastLanguageModel = FastLlamaModel)
|
||||
|
|
|
|||
|
|
@ -58,6 +58,11 @@ def __get_model_name(
|
|||
|
||||
elif load_in_4bit and SUPPORTS_FOURBIT and lower_model_name in FLOAT_TO_INT_MAPPER:
|
||||
|
||||
# Support returning original full -bnb-4bit name if specified specifically
|
||||
# since we'll map it to the dynamic version instead
|
||||
if lower_model_name.endswith("-bnb-4bit"):
|
||||
return lower_model_name
|
||||
|
||||
new_model_name = FLOAT_TO_INT_MAPPER[lower_model_name]
|
||||
# logger.warning_once(
|
||||
# f"Unsloth: You passed in `{model_name}` and `load_in_4bit = True`.\n"\
|
||||
|
|
|
|||
|
|
@ -16,30 +16,17 @@ __all__ = [
|
|||
"PatchFastRL",
|
||||
]
|
||||
|
||||
METRICS_MOVE_TO_END = [
|
||||
"nll",
|
||||
"aux",
|
||||
"beta",
|
||||
"alpha",
|
||||
]
|
||||
import torch
|
||||
try:
|
||||
from transformers.utils.notebook import (
|
||||
IntervalStrategy,
|
||||
NotebookTrainingTracker,
|
||||
NotebookProgressCallback,
|
||||
)
|
||||
HAS_NOTEBOOK = True
|
||||
except:
|
||||
HAS_NOTEBOOK = False
|
||||
pass
|
||||
from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, Union
|
||||
import inspect
|
||||
import os
|
||||
import re
|
||||
import functools
|
||||
from unsloth_zoo.compiler import create_new_function
|
||||
|
||||
from unsloth_zoo.logging_utils import PatchRLStatistics
|
||||
from .rl_replacements import (
|
||||
RL_EXTRA_ARGS,
|
||||
RL_FUNCTIONS,
|
||||
)
|
||||
|
||||
def PatchRL(FastLanguageModel):
|
||||
|
||||
|
|
@ -78,267 +65,441 @@ def PatchRL(FastLanguageModel):
|
|||
trainers = [x for x in trainers if x.endswith("_trainer")]
|
||||
unwrap = "unwrap_model_for_generation"
|
||||
for trainer in trainers:
|
||||
if hasattr(eval(f"trl.trainer.{trainer}"), unwrap):
|
||||
exec(f"trl.trainer.{trainer}.{unwrap} = unsloth_{unwrap}")
|
||||
try: current_trainer = eval(f"trl.trainer.{trainer}")
|
||||
except: continue
|
||||
if hasattr(current_trainer, unwrap):
|
||||
try: exec(f"trl.trainer.{trainer}.{unwrap} = unsloth_{unwrap}")
|
||||
except: continue
|
||||
pass
|
||||
pass
|
||||
|
||||
|
||||
def NotebookProgressCallback_on_train_begin(Trainer_metrics):
|
||||
def _NotebookProgressCallback_on_train_begin(self, args, state, control, **kwargs):
|
||||
self.first_column = "Epoch" if args.eval_strategy == IntervalStrategy.EPOCH else "Step"
|
||||
self.training_loss = 0
|
||||
self.last_log = 0
|
||||
column_names = [self.first_column] + ["Training Loss"]
|
||||
if args.eval_strategy != IntervalStrategy.NO:
|
||||
column_names.append("Validation Loss")
|
||||
column_names += [x.replace("/", " / ") for x in Trainer_metrics]
|
||||
self.training_tracker = NotebookTrainingTracker(state.max_steps, column_names)
|
||||
pass
|
||||
return _NotebookProgressCallback_on_train_begin
|
||||
RLTrainer_replacement = '''
|
||||
import os
|
||||
from typing import *
|
||||
from dataclasses import dataclass, field
|
||||
from packaging.version import Version
|
||||
import torch
|
||||
from contextlib import nullcontext
|
||||
|
||||
@dataclass
|
||||
class Unsloth{RLConfig_name}({RLConfig_name}):
|
||||
"""
|
||||
{__RLConfig_doc__}
|
||||
"""
|
||||
sampling_params: Optional[Any] = field(
|
||||
default = None,
|
||||
metadata = {{'help': 'vLLM SamplingParams'}},
|
||||
)
|
||||
def __init__({RLConfig_arguments},
|
||||
sampling_params = None,
|
||||
**kwargs,
|
||||
):
|
||||
{RLConfig_extra_args}
|
||||
super().__init__({RLConfig_call_args}{RLConfig_kwargs})
|
||||
pass
|
||||
|
||||
{RLTrainer_extras}
|
||||
|
||||
def NotebookProgressCallback_on_log(Trainer_metrics):
|
||||
def _NotebookProgressCallback_on_log(self, args, state, control, logs=None, **kwargs):
|
||||
# Only for when there is no evaluation
|
||||
if args.eval_strategy == IntervalStrategy.NO and "loss" in logs:
|
||||
values = {"Training Loss": logs["loss"]}
|
||||
for metric in Trainer_metrics:
|
||||
# Sometimes metric is not inside logs
|
||||
try: values[metric.replace("/", " / ")] = logs[metric]
|
||||
except: pass
|
||||
pass
|
||||
# First column is necessarily Step since we're not in epoch eval strategy
|
||||
values["Step"] = state.global_step
|
||||
self.training_tracker.write_line(values)
|
||||
pass
|
||||
pass
|
||||
return _NotebookProgressCallback_on_log
|
||||
class Unsloth{RLTrainer_name}(_Unsloth{RLTrainer_name}):
|
||||
"""
|
||||
{__RLTrainer_doc__}
|
||||
"""
|
||||
def __init__({RLTrainer_arguments},
|
||||
**kwargs
|
||||
):
|
||||
if args is None: args = Unsloth{RLConfig_name}()
|
||||
{RLTrainer_extra_args}
|
||||
super().__init__({RLTrainer_call_args}{RLTrainer_kwargs})
|
||||
{RLTrainer_post}
|
||||
pass
|
||||
|
||||
|
||||
def NotebookTrainingTracker_write_line(Trainer_metrics):
|
||||
set_Trainer_metrics = set(Trainer_metrics)
|
||||
def _NotebookTrainingTracker_write_line(self, values):
|
||||
"""
|
||||
Write the values in the inner table.
|
||||
|
||||
Args:
|
||||
values (`Dict[str, float]`): The values to display.
|
||||
"""
|
||||
if self.inner_table is None:
|
||||
self.inner_table = [list(values.keys()), list(values.values())]
|
||||
else:
|
||||
columns = self.inner_table[0]
|
||||
new_values = {}
|
||||
for key, value in values.items():
|
||||
lowered = key.lower()
|
||||
if lowered in set_Trainer_metrics:
|
||||
new_values[lowered.replace("/", " / ")] = value
|
||||
else:
|
||||
new_values[key] = value
|
||||
pass
|
||||
values = new_values
|
||||
|
||||
self.inner_table[0] = columns
|
||||
if len(self.inner_table) > 1:
|
||||
last_values = self.inner_table[-1]
|
||||
first_column = self.inner_table[0][0]
|
||||
if last_values[0] != values[first_column]:
|
||||
# write new line
|
||||
self.inner_table.append([values[c] if c in values else "No Log" for c in columns])
|
||||
else:
|
||||
# update last line
|
||||
new_values = values
|
||||
for c in columns:
|
||||
if c not in new_values.keys():
|
||||
new_values[c] = last_values[columns.index(c)]
|
||||
self.inner_table[-1] = [new_values[c] for c in columns]
|
||||
else:
|
||||
# Edit for evaluation purposes
|
||||
self.inner_table.append([values[c] if c in values else 0 for c in columns])
|
||||
pass
|
||||
pass
|
||||
pass
|
||||
return _NotebookTrainingTracker_write_line
|
||||
pass
|
||||
|
||||
|
||||
def _PatchRLStatistics(metrics, algorithm):
|
||||
if HAS_NOTEBOOK:
|
||||
if len(metrics) == 0:
|
||||
raise RuntimeError(f"Unsloth: RL statistics for {algorithm} failed with no metrics seen?")
|
||||
from transformers.trainer import is_in_notebook
|
||||
if is_in_notebook():
|
||||
# Patch DPO notebook printing
|
||||
NotebookTrainingTracker.write_line = NotebookTrainingTracker_write_line(metrics)
|
||||
from transformers.trainer import DEFAULT_PROGRESS_CALLBACK
|
||||
DEFAULT_PROGRESS_CALLBACK.on_train_begin = NotebookProgressCallback_on_train_begin(metrics)
|
||||
DEFAULT_PROGRESS_CALLBACK.on_log = NotebookProgressCallback_on_log(metrics)
|
||||
pass
|
||||
pass
|
||||
pass
|
||||
|
||||
|
||||
@functools.cache
|
||||
def get_trl_metrics():
|
||||
# Gets metrics so we can output them in notebooks
|
||||
|
||||
import trl.trainer
|
||||
trainers = dir(trl.trainer)
|
||||
trainers = [x for x in trainers if x.endswith("_trainer")]
|
||||
filepath = inspect.getfile(trl.trainer)
|
||||
filepath = os.path.split(filepath)[0]
|
||||
|
||||
all_metrics = dict()
|
||||
for trainer in trainers:
|
||||
filename = os.path.join(filepath, f"{trainer}.py")
|
||||
if not os.path.exists(filename): continue
|
||||
with open(filename, "r") as file: file = file.read()
|
||||
|
||||
# Get metrics['kl'] or stats['kl']
|
||||
metrics = re.findall(r"metrics\[[\"\']([^\"\']{1,})[\"\']\]", file)
|
||||
stats = re.findall(r"stats\[[\"\']([^\"\']{1,})[\"\']\]", file)
|
||||
metrics = metrics + stats
|
||||
|
||||
# Get optional f-strings
|
||||
metrics_f = re.findall(r"metrics\[f[\"\']\{[^\}]{1,}\}([^\"\']{1,})[\"\']\]", file)
|
||||
stats_f = re.findall(r"stats\[f[\"\']\{[^\}]{1,}\}([^\"\']{1,})[\"\']\]", file)
|
||||
metrics_f = metrics_f + stats_f
|
||||
# Filter out prefixes if seen
|
||||
# metrics[f"{prefix}rewards/chosen"]
|
||||
left_prefix = 'prefix = "eval_" if train_eval == "eval" else ""' in file
|
||||
if left_prefix: metrics += metrics_f
|
||||
|
||||
# Move all eval_ things to the end and reward to the front
|
||||
beginning = []
|
||||
middle = []
|
||||
end = []
|
||||
for x in metrics:
|
||||
lowered = x.lower()
|
||||
if "reward" in lowered:
|
||||
beginning.append(x)
|
||||
elif x.lower().startswith("eval"):
|
||||
end.append(x)
|
||||
else:
|
||||
# Check if we want to move to the end
|
||||
moved = False
|
||||
for move_end in METRICS_MOVE_TO_END:
|
||||
if move_end in lowered:
|
||||
end.append(x)
|
||||
moved = True
|
||||
break
|
||||
if not moved:
|
||||
middle.append(x)
|
||||
pass
|
||||
pass
|
||||
metrics = beginning + middle + end
|
||||
|
||||
all_metrics[trainer[:trainer.find("_")].upper()] = metrics
|
||||
pass
|
||||
return all_metrics
|
||||
pass
|
||||
|
||||
|
||||
def PatchRLStatistics(algorithm = "GRPO"):
|
||||
# Get notebook statistics columns to show up
|
||||
algorithm = algorithm.upper()
|
||||
all_metrics = get_trl_metrics()
|
||||
if algorithm not in all_metrics:
|
||||
print(
|
||||
f"Unsloth for {algorithm.upper()} is not yet implemented! Just ignore this function.\n"\
|
||||
f"We support: `{list(all_metrics.keys())}`"
|
||||
)
|
||||
pass
|
||||
_PatchRLStatistics(all_metrics[algorithm], algorithm)
|
||||
pass
|
||||
|
||||
'''
|
||||
|
||||
def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
||||
# Patch for vLLM and Unsloth PEFT
|
||||
import trl
|
||||
import trl.trainer
|
||||
|
||||
trainer = eval(f"trl.trainer.{trainer_file}")
|
||||
name = [x for x in dir(trainer) if x.endswith("Trainer") and x != "Trainer" and trainer_file.split("_")[0] in x.lower()]
|
||||
assert(len(name) == 1)
|
||||
RLTrainer_name = name[0]
|
||||
RLTrainer = eval(f"trl.trainer.{trainer_file}.{RLTrainer_name}")
|
||||
|
||||
try:
|
||||
__init__ = inspect.getsource(RLTrainer.__init__)
|
||||
except:
|
||||
# Already patched most likely!
|
||||
trainer = eval(f"trl.trainer.{trainer_file}")
|
||||
except Exception as error:
|
||||
return
|
||||
old__init__ = __init__
|
||||
|
||||
# Get SFTTrainer and SFTConfig names
|
||||
name = [x for x in dir(trainer) if x.endswith("Trainer") and x != "Trainer" and trainer_file.split("_")[0] in x.lower()]
|
||||
config = [x for x in dir(trainer) if x.endswith("Config") and x != "Config" and trainer_file.split("_")[0] in x.lower()]
|
||||
if len(name) != 1: return
|
||||
if len(config) != 1: return
|
||||
|
||||
# Get SFTTrainer, SFTConfig
|
||||
RLTrainer_name = name[0]
|
||||
RLConfig_name = config[0]
|
||||
try: RLTrainer = eval(f"trl.trainer.{trainer_file}.{RLTrainer_name}")
|
||||
except: return
|
||||
try: RLConfig = eval(f"trl.trainer.{trainer_file}.{RLConfig_name}" )
|
||||
except: return
|
||||
|
||||
# Check name
|
||||
if RLTrainer.__name__.startswith("Unsloth"): return
|
||||
if RLConfig .__name__.startswith("Unsloth"): return
|
||||
|
||||
all_imports = dir(trainer)
|
||||
assert("Union" in all_imports)
|
||||
imports = [x for x in all_imports if not x.startswith("_")]
|
||||
imports += ["Trainer"]
|
||||
|
||||
spaces = __init__.find("def")
|
||||
__init__ = __init__.split("\n")
|
||||
__init__ = "\n".join(x[spaces:] for x in __init__)
|
||||
# Get default arguments
|
||||
EMPTY = inspect.Parameter.empty
|
||||
processed = []
|
||||
for RLobject in [RLTrainer, RLConfig]:
|
||||
parameters = inspect.signature(RLobject.__init__).parameters
|
||||
types = (bool, type(None), int, float, str,)
|
||||
arguments = ["self"]
|
||||
call_args = []
|
||||
for k, v in parameters.items():
|
||||
if k == "self": continue
|
||||
v = v.default
|
||||
if v == "\n": v = re.escape("\n")
|
||||
if v is EMPTY: arguments.append(k)
|
||||
elif type(v) is str: arguments.append(f"{k} = '{v}'")
|
||||
elif type(v) in types: arguments.append(f"{k} = {v}")
|
||||
else: continue
|
||||
call_args.append(f"{k} = {k}")
|
||||
pass
|
||||
arguments = f"\n{' '*8}" + f",\n{' '*8}".join(arguments)
|
||||
call_args = f"\n{' '*12}" + f",\n{' '*12}".join(call_args)
|
||||
processed.append((arguments, call_args,))
|
||||
pass
|
||||
|
||||
# Replace vLLM sections since we already have it done!
|
||||
vllm_part = re.findall(
|
||||
r"(\n[\s]{4}"\
|
||||
r"if (self|args)\.use_vllm\:.+?"\
|
||||
r"\n[\s]{4,}"\
|
||||
"else:\n)",
|
||||
__init__,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
# Process RLTrainer first
|
||||
arguments, call_args = processed[0]
|
||||
RLTrainer_post = ""
|
||||
|
||||
# Add tokenizer if not seen
|
||||
if "tokenizer" not in parameters and "processing_class" in parameters:
|
||||
arguments += f",\n{' '*8}tokenizer = None"
|
||||
call_args = call_args.replace(
|
||||
"processing_class = processing_class",
|
||||
"processing_class = tokenizer if tokenizer is not None else processing_class",
|
||||
)
|
||||
pass
|
||||
|
||||
# Edit bf16, fp16 by checking model's torch_dtype directly
|
||||
extra_args = ""
|
||||
if "args" in call_args and "model" in call_args:
|
||||
mixed_precision = \
|
||||
"use_bf16 = getattr(args, 'bf16', False)\n"\
|
||||
"use_fp16 = getattr(args, 'fp16', False)\n"\
|
||||
"dtype = getattr(model.config, 'torch_dtype', None)\n"\
|
||||
"if dtype is None: dtype = model.get_input_embeddings().dtype\n"\
|
||||
"from unsloth_zoo.utils import _get_dtype\n"\
|
||||
"dtype = _get_dtype(dtype)\n"\
|
||||
"float16 = dtype == torch.float16\n"\
|
||||
"if 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 float16 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 not use_bf16 and not use_fp16:\n"\
|
||||
" args.fp16 = float16\n"\
|
||||
" args.bf16 = not float16\n"\
|
||||
" os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'\n"
|
||||
extra_args += mixed_precision
|
||||
pass
|
||||
|
||||
# Check if per_device_eval_batch_size (default 8) bigger than bsz
|
||||
# Also use FP16 / BF16 evaluation
|
||||
if "args" in call_args:
|
||||
# Check eval_dataset first
|
||||
if "eval_dataset" in call_args:
|
||||
check_eval_dataset = \
|
||||
"if getattr(args, 'eval_dataset', None) is not None and "\
|
||||
"getattr(args, 'eval_strategy', 'no') == 'no':\n"\
|
||||
" args.eval_strategy = 'steps'\n"\
|
||||
" if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1\n"
|
||||
extra_args += check_eval_dataset
|
||||
pass
|
||||
|
||||
# Check if gradient accumulation bug fix is applied
|
||||
check_ga = \
|
||||
"ga_steps = getattr(args, 'gradient_accumulation_steps', None)\n"\
|
||||
"if ga_steps is not None and ga_steps > 1:\n"\
|
||||
" from transformers import __version__ as transformers_version\n"\
|
||||
" if Version(transformers_version) <= Version('4.45.2'):\n"\
|
||||
" print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\\n'\n"\
|
||||
" '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')\n"
|
||||
extra_args += check_ga
|
||||
|
||||
eval_changes = \
|
||||
"if getattr(args, 'eval_strategy', 'no') != 'no':\n"\
|
||||
" eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)\n"\
|
||||
" if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size\n"\
|
||||
" if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps\n"\
|
||||
"fp16_full_eval = getattr(args, 'fp16_full_eval', False)\n"\
|
||||
"bf16_full_eval = getattr(args, 'bf16_full_eval', False)\n"\
|
||||
"if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True\n"\
|
||||
"if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False\n"\
|
||||
"if not bf16_full_eval and not fp16_full_eval: args.bf16_full_eval = args.bf16; args.fp16_full_eval = args.fp16\n"
|
||||
extra_args += eval_changes
|
||||
pass
|
||||
|
||||
# Check max_seq_length
|
||||
if "model" in call_args:
|
||||
length_check = \
|
||||
"if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):\n"\
|
||||
" pass\n"\
|
||||
"else:\n"\
|
||||
" model_max_seq_length = getattr(model, 'max_seq_length', None)\n"\
|
||||
" args_max_seq_length = getattr(args, 'max_seq_length', None)\n"\
|
||||
" if args_max_seq_length is None and model_max_seq_length is not None:\n"\
|
||||
" max_seq_length = model.max_seq_length\n"\
|
||||
" if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length\n"
|
||||
" elif args_max_seq_length is not None and model_max_seq_length is not None:\n"\
|
||||
" if args_max_seq_length > model_max_seq_length:\n"\
|
||||
" print('Unsloth: You set `max_seq_length` as ' + str(args_max_seq_length) + ' but \n"\
|
||||
" the maximum the model supports is ' + str(model_max_seq_length) + '. We shall reduce it.')\n"\
|
||||
" args.max_seq_length = model_max_seq_length\n"
|
||||
extra_args += length_check
|
||||
pass
|
||||
|
||||
# Enable for training and move padding side of tokenizer to right
|
||||
if "model" in call_args:
|
||||
training_check = \
|
||||
"if model is not None and hasattr(model, 'for_training'):\n"\
|
||||
" model.for_training()\n"\
|
||||
"if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'\n"\
|
||||
"if 'processing_class' in locals():\n"\
|
||||
" if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'\n"\
|
||||
" if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): "\
|
||||
"processing_class.tokenizer.padding_side = 'right'\n"
|
||||
extra_args += training_check
|
||||
pass
|
||||
|
||||
# Check NEFTune
|
||||
if "model" in call_args:
|
||||
neftune_check = \
|
||||
"if hasattr(self, 'neftune_hook_handle'):\n"\
|
||||
" self.neftune_hook_handle.remove()\n"\
|
||||
" if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle\n"\
|
||||
"if getattr(args, 'neftune_noise_alpha', None) is not None:\n"\
|
||||
" model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha\n"\
|
||||
"pass\n"
|
||||
RLTrainer_post += neftune_check
|
||||
pass
|
||||
|
||||
# Add statistics as well!
|
||||
extra_args += \
|
||||
"from unsloth_zoo.logging_utils import PatchRLStatistics\n"\
|
||||
f"PatchRLStatistics('{trainer_file}')\n"
|
||||
|
||||
# Patch optional args
|
||||
if trainer_file in RL_EXTRA_ARGS:
|
||||
process_extra_args = RL_EXTRA_ARGS[trainer_file]
|
||||
for process_extra_arg in process_extra_args:
|
||||
extra_args += process_extra_arg(call_args, extra_args)
|
||||
pass
|
||||
|
||||
# Create RLTrainer args
|
||||
extra_args = extra_args.split("\n")
|
||||
extra_args = "\n".join(" "*8 + x for x in extra_args)
|
||||
RLTrainer_post = RLTrainer_post.split("\n")
|
||||
RLTrainer_post = "\n".join(" "*8 + x for x in RLTrainer_post)
|
||||
RLTrainer_arguments = arguments
|
||||
RLTrainer_extra_args = extra_args
|
||||
RLTrainer_call_args = call_args
|
||||
|
||||
# Fix RLConfig next
|
||||
arguments, call_args = processed[1]
|
||||
extra_args = ""
|
||||
|
||||
# Edit GA / bsz and weight_decay
|
||||
replacements = {
|
||||
"output_dir" : None,
|
||||
"logging_nan_inf_filter" : False,
|
||||
"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,
|
||||
"per_device_eval_batch_size" : 4,
|
||||
"eval_accumulation_steps" : 2,
|
||||
"torch_empty_cache_steps" : 250,
|
||||
"logging_steps" : 1,
|
||||
}
|
||||
for k, v in replacements.items():
|
||||
x = f"{k}( = [^,\n]{{1,}})?,\n"
|
||||
y = f"'{v}'" if type(v) is str else f"{v}"
|
||||
y = f"{k} = {y},\n"
|
||||
arguments = re.sub(x, y, arguments)
|
||||
pass
|
||||
|
||||
# Warn on too large or too small learning rate
|
||||
if " learning_rate" in call_args:
|
||||
learning_rate_check = \
|
||||
"if learning_rate < 1e-7: raise FloatingPointError(f'Unsloth: Your learning rate of `{learning_rate}` is too small and less than 1e-7! "\
|
||||
"Consider increasing it, otherwise gradient updates will be close to 0!')\n"\
|
||||
"if learning_rate > 1: raise OverflowError(f'Unsloth: Your learning rate of `{learning_rate}` is way too larger > 1! "\
|
||||
"Consider decreasing it to 1e-1, otherwise gradient updates will explode!')\n"
|
||||
extra_args += learning_rate_check
|
||||
pass
|
||||
|
||||
# Add output_dir saving
|
||||
if "output_dir" in call_args:
|
||||
# Default checks
|
||||
saving_check = \
|
||||
"if output_dir is None and save_strategy == 'steps' and save_steps == 500:\n"\
|
||||
" output_dir = 'unsloth_training_checkpoints'\n"\
|
||||
" save_strategy = 'no'\n"
|
||||
extra_args += saving_check
|
||||
pass
|
||||
|
||||
# Edit dataset_num_proc
|
||||
if "dataset_num_proc" in call_args:
|
||||
num_proc_check = \
|
||||
"if dataset_num_proc is None:\n"\
|
||||
" from multiprocessing import cpu_count\n"\
|
||||
" dataset_num_proc = cpu_count()\n"
|
||||
extra_args += num_proc_check
|
||||
pass
|
||||
|
||||
# Edit report_to and default it to nothing if max_steps is like 60
|
||||
|
||||
# Create RLConfig args
|
||||
extra_args = extra_args.split("\n")
|
||||
extra_args = "\n".join(" "*8 + x for x in extra_args)
|
||||
RLConfig_arguments = arguments
|
||||
RLConfig_extra_args = extra_args
|
||||
RLConfig_call_args = call_args
|
||||
|
||||
# Patch vLLM and other functions
|
||||
RLTrainer_extras = patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, imports)
|
||||
if RLTrainer_extras is None:
|
||||
RLTrainer_extras = f"_Unsloth{RLTrainer_name} = {RLTrainer_name}"
|
||||
|
||||
# Create full module
|
||||
exec(f"from trl.trainer import ({RLTrainer_name}, {RLConfig_name},)")
|
||||
__RLTrainer_doc__ = eval(f"trl.trainer.{RLTrainer_name}").__doc__
|
||||
__RLConfig_doc__ = eval(f"trl.trainer.{RLConfig_name}") .__doc__
|
||||
|
||||
RLTrainer_source = RLTrainer_replacement.format(
|
||||
RLTrainer_name = RLTrainer_name,
|
||||
__RLTrainer_doc__ = __RLTrainer_doc__,
|
||||
RLTrainer_arguments = RLTrainer_arguments,
|
||||
RLTrainer_extra_args = RLTrainer_extra_args,
|
||||
RLTrainer_call_args = RLTrainer_call_args,
|
||||
RLTrainer_kwargs = ",**kwargs"[1 if RLTrainer_call_args.endswith(",") else 0:],
|
||||
|
||||
RLConfig_name = RLConfig_name,
|
||||
__RLConfig_doc__ = __RLConfig_doc__,
|
||||
RLConfig_arguments = RLConfig_arguments,
|
||||
RLConfig_extra_args = RLConfig_extra_args,
|
||||
RLConfig_call_args = RLConfig_call_args,
|
||||
RLConfig_kwargs = ",**kwargs"[1 if RLConfig_call_args .endswith(",") else 0:],
|
||||
|
||||
RLTrainer_extras = RLTrainer_extras,
|
||||
RLTrainer_post = RLTrainer_post,
|
||||
)
|
||||
if (len(vllm_part) != 1): return
|
||||
|
||||
vllm_part, args = vllm_part[0][0], vllm_part[0][1]
|
||||
# Strip all comments
|
||||
new_vllm_part = re.sub(r"\#[^\n]{1,}\n", "", vllm_part)
|
||||
|
||||
# Get SamplingParams
|
||||
sampling_params = re.findall(
|
||||
r"\n[\s]{4,}(self\.[^\s]{1,}[\s]{0,}\=[\s]{0,}"\
|
||||
r"SamplingParams\(.+?\))",
|
||||
new_vllm_part,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
# Create new function
|
||||
created_module = create_new_function(
|
||||
f"Unsloth{RLTrainer_name}",
|
||||
RLTrainer_source,
|
||||
f"trl.trainer.{trainer_file}",
|
||||
imports,
|
||||
overwrite = False,
|
||||
)
|
||||
if len(sampling_params) != 1: return
|
||||
|
||||
# Patch Trainer
|
||||
exec(f"trl.{RLTrainer_name} = created_module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{RLTrainer_name} = created_module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{trainer_file}.{RLTrainer_name} = created_module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
|
||||
# Patch Config
|
||||
exec(f"trl.{RLConfig_name} = created_module.Unsloth{RLConfig_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{RLConfig_name} = created_module.Unsloth{RLConfig_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{trainer_file}.{RLConfig_name} = created_module.Unsloth{RLConfig_name}", locals(), globals())
|
||||
pass
|
||||
|
||||
sampling_params = sampling_params[0]
|
||||
# Replace with our vLLM engine
|
||||
sampling_params = \
|
||||
" "*8 + "self.llm = model.vllm_engine; self._last_loaded_step = 0; " + \
|
||||
sampling_params # Add spaces
|
||||
new_vllm_part = f"\n if {args}.use_vllm:\n{sampling_params}\n else:\n"
|
||||
__init__ = __init__.replace(vllm_part, new_vllm_part)
|
||||
|
||||
def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, imports):
|
||||
init = inspect.getsource(RLTrainer.__init__)
|
||||
old_init = init
|
||||
|
||||
# Remove peft_config
|
||||
__init__ = __init__.replace("elif peft_config is None:", "elif False:")
|
||||
__init__ = __init__.replace("elif peft_config is not None:", "elif False:")
|
||||
__init__ = __init__.replace("if peft_config is None:", "if False:")
|
||||
__init__ = __init__.replace("if peft_config is not None:", "if False:")
|
||||
__init__ = __init__.replace("get_peft_model(model, peft_config)", "model")
|
||||
init = init.replace("elif peft_config is None:", "elif False:")
|
||||
init = init.replace("elif peft_config is not None:", "elif False:")
|
||||
init = init.replace("if peft_config is None:", "if False:")
|
||||
init = init.replace("if peft_config is not None:", "if False:")
|
||||
init = init.replace("get_peft_model(model, peft_config)", "model")
|
||||
|
||||
# Add spaces back into __init__
|
||||
__init__ = __init__.split("\n")
|
||||
__init__ = "\n".join(' '*spaces + x for x in __init__)
|
||||
# Set use_vllm if not set
|
||||
if "args.use_vllm" in init and "model" in init and "args" in init:
|
||||
# .*? matches first match. .+? matches final match.
|
||||
replacer = re.findall(
|
||||
"def __init__\(.*?\).*?\:\n",
|
||||
init,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(replacer) != 0:
|
||||
replacer = replacer[0]
|
||||
vllm_setter = "\n" + " "*8 + \
|
||||
"if hasattr(model, 'vllm_engine') and "\
|
||||
"getattr(args, 'use_vllm') and getattr(args, 'use_vllm', False): "\
|
||||
"args.use_vllm = True\n"
|
||||
init = init.replace(replacer, replacer + vllm_setter)
|
||||
pass
|
||||
pass
|
||||
|
||||
vllm_part = re.findall(
|
||||
r"(\n[\s]{8}"\
|
||||
r"if (self|args)\.use_vllm\:.*?"\
|
||||
r"\n[\s]{8}"\
|
||||
"else:\n)",
|
||||
init,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(vllm_part) == 1:
|
||||
vllm_part, args = vllm_part[0][0], vllm_part[0][1]
|
||||
# Strip all comments
|
||||
new_vllm_part = re.sub(r"\#[^\n]{1,}\n", "", vllm_part)
|
||||
|
||||
# Get SamplingParams
|
||||
sampling_params = re.findall(
|
||||
r"\n[\s]{4,}(self\.[^\s]{1,}[\s]{0,}\=[\s]{0,}"\
|
||||
r"SamplingParams\(.+?\))",
|
||||
new_vllm_part,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(sampling_params) == 1:
|
||||
sampling_params = sampling_params[0]
|
||||
# Replace with our vLLM engine
|
||||
sampling_params = \
|
||||
" "*12 + "self.llm = model.vllm_engine; self._last_loaded_step = 0; " + \
|
||||
sampling_params # Add spaces
|
||||
new_vllm_part = \
|
||||
f"\n{' '*8}if {args}.use_vllm:\n{sampling_params} "\
|
||||
f"if getattr(args, 'sampling_params', None) is None else "\
|
||||
f"getattr(args, 'sampling_params', None)\n{' '*8}else:\n"
|
||||
init = init.replace(vllm_part, new_vllm_part)
|
||||
pass
|
||||
pass
|
||||
|
||||
# Search for vLLM calling in all child functions
|
||||
functions = dir(RLTrainer)
|
||||
RLTrainer_source = inspect.getsource(RLTrainer)
|
||||
functions = [x for x in functions if f"def {x}" in RLTrainer_source]
|
||||
|
||||
changed = {"__init__" : (old__init__, __init__,)}
|
||||
changed = {"__init__" : (old_init, init,)}
|
||||
edit_functions = RL_FUNCTIONS.get(trainer_file, [])
|
||||
|
||||
for function in functions:
|
||||
if not hasattr(RLTrainer, function): continue
|
||||
fx = getattr(RLTrainer, function)
|
||||
try:
|
||||
source = inspect.getsource(fx)
|
||||
except:
|
||||
continue
|
||||
try: source = inspect.getsource(fx)
|
||||
except: continue
|
||||
original_source = source
|
||||
|
||||
# Check for function
|
||||
for edit_function in edit_functions:
|
||||
source = edit_function(function, source)
|
||||
pass
|
||||
|
||||
# llm_model = self.llm.llm_engine.model_executor.driver_worker.model_runner.model
|
||||
source = re.sub(
|
||||
r"(\n[\s]{4,}).+?model_executor\.driver_worker.+?\n",
|
||||
|
|
@ -386,22 +547,9 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
RLTrainer_source = RLTrainer_source.replace(old, new)
|
||||
pass
|
||||
RLTrainer_source = RLTrainer_source.replace(
|
||||
f"class {RLTrainer_name}", f"class Unsloth{RLTrainer_name}", 1
|
||||
f"class {RLTrainer_name}", f"class _Unsloth{RLTrainer_name}", 1
|
||||
)
|
||||
|
||||
# Create new class in compiled cache and import it
|
||||
module = create_new_function(
|
||||
RLTrainer_name,
|
||||
RLTrainer_source,
|
||||
f"trl.trainer.{trainer_file}",
|
||||
imports,
|
||||
)
|
||||
|
||||
# Patch over modules
|
||||
exec(f"trl.{RLTrainer_name} = module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{RLTrainer_name} = module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
exec(f"trl.trainer.{trainer_file}.{RLTrainer_name} = module.Unsloth{RLTrainer_name}", locals(), globals())
|
||||
return module
|
||||
return RLTrainer_source
|
||||
pass
|
||||
|
||||
|
||||
|
|
@ -416,8 +564,8 @@ def patch_trl_rl_trainers():
|
|||
pass
|
||||
|
||||
|
||||
def PatchFastRL(algorithm = "GRPO", FastLanguageModel = None):
|
||||
def PatchFastRL(algorithm = None, FastLanguageModel = None):
|
||||
if FastLanguageModel is not None: PatchRL(FastLanguageModel)
|
||||
patch_trl_rl_trainers()
|
||||
PatchRLStatistics(algorithm)
|
||||
if algorithm is not None: PatchRLStatistics(algorithm)
|
||||
pass
|
||||
|
|
|
|||
186
unsloth/models/rl_replacements.py
Normal file
186
unsloth/models/rl_replacements.py
Normal file
|
|
@ -0,0 +1,186 @@
|
|||
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__all__ = [
|
||||
"RL_EXTRA_ARGS",
|
||||
"RL_FUNCTIONS",
|
||||
]
|
||||
|
||||
import re
|
||||
import inspect
|
||||
from collections import defaultdict
|
||||
RL_EXTRA_ARGS = defaultdict(list)
|
||||
RL_FUNCTIONS = defaultdict(list)
|
||||
|
||||
|
||||
# Check untrained tokens
|
||||
def sft_trainer_fix_untraiend_tokens(call_args, extra_args):
|
||||
if "model" in call_args and "train_dataset" in call_args:
|
||||
fix_tokenizer = \
|
||||
"IGNORED_TOKENIZER_NAMES = os.environ.get('UNSLOTH_IGNORED_TOKENIZER_NAMES', '').split('\\n')\n"\
|
||||
"from unsloth_zoo.tokenizer_utils import fix_untrained_tokens\n"\
|
||||
"from unsloth_zoo.training_utils import fix_zero_training_loss\n"\
|
||||
"if 'tokenizer' not in locals(): tokenizer = processing_class\n"\
|
||||
"fix_untrained_tokens(model, tokenizer, train_dataset, IGNORED_TOKENIZER_NAMES, eps = 1e-16)\n"\
|
||||
"fix_zero_training_loss(model, tokenizer, train_dataset)\n"
|
||||
return fix_tokenizer
|
||||
return ""
|
||||
pass
|
||||
RL_EXTRA_ARGS["sft_trainer"].append(sft_trainer_fix_untraiend_tokens)
|
||||
|
||||
|
||||
# Remove DPO columns which might randomnly be tokenized
|
||||
def dpo_trainer_fix_columns(call_args, extra_args):
|
||||
if "model" in call_args and "train_dataset" in call_args:
|
||||
fix_dpo = \
|
||||
"if hasattr(train_dataset, 'column_names'):\n"\
|
||||
" column_names = set(train_dataset.column_names)\n"\
|
||||
" check = ['chosen', 'rejected', 'prompt', 'chosen_input_ids', 'chosen_attention_mask',\n"\
|
||||
" 'chosen_labels', 'rejected_input_ids', 'rejected_attention_mask', 'rejected_labels',\n"\
|
||||
" 'prompt_input_ids', 'prompt_attention_mask']\n"\
|
||||
" if all(x in column_names for x in check):\n"\
|
||||
" train_dataset = train_dataset.remove_columns(['chosen', 'rejected', 'prompt'])\n"\
|
||||
" del check, column_names\n"
|
||||
return fix_dpo
|
||||
return ""
|
||||
pass
|
||||
RL_EXTRA_ARGS["dpo_trainer"].append(dpo_trainer_fix_columns)
|
||||
|
||||
|
||||
# Fix tokenizer double BOS
|
||||
def sft_trainer_prepare_dataset(function_name, function):
|
||||
if function_name != "_prepare_non_packed_dataloader" and \
|
||||
function_name != "_prepare_dataset": return function
|
||||
|
||||
check_text = \
|
||||
"if 'tokenizer' not in locals(): tokenizer = processing_class\n"\
|
||||
"if 'formatting_func' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `formatting_func` does not exist!')\n"\
|
||||
"if 'dataset_text_field' not in locals() and 'args' in locals(): dataset_text_field = args.dataset_text_field\n"\
|
||||
"if 'dataset_text_field' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `dataset_text_field` does not exist!')\n"\
|
||||
"test_text = dataset[0][dataset_text_field] if (formatting_func is None and dataset_text_field is not None) else formatting_func(dataset[0])[0]\n"\
|
||||
"chat_template = getattr(tokenizer, 'chat_template', None)\n"\
|
||||
"chat_template = '' if chat_template is None else chat_template\n"\
|
||||
"has_bos_token_already = (test_text.startswith(tokenizer.bos_token) or tokenizer.bos_token in chat_template) "\
|
||||
"if getattr(tokenizer, 'bos_token', None) is not None else False\n"\
|
||||
"if 'add_special_tokens' not in locals() and has_bos_token_already:\n"\
|
||||
" from functools import partial\n"\
|
||||
" tokenizer = partial(tokenizer, add_special_tokens = False)\n"\
|
||||
" processing_class = tokenizer\n"\
|
||||
"else:\n"\
|
||||
" add_special_tokens = False if has_bos_token_already else add_special_tokens\n"
|
||||
|
||||
check_text = check_text.split("\n")
|
||||
check_text = "\n".join(" "*8 + x for x in check_text)
|
||||
check_text = check_text.rstrip() + "\n"
|
||||
|
||||
# .*? matches first match. .+? matches final match.
|
||||
replacer = re.findall(
|
||||
r"def {function_name}\(.*?\).*?\:\n",
|
||||
function,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(replacer) != 0:
|
||||
replacer = replacer[0]
|
||||
function = function.replace(replacer, replacer + check_text)
|
||||
pass
|
||||
return function
|
||||
pass
|
||||
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_prepare_dataset)
|
||||
|
||||
|
||||
# Ignore mean_token_accuracy since it needs logits
|
||||
# We override it directly with our version
|
||||
def _sft_trainer_compute_loss(self, model, inputs, return_outputs = False, num_items_in_batch = None):
|
||||
(loss, outputs) = super().compute_loss(
|
||||
model,
|
||||
inputs,
|
||||
return_outputs = return_outputs,
|
||||
num_items_in_batch = num_items_in_batch,
|
||||
)
|
||||
return (loss, outputs) if return_outputs else loss
|
||||
pass
|
||||
|
||||
def sft_trainer_compute_loss(function_name, function):
|
||||
if function_name != "compute_loss": return function
|
||||
|
||||
function = inspect.getsource(_sft_trainer_compute_loss)
|
||||
function = function.replace("def _sft_trainer_compute_loss", "def compute_loss")
|
||||
function = function.split("\n")
|
||||
function = "\n".join(" "*4+x for x in function)
|
||||
return function
|
||||
pass
|
||||
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_compute_loss)
|
||||
|
||||
|
||||
# Autocast precision for GRPO
|
||||
def grpo_trainer__prepare_inputs(function_name, function):
|
||||
if function_name != "_prepare_inputs": return function
|
||||
|
||||
if "with torch.inference_mode()" not in function: return function
|
||||
|
||||
# Add mixed precision training
|
||||
function = function.replace(
|
||||
"with torch.inference_mode():",
|
||||
|
||||
"with torch.inference_mode(), "\
|
||||
"torch.amp.autocast(device_type = 'cuda', "\
|
||||
"dtype = torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16) "\
|
||||
"if not torch.is_autocast_enabled('cuda') else nullcontext():",
|
||||
)
|
||||
|
||||
# Disable attaching a float32 conversion hook which upcasts logits to FP32
|
||||
function = function.replace(
|
||||
"self.accelerator.unwrap_model(self.model)",
|
||||
"self.accelerator.unwrap_model(self.model, keep_fp32_wrapper = False)",
|
||||
)
|
||||
return function
|
||||
pass
|
||||
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__prepare_inputs)
|
||||
|
||||
|
||||
# Remove _move_model_to_vllm
|
||||
def grpo_trainer__move_model_to_vllm(function_name, function):
|
||||
if function_name != "_move_model_to_vllm": return function
|
||||
|
||||
# .*? matches first match. .+? matches final match.
|
||||
replacement = "def _move_model_to_vllm(self, *args, **kwargs): return None\n"
|
||||
return " "*function.find("def") + replacement
|
||||
pass
|
||||
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__move_model_to_vllm)
|
||||
|
||||
|
||||
# Edit _get_per_token_logps to handle mixed precision
|
||||
def grpo_trainer__get_per_token_logps(function_name, function):
|
||||
if function_name != "_get_per_token_logps": return function
|
||||
|
||||
# Edit model to autocast it
|
||||
# .*? matches first match. .+? matches final match.
|
||||
original = re.findall(
|
||||
r"\n([ ]{4,})(logits = model\(.*?\))",
|
||||
function,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(original) != 0:
|
||||
spaces, original = original[0]
|
||||
spaces = len(spaces)
|
||||
replacer = \
|
||||
"if not hasattr(self, '_autocast_dtype'):\n" + \
|
||||
" "*(spaces + 4) + "self._autocast_dtype = torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16\n" + \
|
||||
" "*(spaces + 0) + "with torch.amp.autocast(device_type = 'cuda', dtype = self._autocast_dtype):\n" + \
|
||||
" "*(spaces + 4) + original
|
||||
function = function.replace(original, replacer)
|
||||
pass
|
||||
return function
|
||||
pass
|
||||
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__get_per_token_logps)
|
||||
|
|
@ -59,6 +59,7 @@ IGNORED_TOKENIZER_NAMES = frozenset(
|
|||
[x.lower() for x in IGNORED_TOKENIZER_NAMES] + \
|
||||
[x.lower()+"-bnb-4bit" for x in IGNORED_TOKENIZER_NAMES]
|
||||
)
|
||||
os.environ["UNSLOTH_IGNORED_TOKENIZER_NAMES"] = "\n".join(IGNORED_TOKENIZER_NAMES)
|
||||
|
||||
# Check environments
|
||||
keynames = "\n" + "\n".join(os.environ.keys())
|
||||
|
|
@ -907,44 +908,25 @@ except:
|
|||
pass
|
||||
|
||||
|
||||
def patch_trl_tokenizer_processing_class(trainer_name):
|
||||
# New TRL removes tokenizer!
|
||||
# We return it back!
|
||||
exec(f"from trl import {trainer_name}", globals())
|
||||
if str(eval(f"{trainer_name}").__name__).startswith("Unsloth"): return None
|
||||
parameters = eval(f"inspect.signature({trainer_name}).parameters")
|
||||
if "tokenizer" in parameters: return None
|
||||
|
||||
args = {
|
||||
key : \
|
||||
value.default \
|
||||
if type(value.default) is not str else \
|
||||
f"'{value.default}'" \
|
||||
for key, value in parameters.items()
|
||||
}
|
||||
args["tokenizer"] = None
|
||||
new_args = args.copy()
|
||||
del new_args["tokenizer"]
|
||||
del new_args["processing_class"]
|
||||
new_args = ",\n".join(f"{' '*12}{key} = {key}" for key in new_args) + \
|
||||
f",\n{' '*12}processing_class = tokenizer if tokenizer else processing_class"
|
||||
args = ",\n".join(f"{' '*8}{key} = {value}" for key, value in args.items())
|
||||
args = f"def __init__(\n" + f"{' '*8}self,\n" + args + "):"
|
||||
args += f"\n{' '*8}\n{' '*8}super().__init__(\n{new_args}\n{' '*8})"
|
||||
new_class = f"""class Unsloth{trainer_name}({trainer_name}):\n{' '*4}{args}\n"""
|
||||
return new_class
|
||||
pass
|
||||
|
||||
|
||||
def patch_sft_trainer_tokenizer():
|
||||
"""
|
||||
Patches the trainer with changes
|
||||
"""
|
||||
for function_name, replacer in (
|
||||
("_prepare_non_packed_dataloader", "def tokenize(element):",),
|
||||
try:
|
||||
sft_trainer = eval(f"trl.trainer.sft_trainer.SFTTrainer")
|
||||
except:
|
||||
return
|
||||
all_imports = dir(trl.trainer.sft_trainer)
|
||||
|
||||
for (function_name, replacer,) in (
|
||||
# ("_prepare_non_packed_dataloader", "def tokenize(element):",),
|
||||
("_prepare_non_packed_dataloader", None,),
|
||||
("_prepare_dataset", None,),
|
||||
# ("_prepare_packed_dataloader", "if dataset_text_field is not None",),
|
||||
):
|
||||
function = getsource(eval(f"trl.trainer.sft_trainer.SFTTrainer.{function_name}"))
|
||||
if not hasattr(sft_trainer, function_name): continue
|
||||
|
||||
function = getsource(eval(f"sft_trainer.{function_name}"))
|
||||
where = function.find("def")
|
||||
function = function.split("\n")
|
||||
function = "\n".join(x[where:] for x in function)
|
||||
|
|
@ -953,20 +935,41 @@ def patch_sft_trainer_tokenizer():
|
|||
"\n"\
|
||||
"if 'tokenizer' not in locals(): tokenizer = processing_class\n"\
|
||||
"if 'formatting_func' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `formatting_func` does not exist!')\n"\
|
||||
"if 'dataset_text_field' not in locals() and 'args' in locals(): dataset_text_field = args.dataset_text_field\n"\
|
||||
"if 'dataset_text_field' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `dataset_text_field` does not exist!')\n"\
|
||||
"test_text = dataset[0][dataset_text_field] if (formatting_func is None and dataset_text_field is not None) else formatting_func(dataset[0])[0]\n"\
|
||||
"chat_template = getattr(tokenizer, 'chat_template', None)\n"\
|
||||
"chat_template = '' if chat_template is None else chat_template\n"\
|
||||
"has_bos_token_already = (test_text.startswith(tokenizer.bos_token) or tokenizer.bos_token in chat_template) "\
|
||||
"if getattr(tokenizer, 'bos_token', None) is not None else False\n"\
|
||||
"add_special_tokens = False if has_bos_token_already else add_special_tokens\n\n"
|
||||
"if 'add_special_tokens' not in locals() and has_bos_token_already:\n"\
|
||||
" from functools import partial\n"\
|
||||
" tokenizer = partial(tokenizer, add_special_tokens = False)\n"\
|
||||
" processing_class = tokenizer\n"\
|
||||
"else:\n"\
|
||||
" add_special_tokens = False if has_bos_token_already else add_special_tokens\n\n"
|
||||
|
||||
check_text = check_text.split("\n")
|
||||
check_text = "\n".join(" "*where + x for x in check_text)
|
||||
check_text = check_text.rstrip() + "\n"
|
||||
|
||||
function = function.replace(replacer, check_text + replacer)
|
||||
exec(function, globals())
|
||||
if replacer is None:
|
||||
# .*? matches first match. .+? matches final match.
|
||||
replacer = re.findall(
|
||||
f"def {function_name}\(.*?\).*?\:\n",
|
||||
function,
|
||||
flags = re.MULTILINE | re.DOTALL,
|
||||
)
|
||||
if len(replacer) == 0: continue
|
||||
replacer = replacer[0]
|
||||
function = function.replace(replacer, replacer + check_text)
|
||||
else:
|
||||
function = function.replace(replacer, check_text + replacer)
|
||||
pass
|
||||
|
||||
x = [x for x in all_imports if x in function]
|
||||
exec(f"from trl.trainer.sft_trainer import ({','.join(x)})", locals())
|
||||
exec(function, locals(), globals())
|
||||
exec(f"trl.trainer.sft_trainer.SFTTrainer.{function_name} = {function_name}", globals())
|
||||
pass
|
||||
|
||||
|
|
@ -1053,16 +1056,5 @@ def patch_sft_trainer_tokenizer():
|
|||
pass
|
||||
pass
|
||||
|
||||
# Fix TRL trainers with removed tokenizer args (got replaced with processing_class)
|
||||
for trainer_name in ("SFTTrainer", "DPOTrainer", "KTOTrainer"):
|
||||
trainer_text = patch_trl_tokenizer_processing_class(trainer_name)
|
||||
if trainer_text is None: continue
|
||||
try:
|
||||
exec(trainer_text, globals())
|
||||
except:
|
||||
raise RuntimeError(f"Unsloth: Please file a bug report! Error patching {trainer_name}")
|
||||
exec(f"trl.trainer.{trainer_name} = Unsloth{trainer_name}", globals())
|
||||
pass
|
||||
|
||||
# FInally patch TRL tokenizer things
|
||||
patch_sft_trainer_tokenizer()
|
||||
# Finally patch TRL tokenizer things -> moved to RL
|
||||
# patch_sft_trainer_tokenizer()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue