From 9b5a09e4cfafb8c95620c128170049b4e8461462 Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Wed, 29 Jan 2025 14:34:19 -0800 Subject: [PATCH 1/8] make sure model is in eval mode before generating --- trl/trainer/grpo_trainer.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index b0a36f7590d..5222449992a 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -419,9 +419,11 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N else: # Regular generation path with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model: + unwrapped_model.eval() # Needed to make sure use_cache works with gradient_checkpointing prompt_completion_ids = unwrapped_model.generate( **prompt_inputs, generation_config=self.generation_config ) + model.train() prompt_length = prompt_inputs["input_ids"].size(1) completion_ids = prompt_completion_ids[:, prompt_length:] From bcd74ece42860a1a21dc4f432a09f9ebca042b10 Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Wed, 29 Jan 2025 14:54:01 -0800 Subject: [PATCH 2/8] bit more logging --- trl/trainer/grpo_trainer.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index 5222449992a..da1ffc1e226 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import logging import os import textwrap import warnings @@ -55,6 +56,9 @@ if is_wandb_available(): import wandb + +logger = logging.getLogger(__name__) + # What we call a reward function is a callable that takes a list of prompts and completions and returns a list of # rewards. When it's a string, it's a model ID, so it's loaded as a pretrained model. RewardFunc = Union[str, PreTrainedModel, Callable[[list, list], list[float]]] @@ -383,6 +387,7 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N prompt_inputs["input_ids"] = prompt_inputs["input_ids"][:, -self.max_prompt_length :] prompt_inputs["attention_mask"] = prompt_inputs["attention_mask"][:, -self.max_prompt_length :] + logger.debug("Starting generation") # Generate completions using either vLLM or regular generation if self.args.use_vllm: # First, have main process load weights if needed @@ -442,9 +447,11 @@ def get_per_token_logps(model, input_ids, num_logits_to_keep): per_token_logps.append(token_log_prob) return torch.stack(per_token_logps) + logger.debug("Calculating policy logprobs") num_logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens per_token_logps = get_per_token_logps(model, prompt_completion_ids, num_logits_to_keep) + logger.debug("Calculating reference logprobs") with torch.inference_mode(): if self.ref_model is not None: ref_per_token_logps = get_per_token_logps(self.ref_model, prompt_completion_ids, num_logits_to_keep) From 8e8b98856589105ce7d73c4189b43dfb3bcfa486 Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Wed, 29 Jan 2025 14:55:37 -0800 Subject: [PATCH 3/8] switch to transformers logging --- trl/trainer/grpo_trainer.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index da1ffc1e226..448c5175d4d 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import logging import os import textwrap import warnings @@ -38,7 +37,7 @@ is_wandb_available, ) from transformers.integrations.deepspeed import is_deepspeed_zero3_enabled -from transformers.utils import is_peft_available +from transformers.utils import is_peft_available, logging from ..data_utils import apply_chat_template, is_conversational, maybe_apply_chat_template from ..import_utils import is_vllm_available @@ -57,7 +56,7 @@ import wandb -logger = logging.getLogger(__name__) +logger = logging.get_logger("GRPOTrainer") # What we call a reward function is a callable that takes a list of prompts and completions and returns a list of # rewards. When it's a string, it's a model ID, so it's loaded as a pretrained model. @@ -387,7 +386,7 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N prompt_inputs["input_ids"] = prompt_inputs["input_ids"][:, -self.max_prompt_length :] prompt_inputs["attention_mask"] = prompt_inputs["attention_mask"][:, -self.max_prompt_length :] - logger.debug("Starting generation") + logger.info("Starting generation") # Generate completions using either vLLM or regular generation if self.args.use_vllm: # First, have main process load weights if needed @@ -447,11 +446,11 @@ def get_per_token_logps(model, input_ids, num_logits_to_keep): per_token_logps.append(token_log_prob) return torch.stack(per_token_logps) - logger.debug("Calculating policy logprobs") + logger.info("Calculating policy logprobs") num_logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens per_token_logps = get_per_token_logps(model, prompt_completion_ids, num_logits_to_keep) - logger.debug("Calculating reference logprobs") + logger.info("Calculating reference logprobs") with torch.inference_mode(): if self.ref_model is not None: ref_per_token_logps = get_per_token_logps(self.ref_model, prompt_completion_ids, num_logits_to_keep) From b6a92fe0b021d7b5a6aaddff43474ce0b981360b Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Fri, 31 Jan 2025 16:26:16 -0800 Subject: [PATCH 4/8] supporting microbatchces for ocmputing loss terms --- trl/trainer/grpo_config.py | 9 +++ trl/trainer/grpo_trainer.py | 157 +++++++++++++++++++++--------------- 2 files changed, 102 insertions(+), 64 deletions(-) diff --git a/trl/trainer/grpo_config.py b/trl/trainer/grpo_config.py index 0fd0d9f5d28..814b8b3b696 100644 --- a/trl/trainer/grpo_config.py +++ b/trl/trainer/grpo_config.py @@ -174,3 +174,12 @@ class GRPOConfig(TrainingArguments): default=0.04, metadata={"help": "KL coefficient."}, ) + + per_device_loss_batch_size: int = field( + default=8, + metadata={ + "help": "Micro batch size per GPU/TPU/MPS/NPU core/CPU for computing loss terms. " + "These microbatches will be accumulated over, resulting in the same effective " + "batch size as `per_device_train_batch_size*num_generations`, but with lower mem footprint." + }, + ) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index 2c4137ece2f..7aed43b9815 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -430,76 +430,105 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N model.train() logger.info("Finished generation") + bsz = prompt_inputs["input_ids"].size(0) + loss_bsz = self.args.per_device_loss_batch_size or bsz prompt_length = prompt_inputs["input_ids"].size(1) - completion_ids = prompt_completion_ids[:, prompt_length:] - - # Get the per-token log probabilities for the completions for the model and the reference model - def get_per_token_logps(model, input_ids, logits_to_keep): - # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded - logits = model(input_ids, logits_to_keep=logits_to_keep + 1).logits # (B, L, V) - logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - - # Compute the log probabilities for the input tokens. Use a loop to reduce memory peak. - per_token_logps = [] - for logits_row, input_ids_row in zip(logits, input_ids[:, -logits_to_keep:]): - log_probs = logits_row.log_softmax(dim=-1) - token_log_prob = torch.gather(log_probs, dim=1, index=input_ids_row.unsqueeze(1)).squeeze(1) - per_token_logps.append(token_log_prob) - return torch.stack(per_token_logps) - - logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens - per_token_logps = get_per_token_logps(model, prompt_completion_ids, logits_to_keep) - - with torch.inference_mode(): - if self.ref_model is not None: - ref_per_token_logps = get_per_token_logps(self.ref_model, prompt_completion_ids, logits_to_keep) - else: - with self.accelerator.unwrap_model(model).disable_adapter(): - ref_per_token_logps = get_per_token_logps(model, prompt_completion_ids, logits_to_keep) + prompts = [prompt for prompt in prompts for _ in range(self.num_generations)] - # Compute the KL divergence between the model and the reference model - per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 + # iterate through bsz in loss_bsz chunks and accumulate: + # - rewards + # - per-token log probabilities + # - per-token KL divergences + all_rewards_per_func = [] + all_per_token_logps = [] + all_per_token_kl = [] + for i in range(0, bsz, loss_bsz): + current_batch_idcs = slice(i, i + loss_bsz) + + completion_ids = prompt_completion_ids[current_batch_idcs, prompt_length:] + current_bsz = completion_ids.size(0) + + # Get the per-token log probabilities for the completions for the model and the reference model + def get_per_token_logps(model, input_ids, logits_to_keep): + # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded + logits = model(input_ids, logits_to_keep=logits_to_keep + 1).logits # (B, L, V) + logits = logits[ + :, :-1, : + ] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + + # Compute the log probabilities for the input tokens. Use a loop to reduce memory peak. + per_token_logps = [] + for logits_row, input_ids_row in zip(logits, input_ids[:, -logits_to_keep:]): + log_probs = logits_row.log_softmax(dim=-1) + token_log_prob = torch.gather(log_probs, dim=1, index=input_ids_row.unsqueeze(1)).squeeze(1) + per_token_logps.append(token_log_prob) + return torch.stack(per_token_logps) + + logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens + per_token_logps = get_per_token_logps(model, prompt_completion_ids, logits_to_keep) + + with torch.inference_mode(): + if self.ref_model is not None: + ref_per_token_logps = get_per_token_logps(self.ref_model, prompt_completion_ids, logits_to_keep) + else: + with self.accelerator.unwrap_model(model).disable_adapter(): + ref_per_token_logps = get_per_token_logps(model, prompt_completion_ids, logits_to_keep) - # Mask everything after the first EOS token - is_eos = completion_ids == self.processing_class.eos_token_id - eos_idx = torch.full((is_eos.size(0),), is_eos.size(1), dtype=torch.long, device=device) - eos_idx[is_eos.any(dim=1)] = is_eos.int().argmax(dim=1)[is_eos.any(dim=1)] - sequence_indices = torch.arange(is_eos.size(1), device=device).expand(is_eos.size(0), -1) - completion_mask = (sequence_indices <= eos_idx.unsqueeze(1)).int() + # Compute the KL divergence between the model and the reference model + per_token_kl = ( + torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 + ) - # Decode the generated completions - completions = self.processing_class.batch_decode(completion_ids, skip_special_tokens=True) - if is_conversational(inputs[0]): - completions = [[{"role": "assistant", "content": completion}] for completion in completions] + # Mask everything after the first EOS token + is_eos = completion_ids == self.processing_class.eos_token_id + eos_idx = torch.full((is_eos.size(0),), is_eos.size(1), dtype=torch.long, device=device) + eos_idx[is_eos.any(dim=1)] = is_eos.int().argmax(dim=1)[is_eos.any(dim=1)] + sequence_indices = torch.arange(is_eos.size(1), device=device).expand(is_eos.size(0), -1) + completion_mask = (sequence_indices <= eos_idx.unsqueeze(1)).int() + + # Decode the generated completions + completions = self.processing_class.batch_decode(completion_ids, skip_special_tokens=True) + if is_conversational(inputs[0]): + completions = [[{"role": "assistant", "content": completion}] for completion in completions] + + # Compute the rewards + + rewards_per_func = torch.zeros(current_bsz, len(self.reward_funcs), device=device) + for i, (reward_func, reward_processing_class) in enumerate( + zip(self.reward_funcs, self.reward_processing_classes) + ): + if isinstance(reward_func, PreTrainedModel): + if is_conversational(inputs[0]): + messages = [{"messages": p + c} for p, c in zip(prompts[current_batch_idcs], completions)] + texts = [apply_chat_template(x, reward_processing_class)["text"] for x in messages] + else: + texts = [p + c for p, c in zip(prompts[current_batch_idcs], completions)] + reward_inputs = reward_processing_class( + texts, return_tensors="pt", padding=True, padding_side="right", add_special_tokens=False + ) + reward_inputs = super()._prepare_inputs(reward_inputs) + with torch.inference_mode(): + rewards_per_func[:, i] = reward_func(**reward_inputs).logits[:, 0] # Shape (B*G,) + else: + # Repeat all input columns (but "prompt" and "completion") to match the number of generations + reward_kwargs = {key: [] for key in inputs[0].keys() if key not in ["prompt", "completion"]} + for key in reward_kwargs: + for example in inputs: + # Repeat each value in the column for `num_generations` times + reward_kwargs[key].extend([example[key]] * self.num_generations) + output_reward_func = reward_func( + prompts=prompts[current_batch_idcs], completions=completions, **reward_kwargs + ) + rewards_per_func[:, i] = torch.tensor(output_reward_func, dtype=torch.float32, device=device) - # Compute the rewards - prompts = [prompt for prompt in prompts for _ in range(self.num_generations)] + all_rewards_per_func.append(rewards_per_func) + all_per_token_logps.append(per_token_logps) + all_per_token_kl.append(per_token_kl) - rewards_per_func = torch.zeros(len(prompts), len(self.reward_funcs), device=device) - for i, (reward_func, reward_processing_class) in enumerate( - zip(self.reward_funcs, self.reward_processing_classes) - ): - if isinstance(reward_func, PreTrainedModel): - if is_conversational(inputs[0]): - messages = [{"messages": p + c} for p, c in zip(prompts, completions)] - texts = [apply_chat_template(x, reward_processing_class)["text"] for x in messages] - else: - texts = [p + c for p, c in zip(prompts, completions)] - reward_inputs = reward_processing_class( - texts, return_tensors="pt", padding=True, padding_side="right", add_special_tokens=False - ) - reward_inputs = super()._prepare_inputs(reward_inputs) - with torch.inference_mode(): - rewards_per_func[:, i] = reward_func(**reward_inputs).logits[:, 0] # Shape (B*G,) - else: - # Repeat all input columns (but "prompt" and "completion") to match the number of generations - reward_kwargs = {key: [] for key in inputs[0].keys() if key not in ["prompt", "completion"]} - for key in reward_kwargs: - for example in inputs: - # Repeat each value in the column for `num_generations` times - reward_kwargs[key].extend([example[key]] * self.num_generations) - output_reward_func = reward_func(prompts=prompts, completions=completions, **reward_kwargs) - rewards_per_func[:, i] = torch.tensor(output_reward_func, dtype=torch.float32, device=device) + # Concatenate and compute the loss + rewards_per_func = torch.concatenate(all_rewards_per_func, dim=0) + per_token_logps = torch.concatenate(all_per_token_logps, dim=0) + per_token_kl = torch.concatenate(all_per_token_kl, dim=0) # Sum the rewards from all reward functions rewards = rewards_per_func.sum(dim=1) From fd751f3f9b80bce954b689305ed13ddd981bcd8c Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Sat, 1 Feb 2025 15:54:18 -0800 Subject: [PATCH 5/8] fix issue with reward func kwargs --- trl/trainer/grpo_trainer.py | 26 ++++++++++++++------------ 1 file changed, 14 insertions(+), 12 deletions(-) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index 7aed43b9815..67200ad92ad 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -435,6 +435,14 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N prompt_length = prompt_inputs["input_ids"].size(1) prompts = [prompt for prompt in prompts for _ in range(self.num_generations)] + # Prepare reward kwargs before looping through microbatches + if any(not isinstance(reward_func, PreTrainedModel) for reward_func in self.reward_funcs): + # Repeat all input columns (but "prompt" and "completion") to match the number of generations + all_reward_kwargs = {key: [] for key in inputs[0].keys() if key not in ["prompt", "completion"]} + for key in all_reward_kwargs: + for example in inputs: + all_reward_kwargs[key].extend([example[key]] * self.num_generations) + # iterate through bsz in loss_bsz chunks and accumulate: # - rewards # - per-token log probabilities @@ -443,9 +451,9 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N all_per_token_logps = [] all_per_token_kl = [] for i in range(0, bsz, loss_bsz): - current_batch_idcs = slice(i, i + loss_bsz) + current_batch_span = slice(i, i + loss_bsz) - completion_ids = prompt_completion_ids[current_batch_idcs, prompt_length:] + completion_ids = prompt_completion_ids[current_batch_span, prompt_length:] current_bsz = completion_ids.size(0) # Get the per-token log probabilities for the completions for the model and the reference model @@ -492,17 +500,16 @@ def get_per_token_logps(model, input_ids, logits_to_keep): completions = [[{"role": "assistant", "content": completion}] for completion in completions] # Compute the rewards - rewards_per_func = torch.zeros(current_bsz, len(self.reward_funcs), device=device) for i, (reward_func, reward_processing_class) in enumerate( zip(self.reward_funcs, self.reward_processing_classes) ): if isinstance(reward_func, PreTrainedModel): if is_conversational(inputs[0]): - messages = [{"messages": p + c} for p, c in zip(prompts[current_batch_idcs], completions)] + messages = [{"messages": p + c} for p, c in zip(prompts[current_batch_span], completions)] texts = [apply_chat_template(x, reward_processing_class)["text"] for x in messages] else: - texts = [p + c for p, c in zip(prompts[current_batch_idcs], completions)] + texts = [p + c for p, c in zip(prompts[current_batch_span], completions)] reward_inputs = reward_processing_class( texts, return_tensors="pt", padding=True, padding_side="right", add_special_tokens=False ) @@ -510,14 +517,9 @@ def get_per_token_logps(model, input_ids, logits_to_keep): with torch.inference_mode(): rewards_per_func[:, i] = reward_func(**reward_inputs).logits[:, 0] # Shape (B*G,) else: - # Repeat all input columns (but "prompt" and "completion") to match the number of generations - reward_kwargs = {key: [] for key in inputs[0].keys() if key not in ["prompt", "completion"]} - for key in reward_kwargs: - for example in inputs: - # Repeat each value in the column for `num_generations` times - reward_kwargs[key].extend([example[key]] * self.num_generations) + reward_kwargs = {k: v[current_batch_span] for k, v in all_reward_kwargs.items()} output_reward_func = reward_func( - prompts=prompts[current_batch_idcs], completions=completions, **reward_kwargs + prompts=prompts[current_batch_span], completions=completions, **reward_kwargs ) rewards_per_func[:, i] = torch.tensor(output_reward_func, dtype=torch.float32, device=device) From ed038fec01502274110ec07db29e7481fc9254bb Mon Sep 17 00:00:00 2001 From: Siddharth Narayanan Date: Sat, 1 Feb 2025 16:08:22 -0800 Subject: [PATCH 6/8] rename arg to micro bsz --- trl/trainer/grpo_config.py | 2 +- trl/trainer/grpo_trainer.py | 7 ++++--- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/trl/trainer/grpo_config.py b/trl/trainer/grpo_config.py index 814b8b3b696..5d41071f320 100644 --- a/trl/trainer/grpo_config.py +++ b/trl/trainer/grpo_config.py @@ -175,7 +175,7 @@ class GRPOConfig(TrainingArguments): metadata={"help": "KL coefficient."}, ) - per_device_loss_batch_size: int = field( + per_device_micro_batch_size: int = field( default=8, metadata={ "help": "Micro batch size per GPU/TPU/MPS/NPU core/CPU for computing loss terms. " diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index 67200ad92ad..cee65f127d4 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -148,6 +148,7 @@ class GRPOTrainer(Trainer): """ _tag_names = ["trl", "grpo"] + args: GRPOConfig # helps with type hinting def __init__( self, @@ -431,7 +432,7 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N logger.info("Finished generation") bsz = prompt_inputs["input_ids"].size(0) - loss_bsz = self.args.per_device_loss_batch_size or bsz + micro_bsz = self.args.per_device_micro_batch_size or bsz # default to full batch if not set prompt_length = prompt_inputs["input_ids"].size(1) prompts = [prompt for prompt in prompts for _ in range(self.num_generations)] @@ -450,8 +451,8 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N all_rewards_per_func = [] all_per_token_logps = [] all_per_token_kl = [] - for i in range(0, bsz, loss_bsz): - current_batch_span = slice(i, i + loss_bsz) + for i in range(0, bsz, micro_bsz): + current_batch_span = slice(i, i + micro_bsz) completion_ids = prompt_completion_ids[current_batch_span, prompt_length:] current_bsz = completion_ids.size(0) From c763f072c892593f3542cd4cfc758ce09cb2f987 Mon Sep 17 00:00:00 2001 From: Siddharth M Narayanan Date: Sat, 1 Feb 2025 19:45:35 -0600 Subject: [PATCH 7/8] fix microbatching --- trl/trainer/grpo_trainer.py | 40 ++++++++++++++++++++++--------------- 1 file changed, 24 insertions(+), 16 deletions(-) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index cee65f127d4..f3a2fba4213 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -403,24 +403,24 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N all_prompts_text = gather_object(prompts_text) if self.accelerator.is_main_process: outputs = self.llm.generate(all_prompts_text, sampling_params=self.sampling_params, use_tqdm=False) - completion_ids = [out.token_ids for completions in outputs for out in completions.outputs] + micro_completion_ids = [out.token_ids for completions in outputs for out in completions.outputs] else: - completion_ids = [None] * len(all_prompts_text) * self.num_generations + micro_completion_ids = [None] * len(all_prompts_text) * self.num_generations # Broadcast the completions from the main process to all processes, ensuring each process receives its # corresponding slice. - completion_ids = broadcast_object_list(completion_ids, from_process=0) + micro_completion_ids = broadcast_object_list(micro_completion_ids, from_process=0) process_slice = slice( self.accelerator.process_index * len(prompts) * self.num_generations, (self.accelerator.process_index + 1) * len(prompts) * self.num_generations, ) - completion_ids = completion_ids[process_slice] + micro_completion_ids = micro_completion_ids[process_slice] # Pad the completions, and concatenate them with the prompts - completion_ids = [torch.tensor(ids, device=device) for ids in completion_ids] - completion_ids = pad(completion_ids, padding_value=self.processing_class.pad_token_id) + micro_completion_ids = [torch.tensor(ids, device=device) for ids in micro_completion_ids] + micro_completion_ids = pad(micro_completion_ids, padding_value=self.processing_class.pad_token_id) prompt_inputs_repeated = torch.repeat_interleave(prompt_inputs["input_ids"], self.num_generations, dim=0) - prompt_completion_ids = torch.cat([prompt_inputs_repeated, completion_ids], dim=1) + prompt_completion_ids = torch.cat([prompt_inputs_repeated, micro_completion_ids], dim=1) else: # Regular generation path with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model: @@ -431,7 +431,7 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N model.train() logger.info("Finished generation") - bsz = prompt_inputs["input_ids"].size(0) + bsz = prompt_completion_ids.size(0) micro_bsz = self.args.per_device_micro_batch_size or bsz # default to full batch if not set prompt_length = prompt_inputs["input_ids"].size(1) prompts = [prompt for prompt in prompts for _ in range(self.num_generations)] @@ -448,14 +448,17 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N # - rewards # - per-token log probabilities # - per-token KL divergences + # - masks of completion tokens all_rewards_per_func = [] all_per_token_logps = [] all_per_token_kl = [] + all_completion_mask = [] for i in range(0, bsz, micro_bsz): current_batch_span = slice(i, i + micro_bsz) - completion_ids = prompt_completion_ids[current_batch_span, prompt_length:] - current_bsz = completion_ids.size(0) + micro_prompt_completion_ids = prompt_completion_ids[current_batch_span] + micro_completion_ids = micro_prompt_completion_ids[:, prompt_length:] + current_bsz = micro_completion_ids.size(0) # last one may be Date: Sun, 2 Feb 2025 00:05:21 -0600 Subject: [PATCH 8/8] rename to completion_ids in vllm block --- trl/trainer/grpo_trainer.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index f3a2fba4213..5e8064fb21d 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -403,24 +403,24 @@ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=N all_prompts_text = gather_object(prompts_text) if self.accelerator.is_main_process: outputs = self.llm.generate(all_prompts_text, sampling_params=self.sampling_params, use_tqdm=False) - micro_completion_ids = [out.token_ids for completions in outputs for out in completions.outputs] + completion_ids = [out.token_ids for completions in outputs for out in completions.outputs] else: - micro_completion_ids = [None] * len(all_prompts_text) * self.num_generations + completion_ids = [None] * len(all_prompts_text) * self.num_generations # Broadcast the completions from the main process to all processes, ensuring each process receives its # corresponding slice. - micro_completion_ids = broadcast_object_list(micro_completion_ids, from_process=0) + completion_ids = broadcast_object_list(completion_ids, from_process=0) process_slice = slice( self.accelerator.process_index * len(prompts) * self.num_generations, (self.accelerator.process_index + 1) * len(prompts) * self.num_generations, ) - micro_completion_ids = micro_completion_ids[process_slice] + completion_ids = completion_ids[process_slice] # Pad the completions, and concatenate them with the prompts - micro_completion_ids = [torch.tensor(ids, device=device) for ids in micro_completion_ids] - micro_completion_ids = pad(micro_completion_ids, padding_value=self.processing_class.pad_token_id) + completion_ids = [torch.tensor(ids, device=device) for ids in completion_ids] + completion_ids = pad(completion_ids, padding_value=self.processing_class.pad_token_id) prompt_inputs_repeated = torch.repeat_interleave(prompt_inputs["input_ids"], self.num_generations, dim=0) - prompt_completion_ids = torch.cat([prompt_inputs_repeated, micro_completion_ids], dim=1) + prompt_completion_ids = torch.cat([prompt_inputs_repeated, completion_ids], dim=1) else: # Regular generation path with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model: