diff --git a/slime/backends/megatron_utils/actor.py b/slime/backends/megatron_utils/actor.py index b39d3b57d4..81805cf493 100644 --- a/slime/backends/megatron_utils/actor.py +++ b/slime/backends/megatron_utils/actor.py @@ -4,7 +4,7 @@ from argparse import Namespace from contextlib import nullcontext from pathlib import Path -from typing import Dict, Optional, Tuple, Union +from typing import Any, Dict, Optional, Tuple, Union import ray import torch @@ -217,6 +217,23 @@ def _get_rollout_data( rollout_data["rollout_log_probs"], rollout_data["total_lengths"], rollout_data["response_lengths"] ) ] + + # Move RL fields to GPU if they exist and are not already on GPU + # This is needed for external APIs (like GMI wrapper) that provide pre-computed RL fields + # In normal Slime training, these fields come from GPU forward passes and don't need moving + if getattr(self.args, 'move_rl_fields_to_gpu', False): + for field in ["log_probs", "ref_log_probs", "advantages", "returns", "values"]: + if field in rollout_data and rollout_data[field]: + # Check if first tensor is already on GPU to avoid unnecessary transfers + first_tensor = rollout_data[field][0] + if isinstance(first_tensor, torch.Tensor) and not first_tensor.is_cuda: + rollout_data[field] = [ + torch.tensor(t, dtype=torch.float32, device=torch.cuda.current_device()) + if not isinstance(t, torch.Tensor) or not t.is_cuda + else t.to(device=torch.cuda.current_device()) + for t in rollout_data[field] + ] + return rollout_data def compute_log_prob( @@ -390,6 +407,259 @@ def train_actor( log_perf_data(rollout_id, self.args) Timer().start("train_wait") + def forward_backward_step_only( + self, rollout_id: int, rollout_data_ref: Box, zero_grads: bool = False + ) -> Dict[str, float]: + """ + Perform forward + backward pass only, accumulating gradients WITHOUT optimizer.step(). + + This enables gradient accumulation by calling Megatron's forward_backward_func directly + without the optimizer step. + + Args: + rollout_id: Rollout identifier + rollout_data_ref: Reference to rollout data + zero_grads: If True, zero gradients before forward pass (first accumulation step). + If False, accumulate on top of existing gradients (subsequent steps). + + Returns: + Dictionary with loss, grad_norm, and valid_step information + """ + import math + import os + from functools import partial + + from megatron.core import mpu + from megatron.core.models.gpt import GPTModel + from megatron.core.pipeline_parallel import get_forward_backward_func + from megatron.training.global_vars import get_args + + from .data import get_batch + from .loss import loss_function + + Timer().end("train_wait") + + if self.args.offload: + self.wake_up(("model")) + + with timer("data_preprocess"): + rollout_data = self._get_rollout_data(rollout_data_ref) + + # Create data iterator + data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) + + with timer("forward_backward_only"): + args = get_args() + + # Optionally zero gradients (only for first accumulation step) + if zero_grads: + for model_chunk in self.model: + model_chunk.zero_grad_buffer() + self.optimizer.zero_grad() + + if args.custom_megatron_before_train_step_hook_path: + from slime.utils.misc import load_function + + custom_before_train_step_hook = load_function(args.custom_megatron_before_train_step_hook_path) + custom_before_train_step_hook(args, rollout_id, 0, self.model, self.optimizer, self.opt_param_scheduler) + + def forward_step(data_iterator, model: GPTModel): + """Forward training step.""" + batch = get_batch( + data_iterator, + [ + "tokens", + "packed_seq_params", + "total_lengths", + "response_lengths", + "loss_masks", + "log_probs", + "ref_log_probs", + "values", + "advantages", + "returns", + "rollout_log_probs", + ], + ) + + if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1": + old_stage = os.environ["ROUTING_REPLAY_STAGE"] + os.environ["ROUTING_REPLAY_STAGE"] = "replay_forward" + + output_tensor = model( + input_ids=batch["tokens"], + position_ids=None, + attention_mask=None, + labels=None, + packed_seq_params=batch["packed_seq_params"], + ) + + if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1": + os.environ["ROUTING_REPLAY_STAGE"] = old_stage + + return output_tensor, partial(loss_function, args, batch, num_microbatches[0]) + + # Call Megatron's forward_backward_func directly + forward_backward_func = get_forward_backward_func() + losses_reduced = forward_backward_func( + forward_step_func=forward_step, + data_iterator=data_iterator, + model=self.model, + num_microbatches=num_microbatches[0], + seq_length=args.seq_length, + micro_batch_size=args.micro_batch_size, + decoder_seq_length=args.decoder_seq_length, + forward_only=False, + ) + + # Validate gradients + valid_step = True + grad_norm = None + if not getattr(args, "check_for_nan_in_loss_and_grad", True): + found_inf_flag = self.optimizer.prepare_grads() + if found_inf_flag: + valid_step = False + else: + grad_norm = self.optimizer.get_grad_norm() + if isinstance(grad_norm, torch.Tensor): + valid_step = not (torch.isnan(grad_norm) or torch.isinf(grad_norm)) + else: + valid_step = not (math.isnan(grad_norm) or math.isinf(grad_norm)) + + # Compute loss metrics + loss_dict = {} + if mpu.is_pipeline_last_stage(ignore_virtual=True): + keys = losses_reduced[0]["keys"] + values = None + for x in losses_reduced: + if values is None: + values = x["values"] + else: + values += x["values"] + assert len(keys) + 1 == values.numel() + torch.distributed.all_reduce(values, group=mpu.get_data_parallel_group(with_context_parallel=True)) + + values = values.tolist() + num_samples_or_tokens = values[0] + for key, value in zip(keys, values[1:]): + loss_dict[key] = value * mpu.get_context_parallel_world_size() / num_samples_or_tokens + + # Extract log_probs if present (added separately, not in keys/values tensor) + # Aggregate log_probs from all microbatches + if "log_probs" in losses_reduced[0]: + # log_probs is a list of tensors (one per sample in batch) + # With multiple microbatches, we need to concatenate across microbatches + all_log_probs = [] + for x in losses_reduced: + if "log_probs" in x and x["log_probs"]: + all_log_probs.extend(x["log_probs"]) + loss_dict["log_probs"] = all_log_probs + + Timer().start("train_wait") + + return { + "loss": loss_dict, + "grad_norm": grad_norm if grad_norm is not None else 0.0, + "valid_step": valid_step, + } + + def forward_only_step( + self, rollout_id: int, rollout_data_ref: Box + ) -> Dict[str, Any]: + """ + Perform forward-only pass WITHOUT gradients, returning logprobs per sample. + + This is used for DPO's forward_backward_custom where we need reference logprobs + from a forward pass, then apply custom loss function client-side. + + Unlike forward_backward_step_only, this does NOT compute gradients. + + Args: + rollout_id: Rollout identifier + rollout_data_ref: Reference to rollout data + + Returns: + Dictionary with loss_dict containing log_probs per sample + """ + from megatron.core import mpu + + from .loss import get_log_probs_and_entropy + from .model import forward_only + + Timer().end("train_wait") + + if self.args.offload: + self.wake_up(("model")) + + with timer("data_preprocess"): + rollout_data = self._get_rollout_data(rollout_data_ref) + + # Create data iterator + data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data) + + with timer("forward_only"): + # Call forward_only which sets model to eval mode and does forward pass only + rollout_data_result = forward_only( + get_log_probs_and_entropy, + self.args, + self.model, + data_iterator, + num_microbatches, + store_prefix="", + ) + + # Extract log_probs from rollout_data_result + # forward_only returns dict with "log_probs" key containing list of tensors + loss_dict = {} + if mpu.is_pipeline_last_stage(): + if "log_probs" in rollout_data_result: + loss_dict["log_probs"] = rollout_data_result["log_probs"] + if "entropy" in rollout_data_result: + loss_dict["entropy"] = rollout_data_result["entropy"] + + Timer().start("train_wait") + + return { + "loss": loss_dict, + "grad_norm": 0.0, # No gradients computed in forward-only pass + "valid_step": True, # Always valid since no gradient computation + } + + def apply_optimizer_step(self) -> Dict[str, float]: + """ + Apply optimizer step using accumulated gradients. + + This enables gradient accumulation by applying the optimizer step + and zeroing gradients after. + + Returns: + Dictionary with success status and grad_norm + """ + from megatron.training.global_vars import get_args + + with timer("apply_optimizer_step"): + args = get_args() + + # Apply optimizer step + update_successful, grad_norm, num_zeros_in_grad = self.optimizer.step() + + # Update learning rate scheduler + if update_successful: + self.opt_param_scheduler.step(increment=args.global_batch_size) + + # Zero gradients after applying them + for model_chunk in self.model: + model_chunk.zero_grad_buffer() + self.optimizer.zero_grad() + + # Update CPU weight cache after applying optimizer step + self.update_cpu_params_dict(self.weights["actor"]) + + return { + "success": update_successful, + "grad_norm": grad_norm if grad_norm is not None else 0.0, + } + def save_model(self, iteration: int) -> None: if self.args.debug_rollout_only: return diff --git a/slime/backends/megatron_utils/data.py b/slime/backends/megatron_utils/data.py index 28742de262..8d0149a7ba 100644 --- a/slime/backends/megatron_utils/data.py +++ b/slime/backends/megatron_utils/data.py @@ -179,8 +179,23 @@ def get_data_iterator(args, model, rollout_data): cp_size = mpu.get_context_parallel_world_size() num_local_samples = len(rollout_data["total_lengths"]) - num_local_gbs = args.global_batch_size // dp_size - num_steps_per_rollout = num_local_samples // num_local_gbs + + # FLEXIBLE BATCH SIZE: Support variable batch sizes from tinker-cookbook + # while preserving gradient accumulation semantics + target_local_batch_size = args.global_batch_size // dp_size + + # Handle variable batch sizes gracefully + if num_local_samples <= target_local_batch_size: + # Small batch - process in one step (no gradient accumulation needed) + num_local_gbs = num_local_samples + num_steps_per_rollout = 1 + else: + # Large batch - use gradient accumulation across multiple steps + num_local_gbs = target_local_batch_size + num_steps_per_rollout = (num_local_samples + num_local_gbs - 1) // num_local_gbs # Ceiling division + + # Pass actual batch size for loss scaling (used in loss.py) + rollout_data["_actual_global_batch_size"] = num_local_samples * dp_size def _generate_data_iterator(rollout_data, micro_batch_size, micro_batch_indices=None): data_iterator = [] diff --git a/slime/backends/megatron_utils/loss.py b/slime/backends/megatron_utils/loss.py index 2290fda8dd..6d1301b8a2 100644 --- a/slime/backends/megatron_utils/loss.py +++ b/slime/backends/megatron_utils/loss.py @@ -262,7 +262,6 @@ def compute_advantages_and_returns(args, rollout_data): def policy_loss_function(args, batch, logits, sum_of_sample_mean): - advantages = torch.cat(batch["advantages"], dim=0) old_log_probs = batch["log_probs"] response_lengths = batch["response_lengths"] @@ -279,6 +278,63 @@ def policy_loss_function(args, batch, logits, sum_of_sample_mean): log_probs = log_probs_and_entropy["log_probs"] + # Check if we're on a non-last PP stage (no logits computed) + # Similar to check in compute_advantages_and_returns() at line 137 + if not log_probs or len(log_probs) == 0: + # Return dummy loss - this rank doesn't compute the loss + return torch.tensor(0.0, device=logits.device if logits is not None else "cpu"), {} + + # Handle dynamic batching: batch may contain more samples than the forward pass computed + # First, subset batch data to match the number of computed samples + num_samples_computed = len(log_probs) + old_log_probs = old_log_probs[:num_samples_computed] + response_lengths = response_lengths[:num_samples_computed] + total_lengths = total_lengths[:num_samples_computed] + + # Filter out empty tensors from log_probs and match corresponding batch indices + # Empty tensors occur when get_log_probs_and_entropy() returns [0] for some samples + valid_indices = [i for i, lp in enumerate(log_probs) if lp.shape[0] > 0] + + # Apply filter to all parallel lists + log_probs = [log_probs[i] for i in valid_indices] + old_log_probs = [old_log_probs[i] for i in valid_indices] + response_lengths = [response_lengths[i] for i in valid_indices] + total_lengths = [total_lengths[i] for i in valid_indices] + + # Advantages need special handling since they're already concatenated per-sample + advantages = torch.cat([batch["advantages"][i] for i in valid_indices], dim=0) + + # Keep per-sample logprobs for Tinker API (detach to avoid keeping computation graph) + # IMPORTANT: Create AFTER filtering to ensure we only return valid (non-empty) log_probs + per_sample_log_probs = [lp.clone().detach() for lp in log_probs] + + # Filter loss_masks to match valid indices + loss_masks_filtered = [batch["loss_masks"][i] for i in valid_indices] + + # Recreate sum_of_sample_mean with filtered data + # The original sum_of_sample_mean was created with unfiltered data, so we need to recreate it + from .cp_utils import get_sum_of_sample_mean + sum_of_sample_mean = get_sum_of_sample_mean( + total_lengths, + response_lengths, + loss_masks_filtered, + args.calculate_per_token_loss, + ) + + # Debug logging to file (Ray actors don't show print output) + import os + import torch.distributed as dist + rank = dist.get_rank() if dist.is_initialized() else 0 + with open(f"/tmp/grpo_debug_rank{rank}.log", "a") as f: + f.write(f"\n=== GRPO DEBUG RANK {rank} (AFTER FILTER) ===\n") + f.write(f"num_samples_computed={num_samples_computed}\n") + f.write(f"valid_indices={valid_indices}\n") + f.write(f"old_log_probs shapes: {[t.shape[0] for t in old_log_probs]}\n") + f.write(f"new log_probs shapes: {[t.shape[0] for t in log_probs]}\n") + f.write(f"response_lengths: {response_lengths}\n") + f.write(f"old_log_probs total tokens: {sum(t.shape[0] for t in old_log_probs)}\n") + f.write(f"new log_probs total tokens: {sum(t.shape[0] for t in log_probs)}\n") + if args.advantage_estimator == "gspo": full_log_probs = [ all_gather_with_cp(log_prob, total_length, response_length) @@ -298,7 +354,8 @@ def policy_loss_function(args, batch, logits, sum_of_sample_mean): ppo_kl = torch.cat(ppo_kl, dim=0) log_probs = torch.cat(log_probs, dim=0) else: - old_log_probs = torch.cat(batch["log_probs"], dim=0) + # GRPO path (old_log_probs already subset above if needed) + old_log_probs = torch.cat(old_log_probs, dim=0) log_probs = torch.cat(log_probs, dim=0) ppo_kl = old_log_probs - log_probs @@ -350,6 +407,7 @@ def policy_loss_function(args, batch, logits, sum_of_sample_mean): "entropy_loss": entropy_loss.clone().detach(), "pg_clipfrac": pg_clipfrac.clone().detach(), "ppo_kl": ppo_kl.clone().detach(), + "log_probs": per_sample_log_probs, # Per-sample logprobs for Tinker API } if args.use_kl_loss: @@ -393,6 +451,7 @@ def value_loss_function(args, batch, logits, sum_of_sample_mean): reported_loss = { "value_loss": loss.clone().detach(), "value_clipfrac": values_clipfrac.clone().detach(), + "log_probs": [], # Value loss doesn't produce per-token logprobs } return loss, reported_loss @@ -411,7 +470,10 @@ def sft_loss_function(args, batch, logits, sum_of_sample_mean): with_entropy=False, ) - log_probs = log_probs_and_entropy["log_probs"] + log_probs = log_probs_and_entropy["log_probs"] # List of per-sample tensors + # Keep per-sample logprobs for Tinker API (detach to avoid keeping computation graph) + per_sample_log_probs = [lp.clone().detach() for lp in log_probs] + log_probs = torch.cat(log_probs, dim=0) loss = -sum_of_sample_mean(log_probs) @@ -423,6 +485,7 @@ def sft_loss_function(args, batch, logits, sum_of_sample_mean): loss, { "loss": loss.clone().detach(), + "log_probs": per_sample_log_probs, # Per-sample logprobs for Tinker API }, ) @@ -459,21 +522,33 @@ def loss_function(args, batch, num_microbatches, logits): raise ValueError(f"Unknown loss type: {args.loss_type}") # Here we need to divide by cp_size because to cancel the multiply in Megatron. + # Use actual batch size if provided (for variable batch sizes from tinker-cookbook) + actual_global_batch_size = batch.get("_actual_global_batch_size", args.global_batch_size) loss = ( - loss * num_microbatches / args.global_batch_size * mpu.get_data_parallel_world_size(with_context_parallel=True) + loss * num_microbatches / actual_global_batch_size * mpu.get_data_parallel_world_size(with_context_parallel=True) ) + # Separate log_probs (list of tensors) from scalar metrics + # log_probs cannot be converted to a single tensor, so we handle it separately + log_probs = log.pop("log_probs", None) # Remove log_probs if present + + result_dict = { + "keys": list(log.keys()), + "values": torch.tensor( + [ + num_samples if not args.calculate_per_token_loss else num_tokens, + ] + + list(log.values()), + device=logits.device, + ), + } + + # Add log_probs back to the result dict (not in the tensor) + if log_probs is not None: + result_dict["log_probs"] = log_probs + return ( loss, num_tokens if args.calculate_per_token_loss else 1, - { - "keys": list(log.keys()), - "values": torch.tensor( - [ - num_samples if not args.calculate_per_token_loss else num_tokens, - ] - + list(log.values()), - device=logits.device, - ), - }, + result_dict, ) diff --git a/slime/ray/actor_group.py b/slime/ray/actor_group.py index d056f0e24f..f268ce4494 100644 --- a/slime/ray/actor_group.py +++ b/slime/ray/actor_group.py @@ -138,3 +138,60 @@ def connect(self, critic_group): def set_rollout_manager(self, rollout_manager): return ray.get([actor.set_rollout_manager.remote(rollout_manager) for actor in self._actor_handlers]) + + def forward_backward_only(self, rollout_id, rollout_data_ref, zero_grads=False): + """ + Perform forward + backward pass only, accumulating gradients WITHOUT optimizer.step(). + + This enables gradient accumulation for Tinker API integration. + + Args: + rollout_id: Rollout identifier + rollout_data_ref: Reference to rollout data + zero_grads: If True, zero gradients before forward pass (first accumulation step). + If False, accumulate on top of existing gradients (subsequent steps). + + Returns: + List of results from all actors (loss, grad_norm, valid_step) + """ + return ray.get( + [ + actor.forward_backward_step_only.remote(rollout_id, rollout_data_ref, zero_grads) + for actor in self._actor_handlers + ] + ) + + def forward_only(self, rollout_id, rollout_data_ref): + """ + Perform forward-only pass WITHOUT gradients, returning logprobs per sample. + + This is used for DPO's forward_backward_custom where we need reference logprobs + from a forward pass, then apply custom loss function client-side. + + Unlike forward_backward_only, this does NOT compute gradients. + + Args: + rollout_id: Rollout identifier + rollout_data_ref: Reference to rollout data + + Returns: + List of results from all actors (loss dict with log_probs) + """ + return ray.get( + [ + actor.forward_only_step.remote(rollout_id, rollout_data_ref) + for actor in self._actor_handlers + ] + ) + + def apply_optimizer_step(self): + """ + Apply optimizer step using accumulated gradients. + + This enables gradient accumulation for Tinker API integration. + Must be called after one or more forward_backward_only() calls. + + Returns: + List of results from all actors (success, grad_norm) + """ + return ray.get([actor.apply_optimizer_step.remote() for actor in self._actor_handlers]) diff --git a/slime/ray/rollout.py b/slime/ray/rollout.py index c66546287d..372032f74a 100644 --- a/slime/ray/rollout.py +++ b/slime/ray/rollout.py @@ -77,6 +77,10 @@ def rollout_engines(self): def get_rollout_engines_and_lock(self): return self.rollout_engines, self.rollout_engine_lock, self.num_new_engines + def get_router_address(self): + """Get SGLang router address (for GMI wrapper integration)""" + return self.args.sglang_router_ip, self.args.sglang_router_port + def get_num_rollout_per_epoch(self): assert self.args.rollout_global_dataset return len(self.data_source.dataset) // self.args.rollout_batch_size diff --git a/slime/utils/data.py b/slime/utils/data.py index 5140782f5d..c109a22fcc 100644 --- a/slime/utils/data.py +++ b/slime/utils/data.py @@ -131,8 +131,9 @@ def process_rollout_data(args, rollout_data_ref, dp_rank, dp_size): dist.broadcast_object_list(data, src=0) data = data[0] - # save the unprocessed reward for logging - rollout_data["raw_reward"] = data["raw_reward"] + # save the unprocessed reward for logging (optional for forward-only passes) + if "raw_reward" in data: + rollout_data["raw_reward"] = data["raw_reward"] if "prompt" in data: rollout_data["prompt"] = data["prompt"] @@ -186,6 +187,12 @@ def get_partition(val): "sample_indices", "rollout_log_probs", "prompt", + # Additional RL training fields for forward_backward_only API + "advantages", + "returns", + "log_probs", + "ref_log_probs", + "values", ]: if key not in data: continue