Merge pull request #3628 from pluesclues/alternative_compute_chunked_loss
Chunk Across Batch and Context length for logprob calculations for grpo
This commit is contained in:
parent
4fc06bd7fb
commit
e83cbc9fe0
2 changed files with 267 additions and 64 deletions
|
|
@ -231,11 +231,13 @@ def PatchRL(FastLanguageModel):
|
|||
Trainer.prediction_step = unsloth_prediction_step
|
||||
|
||||
|
||||
grpo_selective_log_softmax = RL_REPLACEMENTS["grpo_selective_log_softmax"]
|
||||
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"]
|
||||
autotune_batch_and_chunks = RL_REPLACEMENTS["grpo_autotune_batch_and_chunks"]
|
||||
|
||||
RLTrainer_replacement = '''
|
||||
import os
|
||||
|
|
@ -247,7 +249,6 @@ import numpy as np
|
|||
from contextlib import nullcontext
|
||||
from torch.nn import functional as F
|
||||
import inspect
|
||||
import psutil
|
||||
from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling as TransformersDataCollatorForLanguageModeling
|
||||
from transformers.training_args import ParallelMode
|
||||
|
||||
|
|
@ -300,11 +301,13 @@ torch_compile_options = {{
|
|||
"triton.cudagraphs" : False,
|
||||
}}
|
||||
|
||||
{grpo_selective_log_softmax_code}
|
||||
{selective_log_softmax_code}
|
||||
{calculate_pad_tokens_in_prompt_code}
|
||||
{create_completion_attention_mask_code}
|
||||
{left_pack_padding_code}
|
||||
{align_logprobs_with_mask_code}
|
||||
{autotune_batch_and_chunks_code}
|
||||
|
||||
{RL_pre}
|
||||
|
||||
|
|
@ -321,10 +324,20 @@ class Unsloth{RLConfig_name}({RLConfig_name}):
|
|||
default = -1,
|
||||
metadata = {{'help': 'Chunk size to reduce memory usage. -1 is most efficient.'}},
|
||||
)
|
||||
unsloth_logit_chunk_multiplier : Optional[int] = field(
|
||||
default = None,
|
||||
metadata = {{'help': 'Multiplier for chunked logit computations.'}},
|
||||
)
|
||||
unsloth_grpo_mini_batch : Optional[int] = field(
|
||||
default = None,
|
||||
metadata = {{'help': 'Mini batch size for GRPO hidden state accumulation. Default is None unless user defines it.'}},
|
||||
)
|
||||
{max_seq_length_pre}
|
||||
def __init__({RLConfig_arguments},
|
||||
vllm_sampling_params = None,
|
||||
unsloth_num_chunks = -1,
|
||||
unsloth_logit_chunk_multiplier = None,
|
||||
unsloth_grpo_mini_batch = None,
|
||||
{max_seq_length_call}
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -332,6 +345,15 @@ class Unsloth{RLConfig_name}({RLConfig_name}):
|
|||
super().__init__({RLConfig_call_args}{RLConfig_kwargs})
|
||||
self.vllm_sampling_params = vllm_sampling_params
|
||||
self.unsloth_num_chunks = unsloth_num_chunks
|
||||
if unsloth_grpo_mini_batch is not None:
|
||||
if self.generation_batch_size >= unsloth_grpo_mini_batch:
|
||||
self.unsloth_grpo_mini_batch = unsloth_grpo_mini_batch
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsloth GRPO mini batch size needs to be less than or equal to the effective generation batch size, "
|
||||
f"which is self.per_device_train_batch_size * gradient_accumulation_steps."
|
||||
)
|
||||
self.unsloth_logit_chunk_multiplier = unsloth_logit_chunk_multiplier
|
||||
{max_seq_length_post}
|
||||
pass
|
||||
|
||||
|
|
@ -1029,6 +1051,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
|
||||
# Selective log softmax and other functions
|
||||
selective_log_softmax_code = inspect.getsource(selective_log_softmax)
|
||||
grpo_selective_log_softmax_code = inspect.getsource(grpo_selective_log_softmax)
|
||||
calculate_pad_tokens_in_prompt_code = inspect.getsource(
|
||||
calculate_pad_tokens_in_prompt
|
||||
)
|
||||
|
|
@ -1037,6 +1060,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
)
|
||||
left_pack_padding_code = inspect.getsource(left_pack_padding)
|
||||
align_logprobs_with_mask_code = inspect.getsource(align_logprobs_with_mask)
|
||||
autotune_batch_and_chunks_code = inspect.getsource(autotune_batch_and_chunks)
|
||||
# Get final source code
|
||||
RLTrainer_source = RLTrainer_replacement.format(
|
||||
RLTrainer_name = RLTrainer_name,
|
||||
|
|
@ -1058,8 +1082,10 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"):
|
|||
max_seq_length_call = max_seq_length_call,
|
||||
max_seq_length_post = max_seq_length_post,
|
||||
selective_log_softmax_code = selective_log_softmax_code,
|
||||
grpo_selective_log_softmax_code = grpo_selective_log_softmax_code,
|
||||
calculate_pad_tokens_in_prompt_code = calculate_pad_tokens_in_prompt_code,
|
||||
create_completion_attention_mask_code = create_completion_attention_mask_code,
|
||||
autotune_batch_and_chunks_code = autotune_batch_and_chunks_code,
|
||||
left_pack_padding_code = left_pack_padding_code,
|
||||
align_logprobs_with_mask_code = align_logprobs_with_mask_code,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -50,7 +50,7 @@ RL_ADDITIONAL_FUNCTIONS = defaultdict(list)
|
|||
|
||||
torch_compile_options = {
|
||||
"epilogue_fusion": True,
|
||||
"max_autotune": True,
|
||||
"max_autotune": False, # I saw speedups, but not sure if this has issues in collab
|
||||
"shape_padding": True,
|
||||
"trace.enabled": False,
|
||||
"triton.cudagraphs": False,
|
||||
|
|
@ -258,18 +258,20 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
|
|||
|
||||
# The new multi-line string that will replace the line above
|
||||
replacement_lines = """
|
||||
max_left_pad = None
|
||||
batch_size = self.args.per_device_train_batch_size if mode == "train" else self.args.per_device_eval_batch_size
|
||||
try:
|
||||
# TRL 0.23.1 and below path
|
||||
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)
|
||||
self.model.for_training()
|
||||
left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt(prompt_completion_ids, logits_to_keep, self.processing_class.pad_token_id)
|
||||
max_left_pad = torch.max(left_pad_tokens_per_prompt).item()
|
||||
except:
|
||||
# TRL 0.24.0 and below path
|
||||
if images is None:
|
||||
# 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)
|
||||
left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt(prompt_completion_ids, logits_to_keep, self.processing_class.pad_token_id)
|
||||
max_left_pad = torch.max(left_pad_tokens_per_prompt).item()
|
||||
self.model.for_training()"""
|
||||
|
||||
function = function.replace(line_to_replace, replacement_lines)
|
||||
|
|
@ -346,17 +348,45 @@ def grpo_trainer__generate_and_score_completions(function_name, function):
|
|||
if self.use_vllm:"""
|
||||
function = function.replace(replace_part, new_replacement)
|
||||
|
||||
# Important note: we disable TRL's importance sampling logic
|
||||
# It is disabled because the LLM path moves left padding to the right.
|
||||
# We must adjust the vLLM sampling_logprob tensor in Unsloth to account for this.
|
||||
string_to_find = "if self.use_vllm and self.vllm_importance_sampling_correction:"
|
||||
|
||||
replacement_string = (
|
||||
"if False and self.use_vllm and self.vllm_importance_sampling_correction:"
|
||||
)
|
||||
|
||||
function = function.replace(string_to_find, replacement_string)
|
||||
|
||||
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:
|
||||
if max_left_pad is not None:
|
||||
output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1)
|
||||
try:
|
||||
if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False):
|
||||
output["sampling_per_token_logps"] = sampling_per_token_logps
|
||||
except NameError:
|
||||
output["sampling_per_token_logps"] = None"""
|
||||
except NameError:
|
||||
output["sampling_per_token_logps"] = None"""
|
||||
|
||||
function = function.replace(string_to_find, replacement_string)
|
||||
|
||||
# This path is for TRL 0.24.0 images is a variable exclusive to this version
|
||||
string_to_find = """ if images is not None:
|
||||
output["num_images"] = num_images"""
|
||||
|
||||
replacement_string = """ if images is not None:
|
||||
output["num_images"] = num_images
|
||||
if max_left_pad is not None:
|
||||
output["max_left_pad"] = torch.tensor(prompt_ids.shape[0] * [max_left_pad]).unsqueeze(-1)
|
||||
try:
|
||||
if self.use_vllm and getattr(self, "vllm_importance_sampling_correction", False):
|
||||
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)
|
||||
|
||||
|
|
@ -532,12 +562,12 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
|
|||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
# All Unsloth code here in this function is licensed under AGPL3
|
||||
# if True: # os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0':
|
||||
# return None, None # logps, entropies Unsloth efficient GRPO
|
||||
if compute_efficient:
|
||||
return None, None
|
||||
else:
|
||||
# Otherwise, calculate normally:
|
||||
if not hasattr(self, "_autocast_dtype"):
|
||||
self._autocast_dtype = (
|
||||
torch.float16
|
||||
|
|
@ -556,47 +586,199 @@ def grpo_trainer__get_per_token_logps_and_entropies(function_name, function):
|
|||
kwargs.get("image_sizes", None),
|
||||
)
|
||||
|
||||
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
|
||||
|
||||
unwrapped_model = self.accelerator.unwrap_model(
|
||||
model, keep_fp32_wrapper = False
|
||||
)
|
||||
|
||||
with torch.amp.autocast(device_type = "cuda", dtype = self._autocast_dtype):
|
||||
with _get_inference_mode_context_manager(model):
|
||||
if pixel_values is None:
|
||||
attention_mask = input_ids != self.processing_class.pad_token_id
|
||||
attention_mask = attention_mask.to(attention_mask.dtype)
|
||||
# We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
|
||||
logits = unwrapped_model(
|
||||
input_ids = input_ids,
|
||||
attention_mask = attention_mask,
|
||||
pixel_values = pixel_values,
|
||||
image_grid_thw = image_grid_thw,
|
||||
pixel_attention_mask = pixel_attention_mask,
|
||||
image_sizes = image_sizes,
|
||||
# logits_to_keep = logits_to_keep + 1,
|
||||
).logits
|
||||
lm_head = self.model.get_output_embeddings().weight
|
||||
|
||||
dtype_bytes = (
|
||||
16 if self._autocast_dtype in [torch.float16, torch.bfloat16] else 32
|
||||
)
|
||||
total_rows = input_ids.shape[0]
|
||||
seq_len = input_ids.shape[1]
|
||||
hidden_dim = lm_head.shape[1]
|
||||
vocab_dim = lm_head.shape[0]
|
||||
|
||||
if self.args.unsloth_grpo_mini_batch is None:
|
||||
B, multiplier = autotune_batch_and_chunks(
|
||||
total_rows,
|
||||
seq_len,
|
||||
hidden_dim,
|
||||
vocab_dim,
|
||||
dtype_bytes,
|
||||
self.args.unsloth_logit_chunk_multiplier,
|
||||
)
|
||||
B = total_rows // B
|
||||
else:
|
||||
B = self.args.unsloth_grpo_mini_batch
|
||||
|
||||
if self.args.unsloth_logit_chunk_multiplier is None:
|
||||
multiplier = max(4, seq_len // 4096)
|
||||
else:
|
||||
multiplier = self.args.unsloth_logit_chunk_multiplier
|
||||
|
||||
all_logprobs_list = []
|
||||
if pixel_values is None:
|
||||
left_pad_tokens_per_prompt = calculate_pad_tokens_in_prompt(
|
||||
input_ids, logits_to_keep, self.processing_class.pad_token_id
|
||||
)
|
||||
max_left_pad = torch.max(left_pad_tokens_per_prompt).item()
|
||||
input_ids = left_pack_padding(
|
||||
input_ids, self.processing_class.pad_token_id
|
||||
)
|
||||
attention_mask = input_ids != self.processing_class.pad_token_id
|
||||
attention_mask = attention_mask.to(attention_mask.dtype)
|
||||
else:
|
||||
max_left_pad = 0
|
||||
|
||||
# input_ids_chunks = torch.chunk(input_ids, chunks = B, dim = 0)
|
||||
attention_mask_chunks = torch.chunk(attention_mask, chunks = B, dim = 0)
|
||||
|
||||
def chunk_optional(tensor, chunks):
|
||||
if tensor is None:
|
||||
return [None] * chunks
|
||||
return torch.chunk(tensor, chunks = chunks, dim = 0)
|
||||
|
||||
import math
|
||||
|
||||
total_samples = input_ids.shape[0]
|
||||
batch_size = math.ceil(total_samples / B)
|
||||
|
||||
input_ids_chunks = []
|
||||
attention_mask_chunks = []
|
||||
pixel_values_chunks = []
|
||||
image_grid_thw_chunks = []
|
||||
pixel_attention_mask_chunks = []
|
||||
|
||||
current_pixel_idx = 0
|
||||
# TRL 0.23.0 batching logic
|
||||
for start in range(0, total_samples, batch_size):
|
||||
end = start + batch_size
|
||||
|
||||
input_ids_chunks.append(input_ids[start:end])
|
||||
attention_mask_chunks.append(attention_mask[start:end])
|
||||
|
||||
if image_grid_thw is not None and pixel_values is not None:
|
||||
grid_slice = image_grid_thw[start:end]
|
||||
image_grid_thw_chunks.append(grid_slice)
|
||||
|
||||
batch_pixel_count = grid_slice.prod(dim = -1).sum().item()
|
||||
|
||||
start_pixel_idx = current_pixel_idx
|
||||
end_pixel_idx = current_pixel_idx + batch_pixel_count
|
||||
|
||||
pixel_values_chunks.append(
|
||||
pixel_values[start_pixel_idx:end_pixel_idx]
|
||||
)
|
||||
|
||||
if pixel_attention_mask is not None:
|
||||
pixel_attention_mask_chunks.append(
|
||||
pixel_attention_mask[start_pixel_idx:end_pixel_idx]
|
||||
)
|
||||
else:
|
||||
logits = unwrapped_model(
|
||||
input_ids = input_ids,
|
||||
attention_mask = attention_mask,
|
||||
pixel_values = pixel_values,
|
||||
image_grid_thw = image_grid_thw,
|
||||
pixel_attention_mask = pixel_attention_mask,
|
||||
image_sizes = image_sizes,
|
||||
logits_to_keep = logits_to_keep + 1,
|
||||
).logits
|
||||
pixel_attention_mask_chunks.append(None)
|
||||
|
||||
current_pixel_idx = end_pixel_idx
|
||||
|
||||
else:
|
||||
pixel_values_chunks.append(None)
|
||||
image_grid_thw_chunks.append(None)
|
||||
pixel_attention_mask_chunks.append(None)
|
||||
|
||||
if image_sizes is not None and not isinstance(image_sizes, torch.Tensor):
|
||||
image_sizes_chunks = [[size] for size in image_sizes]
|
||||
else:
|
||||
image_sizes_chunks = chunk_optional(image_sizes, B)
|
||||
|
||||
temperature = self.temperature
|
||||
logit_softcapping = getattr(model.config, "final_logit_softcapping", 0)
|
||||
if logit_softcapping is None:
|
||||
logit_softcapping = 0
|
||||
logit_scale_multiply = getattr(model.config, "logit_scale", 0)
|
||||
if logit_scale_multiply is None:
|
||||
logit_scale_multiply = 0
|
||||
logit_scale_divide = getattr(model.config, "logits_scaling", 0)
|
||||
if logit_scale_divide is None:
|
||||
logit_scale_divide = 0
|
||||
|
||||
zipped_inputs = zip(
|
||||
input_ids_chunks,
|
||||
attention_mask_chunks,
|
||||
pixel_values_chunks,
|
||||
image_grid_thw_chunks,
|
||||
pixel_attention_mask_chunks,
|
||||
image_sizes_chunks,
|
||||
)
|
||||
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1"
|
||||
|
||||
with _get_inference_mode_context_manager(model):
|
||||
for (
|
||||
input_ids_chunk,
|
||||
attention_mask_chunk,
|
||||
pixel_values_chunk,
|
||||
image_grid_thw_chunk,
|
||||
pixel_attention_mask_chunk,
|
||||
image_sizes_chunk,
|
||||
) in zipped_inputs:
|
||||
with torch.amp.autocast(
|
||||
device_type = "cuda", dtype = self._autocast_dtype
|
||||
):
|
||||
if pixel_values is None:
|
||||
logits_chunk = unwrapped_model(
|
||||
input_ids = input_ids_chunk,
|
||||
attention_mask = attention_mask_chunk,
|
||||
pixel_values = pixel_values_chunk,
|
||||
image_grid_thw = image_grid_thw_chunk,
|
||||
pixel_attention_mask = pixel_attention_mask_chunk,
|
||||
image_sizes = image_sizes_chunk,
|
||||
).logits
|
||||
|
||||
completion_input_ids_chunk = input_ids_chunk[
|
||||
:, -(logits_to_keep + max_left_pad) :
|
||||
]
|
||||
logits_chunk = logits_chunk[
|
||||
:, -(logits_to_keep + max_left_pad + 1) :, :
|
||||
]
|
||||
logits_chunk = logits_chunk[:, :-1, :]
|
||||
else:
|
||||
# Essentially, for VLMs we do not go via the optimized path in models/,
|
||||
# so we don't encounter the Flash Attn left-padding issue.
|
||||
logits_chunk = unwrapped_model(
|
||||
input_ids = input_ids_chunk,
|
||||
attention_mask = attention_mask_chunk,
|
||||
pixel_values = pixel_values_chunk,
|
||||
image_grid_thw = image_grid_thw_chunk,
|
||||
pixel_attention_mask = pixel_attention_mask_chunk,
|
||||
image_sizes = image_sizes_chunk,
|
||||
logits_to_keep = logits_to_keep + 1,
|
||||
).logits
|
||||
|
||||
logits_chunk = logits_chunk[:, :-1, :]
|
||||
completion_input_ids_chunk = input_ids_chunk[
|
||||
:, -logits_to_keep:
|
||||
]
|
||||
|
||||
logprobs_chunk = chunked_hidden_states_selective_log_softmax(
|
||||
logits_chunk,
|
||||
lm_head,
|
||||
completion_input_ids_chunk,
|
||||
chunks = input_ids_chunk.shape[0] * multiplier,
|
||||
logit_scale_multiply = logit_scale_multiply,
|
||||
logit_scale_divide = logit_scale_divide,
|
||||
logit_softcapping = logit_softcapping,
|
||||
temperature = temperature,
|
||||
)
|
||||
# This is needed to avoid race conditions with GPT OSS offload_embbed=True
|
||||
# However, it seems that this line does not slow down or disrupt models.
|
||||
torch.cuda.synchronize()
|
||||
all_logprobs_list.append(logprobs_chunk)
|
||||
logprobs = torch.cat(all_logprobs_list, dim = 0)
|
||||
entropies = None
|
||||
if compute_entropy:
|
||||
from trl.trainer.utils import entropy_from_logits
|
||||
|
||||
entropies = entropy_from_logits(logits)
|
||||
|
||||
os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "0"
|
||||
# logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
return logits.detach(), entropies # logps, entropies
|
||||
|
||||
return logprobs.detach(), entropies # logps, entropies
|
||||
# input_ids = input_ids[:, -logits_to_keep:]
|
||||
# For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves.
|
||||
# See https://github.com/huggingface/trl/issues/2770
|
||||
|
|
@ -708,14 +890,14 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
# ref_per_token_logps = per_token_logps = get_logps_func(model, input_ids, attention_mask, logits_to_keep)
|
||||
# else:
|
||||
# ref_per_token_logps = None
|
||||
ref_hidden_states = inputs.get("ref_per_token_logps", None)
|
||||
ref_logps = inputs.get("ref_per_token_logps", None)
|
||||
# per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1
|
||||
# x - x.detach() allows for preserving gradients from x
|
||||
advantages = inputs["advantages"]
|
||||
# per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)
|
||||
# per_token_loss = -(per_token_loss - self.beta * per_token_kl)
|
||||
# loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean()
|
||||
old_hidden_states = inputs.get("old_per_token_logps", None)
|
||||
old_logps = inputs.get("old_per_token_logps", None)
|
||||
|
||||
input_ids = input_ids[:, -logits_to_keep:]
|
||||
|
||||
|
|
@ -730,24 +912,13 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
if logit_scale_divide is None:
|
||||
logit_scale_divide = 0
|
||||
|
||||
max_left_pad = inputs.get("max_left_pad", 0)
|
||||
if per_token_logps is not None:
|
||||
if ref_hidden_states is not None:
|
||||
ref_hidden_states = ref_hidden_states[
|
||||
:, :-1, :
|
||||
] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
|
||||
if old_hidden_states is not None:
|
||||
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, delta, flat_is_ratio = (
|
||||
grpo_compute_loss_slow(
|
||||
ref_hidden_states,
|
||||
ref_logps,
|
||||
per_token_logps,
|
||||
old_hidden_states,
|
||||
old_logps,
|
||||
input_ids,
|
||||
completion_mask,
|
||||
self.beta,
|
||||
|
|
@ -761,6 +932,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
max_completion_length = self.args.max_completion_length,
|
||||
delta = self.args.delta,
|
||||
temperature = self.args.temperature,
|
||||
max_left_pad = max_left_pad,
|
||||
logit_softcapping = logit_softcapping,
|
||||
logit_scale_multiply = logit_scale_multiply,
|
||||
logit_scale_divide = logit_scale_divide,
|
||||
|
|
@ -781,8 +953,8 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
logits_to_keep = logits_to_keep,
|
||||
completion_mask = completion_mask,
|
||||
advantages = advantages,
|
||||
old_hidden_states = old_hidden_states,
|
||||
ref_hidden_states = ref_hidden_states,
|
||||
old_logps = old_logps,
|
||||
ref_logps = ref_logps,
|
||||
n_chunks = self.args.unsloth_num_chunks,
|
||||
loss_type = self.args.loss_type,
|
||||
importance_sampling_level = self.importance_sampling_level,
|
||||
|
|
@ -791,6 +963,7 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
max_completion_length = self.args.max_completion_length,
|
||||
delta = self.args.delta,
|
||||
temperature = self.args.temperature,
|
||||
max_left_pad = max_left_pad,
|
||||
logit_softcapping = logit_softcapping,
|
||||
logit_scale_multiply = logit_scale_multiply,
|
||||
logit_scale_divide = logit_scale_divide,
|
||||
|
|
@ -809,8 +982,8 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
logits_to_keep = logits_to_keep,
|
||||
completion_mask = completion_mask,
|
||||
advantages = advantages,
|
||||
old_hidden_states = old_hidden_states,
|
||||
ref_hidden_states = ref_hidden_states,
|
||||
old_logps = old_logps,
|
||||
ref_logps = ref_logps,
|
||||
n_chunks = self.args.unsloth_num_chunks,
|
||||
temperature = self.args.temperature,
|
||||
logit_softcapping = logit_softcapping,
|
||||
|
|
@ -827,7 +1000,11 @@ def grpo_trainer_compute_loss(function_name, function):
|
|||
self._metrics["completion_length"].append(completion_length.item())
|
||||
self._metrics["kl"].append(mean_kl.item())
|
||||
|
||||
if self.use_vllm and delta is not None:
|
||||
if (
|
||||
self.use_vllm
|
||||
and delta is not None
|
||||
and getattr(self, "vllm_importance_sampling_correction", False)
|
||||
):
|
||||
mean_delta = (
|
||||
torch.mean(delta)
|
||||
if delta.numel() > 0
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue