From 4b385df264205d11696eac23e2dc9b85c87b69c9 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Fri, 14 Feb 2025 14:56:37 -0800 Subject: [PATCH] Selective Log softmax --- unsloth/models/rl.py | 17 +++----------- unsloth/models/rl_replacements.py | 38 ++++--------------------------- 2 files changed, 7 insertions(+), 48 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 128725a0a9..58b6d8271b 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -24,11 +24,13 @@ import re import torch from unsloth_zoo.compiler import create_new_function from unsloth_zoo.logging_utils import PatchRLStatistics +from unsloth_zoo.rl_replacements import RL_REPLACEMENTS from .rl_replacements import ( RL_EXTRA_ARGS, RL_FUNCTIONS, RL_PRE_ITEMS, ) +selective_log_softmax = RL_REPLACEMENTS["selective_log_softmax"] torch_compile_options = { "epilogue_fusion" : True, @@ -84,19 +86,6 @@ def PatchRL(FastLanguageModel): pass -# https://github.com/huggingface/trl/blob/main/trl/trainer/utils.py#L1674 -@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,) -def selective_log_softmax(logits, index): - logits = logits.to(torch.float32) - selected_logits = torch.gather(logits, dim=-1, index=index.unsqueeze(-1)).squeeze(-1) - # loop to reduce peak mem consumption - # logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits]) - logsumexp_values = torch.logsumexp(logits, dim = -1) - per_token_logps = selected_logits - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x) - return per_token_logps -pass - - RLTrainer_replacement = ''' import os from typing import * @@ -420,7 +409,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): # Selective log softmax selective_log_softmax_code = inspect.getsource(selective_log_softmax) - + # Get final source code RLTrainer_source = RLTrainer_replacement.format( RLTrainer_name = RLTrainer_name, diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index c4a52987ab..d01f6cd45f 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -22,6 +22,7 @@ import re import torch import inspect from collections import defaultdict +from unsloth_zoo.rl_replacements import RL_REPLACEMENTS RL_EXTRA_ARGS = defaultdict(list) RL_FUNCTIONS = defaultdict(list) RL_PRE_ITEMS = defaultdict(list) @@ -193,45 +194,14 @@ def grpo_trainer__get_per_token_logps(function_name, function): pass RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__get_per_token_logps) - -# Custom compiled GRPO loss - creates 3 Triton kernels -@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,) -def grpo_compute_loss(old_logits, new_logits, input_ids, mask, beta, advantages): - old_logits = old_logits.to(torch.float32) - new_logits = new_logits.to(torch.float32) - input_ids = input_ids.unsqueeze(-1) - - # x_i - logsumexp(x_i) - old_x = torch.gather(old_logits, dim = -1, index = input_ids).squeeze(-1) - new_x = torch.gather(new_logits, dim = -1, index = input_ids).squeeze(-1) - old = old_x - torch.logsumexp(old_logits, dim = -1) - new = new_x - torch.logsumexp(new_logits, dim = -1) - - kl_i = torch.exp(old - new) - (old - new) - 1.0 - loss_i = torch.exp(new - new.detach()) * advantages.unsqueeze(1) - loss_i = -(loss_i - beta * kl_i) - - mask = mask.to(torch.float32) - n_mask_per_reward = mask.sum(1) - loss_per_reward = (loss_i * mask).sum(1) / n_mask_per_reward - loss = loss_per_reward.mean() - - # Get metrics as well which are folded - with torch.inference_mode(): - completion_length = n_mask_per_reward.mean() - mean_kl_per_reward = (kl_i * mask).sum(1) / n_mask_per_reward - mean_kl = mean_kl_per_reward.mean() - pass - return loss, completion_length, mean_kl -pass -RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource((grpo_compute_loss))) - +grpo_compute_loss = RL_REPLACEMENTS["grpo_compute_loss"] +RL_PRE_ITEMS["grpo_trainer"].append(inspect.getsource(grpo_compute_loss)) # Edit _get_per_token_logps to handle mixed precision def grpo_trainer_compute_loss(function_name, function): if function_name != "compute_loss": return function - def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None): + def compute_loss(self, model, inputs, return_outputs = False, num_items_in_batch = None): if return_outputs: raise ValueError("The GRPOTrainer does not support returning outputs") # Compute the per-token log probabilities for the model