Merge branch 'main' of https://github.com/unslothai/unsloth
This commit is contained in:
commit
810171d82c
3 changed files with 115 additions and 16 deletions
|
|
@ -183,6 +183,7 @@ selective_log_softmax = RL_REPLACEMENTS["selective_log_softmax"]
|
|||
calculate_pad_tokens_in_prompt = RL_REPLACEMENTS["calculate_pad_tokens_in_prompt"]
|
||||
create_completion_attention_mask = RL_REPLACEMENTS["create_completion_attention_mask"]
|
||||
left_pack_padding = RL_REPLACEMENTS["left_pack_padding"]
|
||||
align_logprobs_with_mask = RL_REPLACEMENTS["align_logprobs_with_mask"]
|
||||
|
||||
RLTrainer_replacement = '''
|
||||
import os
|
||||
|
|
@ -225,6 +226,7 @@ torch_compile_options = {{
|
|||
{calculate_pad_tokens_in_prompt_code}
|
||||
{create_completion_attention_mask_code}
|
||||
{left_pack_padding_code}
|
||||
{align_logprobs_with_mask_code}
|
||||
|
||||
{RL_pre}
|
||||
|
||||
|
|
@ -758,7 +760,6 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
" mask_truncated_completions = True\n"\
|
||||
" epsilon_high = 0.28\n"\
|
||||
" beta = 0.0\n"\
|
||||
" loss_type = 'bnpo'\n"\
|
||||
"\n"
|
||||
extra_args += check_dr_grpo
|
||||
pass
|
||||
|
|
@ -830,7 +831,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
calculate_pad_tokens_in_prompt_code = inspect.getsource(calculate_pad_tokens_in_prompt)
|
||||
create_completion_attention_mask_code = inspect.getsource(create_completion_attention_mask)
|
||||
left_pack_padding_code = inspect.getsource(left_pack_padding)
|
||||
|
||||
align_logprobs_with_mask_code = inspect.getsource(align_logprobs_with_mask)
|
||||
# Get final source code
|
||||
RLTrainer_source = RLTrainer_replacement.format(
|
||||
RLTrainer_name = RLTrainer_name,
|
||||
|
|
@ -859,6 +860,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
calculate_pad_tokens_in_prompt_code = calculate_pad_tokens_in_prompt_code,
|
||||
create_completion_attention_mask_code = create_completion_attention_mask_code,
|
||||
left_pack_padding_code = left_pack_padding_code,
|
||||
align_logprobs_with_mask_code = align_logprobs_with_mask_code,
|
||||
)
|
||||
|
||||
if RLTrainer_name == "SFTTrainer":
|
||||
|
|
@ -942,16 +944,20 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import
|
|||
"if hasattr(model, 'vllm_engine') and hasattr(args, 'use_vllm'):\n" + \
|
||||
" " * 12 + "if (getattr(args, 'use_vllm', False) == False):\n" + \
|
||||
" " * 16 + "args.use_vllm = True\n"
|
||||
#" " * 16 + "args.vllm_importance_sampling_correction = True\n" + \
|
||||
#" " * 16 + "args.vllm_importance_sampling_cap = 2.0\n"
|
||||
|
||||
if "grpo" in trainer_file and trl_version >= Version("0.18.0"):
|
||||
# If model has vllm_engine, then use vllm in colocate mode. Donot wait for server
|
||||
vllm_setter += \
|
||||
" " * 12 + "args.vllm_mode='colocate'\n"
|
||||
|
||||
" " * 12 + "args.vllm_mode='colocate'\n"
|
||||
|
||||
init = init.replace(replacer, replacer + vllm_setter)
|
||||
pass
|
||||
pass
|
||||
|
||||
#breakpoint()
|
||||
|
||||
vllm_part = re.findall(
|
||||
r"(\n[\s]{8}"\
|
||||
r"if (self|args)\.use_vllm\:.*?"\
|
||||
|
|
@ -960,6 +966,7 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import
|
|||
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
|
||||
|
|
|
|||
|
|
@ -217,10 +217,51 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
|
|||
batch_size = self.args.per_device_train_batch_size if mode == "train" else self.args.per_device_eval_batch_size
|
||||
if not has_images:
|
||||
# Left pad prompt before calculation old and ref hidden states
|
||||
prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id)"""
|
||||
prompt_completion_ids = left_pack_padding(prompt_completion_ids, self.processing_class.pad_token_id)
|
||||
|
||||
fast_language_model_types = ["llama", "mistral", "gemma2", "qwen2" , "qwen3", "qwen3_moe"]
|
||||
|
||||
from unsloth import FastVisionModel, FastLanguageModel, FastModel
|
||||
if prompt_inputs.get("pixel_values") is None and self.model.config.model_type not in fast_language_model_types:
|
||||
self.model = FastModel.for_training(self.model)
|
||||
else:
|
||||
self.model = FastVisionModel.for_training(self.model)"""
|
||||
|
||||
function = function.replace(line_to_replace, replacement_lines)
|
||||
|
||||
pattern_to_find = re.compile(
|
||||
r"^\s*if self\.args\.gradient_accumulation_steps % generate_every != 0 or \(\s*"
|
||||
r"self\.use_vllm and self\.vllm_importance_sampling_correction\s*"
|
||||
r"\):",
|
||||
re.MULTILINE
|
||||
)
|
||||
|
||||
replacement_text = """
|
||||
if self.args.gradient_accumulation_steps % generate_every != 0 or (
|
||||
self.use_vllm
|
||||
):"""
|
||||
# Use re.sub() to perform the replacement
|
||||
function, num_replacements = pattern_to_find.subn(replacement_text, function)
|
||||
|
||||
pattern_to_find = re.compile(
|
||||
r"(^\s*)all_logprobs = \[" # Capture indentation (group 1)
|
||||
r".*?" # Match everything inside non-greedily
|
||||
r"for output in outputs\.outputs\s*"
|
||||
r"\]",
|
||||
re.DOTALL | re.MULTILINE
|
||||
)
|
||||
|
||||
replacement_text = (
|
||||
r'\1from trl.scripts.vllm_serve import sanitize_logprob\n'
|
||||
r'\1all_logprobs = [\n'
|
||||
r'\1 [sanitize_logprob(next(iter(logprob.values()))) for logprob in output.logprobs]\n'
|
||||
r'\1 for outputs in all_outputs\n'
|
||||
r'\1 for output in outputs.outputs\n'
|
||||
r'\1]'
|
||||
)
|
||||
|
||||
function, num_replacements = pattern_to_find.subn(replacement_text, function)
|
||||
|
||||
# Always between max_prompt_length and use_vllm
|
||||
found = re.findall(
|
||||
r"\n(([ ]{8,})if self\.max_prompt_length is not None:.*?"\
|
||||
|
|
@ -259,6 +300,20 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
|
|||
if self.use_vllm:"""
|
||||
function = function.replace(replace_part, new_replacement)
|
||||
|
||||
string_to_find = """ if "image_sizes" in prompt_inputs:
|
||||
output["image_sizes"] = prompt_inputs["image_sizes"]"""
|
||||
|
||||
replacement_string = """ if "image_sizes" in prompt_inputs:
|
||||
output["image_sizes"] = prompt_inputs["image_sizes"]
|
||||
|
||||
if self.use_vllm:
|
||||
try:
|
||||
output["sampling_per_token_logps"] = sampling_per_token_logps
|
||||
except NameError:
|
||||
output["sampling_per_token_logps"] = None"""
|
||||
|
||||
function = function.replace(string_to_find, replacement_string)
|
||||
|
||||
if 'wake_up()' not in function:
|
||||
# Sleep functionality has been added to trl in v0.23.0. We do not want to redo this.
|
||||
# https://github.com/huggingface/trl/commit/edbe8234bc7e528f72ac76607de9d3e4753e2709
|
||||
|
|
@ -291,7 +346,6 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
|
|||
pass
|
||||
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__generate_and_score_completions)
|
||||
|
||||
|
||||
# Fix {"reasoning_effort" : "high"} not applied
|
||||
def grpo_trainer_fix_maybe_apply_chat_template(function_name, function):
|
||||
spaces = function.find("def ")
|
||||
|
|
@ -339,7 +393,6 @@ def grpo_trainer_fix_maybe_apply_chat_template(function_name, function):
|
|||
pass
|
||||
RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer_fix_maybe_apply_chat_template)
|
||||
|
||||
|
||||
# Remove _move_model_to_vllm
|
||||
def grpo_trainer__move_model_to_vllm(function_name, function):
|
||||
if function_name != "_move_model_to_vllm": return function
|
||||
|
|
@ -504,6 +557,10 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
completion_ids, completion_mask = inputs["completion_ids"], inputs["completion_mask"]
|
||||
pixel_values, image_grid_thw = inputs.get("pixel_values", None), inputs.get("image_grid_thw", None)
|
||||
pixel_attention_mask, image_sizes = inputs.get('pixel_attention_mask',None), inputs.get('image_sizes',None)
|
||||
num_items_in_batch = inputs.get("num_items_in_batch", None)
|
||||
sampling_per_token_logps = inputs.get("sampling_per_token_logps", None)
|
||||
current_gradient_accumulation_steps = self.current_gradient_accumulation_steps
|
||||
num_processes = self.accelerator.num_processes
|
||||
|
||||
input_ids = torch.cat([prompt_ids, completion_ids], dim=1)
|
||||
bsz, qlen = input_ids.shape
|
||||
|
|
@ -518,7 +575,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep, compute_efficient) \
|
||||
if hasattr(self, "_get_per_token_logps") else \
|
||||
self._get_per_token_logps_and_entropies(model, input_ids, attention_mask, logits_to_keep, batch_size, compute_entropy, compute_efficient)[0] # logps
|
||||
#breakpoint()
|
||||
|
||||
per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep, compute_efficient = True)
|
||||
# Compute the KL divergence between the model and the reference model
|
||||
# _prepare_inputs doesn't return reference log probs anymore. We need to calculate it ourselves.
|
||||
|
|
@ -555,7 +612,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
old_hidden_states = old_hidden_states[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
per_token_logps = per_token_logps[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
|
||||
loss, completion_length, mean_kl = grpo_compute_loss_slow(
|
||||
loss, completion_length, mean_kl, delta, flat_is_ratio = grpo_compute_loss_slow(
|
||||
ref_hidden_states,
|
||||
per_token_logps,
|
||||
old_hidden_states,
|
||||
|
|
@ -575,10 +632,14 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
logit_softcapping = logit_softcapping,
|
||||
logit_scale_multiply = logit_scale_multiply,
|
||||
logit_scale_divide = logit_scale_divide,
|
||||
num_items_in_batch = num_items_in_batch,
|
||||
current_gradient_accumulation_steps = current_gradient_accumulation_steps,
|
||||
num_processes = num_processes,
|
||||
sampling_per_token_logps = sampling_per_token_logps,
|
||||
)
|
||||
else:
|
||||
if hasattr(self.args, "loss_type"):
|
||||
loss, completion_length, mean_kl = grpo_accumulated_loss(
|
||||
loss, completion_length, mean_kl, delta, flat_is_ratio = grpo_accumulated_loss(
|
||||
trainer = self,
|
||||
input_ids = _input_ids,
|
||||
pixel_values = pixel_values,
|
||||
|
|
@ -600,6 +661,10 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
logit_scale_multiply = logit_scale_multiply,
|
||||
logit_scale_divide = logit_scale_divide,
|
||||
attention_mask = attention_mask,
|
||||
num_items_in_batch = num_items_in_batch,
|
||||
current_gradient_accumulation_steps = current_gradient_accumulation_steps,
|
||||
num_processes = num_processes,
|
||||
sampling_per_token_logps = sampling_per_token_logps,
|
||||
)
|
||||
else:
|
||||
# to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17
|
||||
|
|
@ -620,10 +685,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
)
|
||||
pass
|
||||
pass
|
||||
# Log the metrics
|
||||
# completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item()
|
||||
# mean_kl = ((per_token_kl * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean()
|
||||
# self._metrics["kl"].append(self.accelerator.gather_for_metrics(mean_kl).mean().item())
|
||||
|
||||
if "train" in self._metrics:
|
||||
mode = "eval" if self.control.should_evaluate else "train"
|
||||
self._metrics[mode]["completion_length"].append(completion_length.item())
|
||||
|
|
@ -631,6 +693,36 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
else:
|
||||
self._metrics["completion_length"].append(completion_length.item())
|
||||
self._metrics["kl"].append(mean_kl.item())
|
||||
|
||||
if self.use_vllm and delta is not None:
|
||||
mean_delta = torch.mean(delta) if delta.numel() > 0 else torch.tensor(0.0, device=self.model.device)
|
||||
max_delta = torch.max(delta) if delta.numel() > 0 else torch.tensor(0.0, device=self.model.device)
|
||||
self._metrics[mode]["sampling/sampling_logp_difference/mean"].append(
|
||||
self.accelerator.gather(mean_delta).mean().item()
|
||||
)
|
||||
self._metrics[mode]["sampling/sampling_logp_difference/max"].append(
|
||||
self.accelerator.gather(max_delta).max().item()
|
||||
)
|
||||
|
||||
min_importance_sampling_ratio = (
|
||||
torch.min(flat_is_ratio) if flat_is_ratio.numel() > 0 else torch.tensor(0.0, device=self.model.device)
|
||||
)
|
||||
mean_importance_sampling_ratio = (
|
||||
torch.mean(flat_is_ratio) if flat_is_ratio.numel() > 0 else torch.tensor(0.0, device=self.model.device)
|
||||
)
|
||||
max_importance_sampling_ratio = (
|
||||
torch.max(flat_is_ratio) if flat_is_ratio.numel() > 0 else torch.tensor(0.0, device=self.model.device)
|
||||
)
|
||||
self._metrics[mode]["sampling/importance_sampling_ratio/min"].append(
|
||||
nanmin(self.accelerator.gather(min_importance_sampling_ratio)).item()
|
||||
)
|
||||
self._metrics[mode]["sampling/importance_sampling_ratio/mean"].append(
|
||||
self.accelerator.gather(mean_importance_sampling_ratio).nanmean().item()
|
||||
)
|
||||
self._metrics[mode]["sampling/importance_sampling_ratio/max"].append(
|
||||
nanmax(self.accelerator.gather(max_importance_sampling_ratio)).item()
|
||||
)
|
||||
|
||||
return loss
|
||||
pass
|
||||
|
||||
|
|
|
|||
|
|
@ -640,13 +640,13 @@ class FastBaseModel:
|
|||
_, quant_state_dict = get_vllm_state_dict(
|
||||
llm,
|
||||
config = model_config,
|
||||
is_vision_model = True,
|
||||
is_vision_model = is_vlm,
|
||||
)
|
||||
model = convert_vllm_to_huggingface(
|
||||
quant_state_dict,
|
||||
model_config,
|
||||
dtype, bnb_config,
|
||||
is_vision_model = True,
|
||||
is_vision_model = is_vlm,
|
||||
)
|
||||
model.vllm_engine = llm
|
||||
model.fast_generate = model.vllm_engine.generate
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue