Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
import os

import slime.utils.misc as U
from slime.utils.external_utils.command_utils import execute_train_npu

MODEL_NAME = os.environ.get("SLIME_SCRIPT_MODEL_NAME", "Qwen3-VL-2B-Instruct")
Expand Down Expand Up @@ -78,7 +77,6 @@ def execute():
"--weight-decay 0.1 "
"--adam-beta1 0.9 "
"--adam-beta2 0.98 "

"--optimizer-cpu-offload "
"--overlap-cpu-optimizer-d2h-h2d "
"--use-precision-aware-optimizer "
Expand Down Expand Up @@ -120,7 +118,9 @@ def execute():
)

misc_args = (
"--actor-num-nodes 1 " f"--actor-num-gpus-per-node 8 " f"--rollout-num-gpus 8 "
"--actor-num-nodes 1 "
"--actor-num-gpus-per-node 8 "
"--rollout-num-gpus 8 "
"--no-gradient-accumulation-fusion "
"--use-flash-attn "
)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import os
import tempfile

import slime.utils.misc as U
from slime.utils.external_utils.command_utils import execute_train_npu

MODEL_NAME = os.environ.get("SLIME_SCRIPT_MODEL_NAME", "Qwen3-VL-2B-Instruct")
Expand Down Expand Up @@ -29,6 +29,10 @@ def get_megatron_model_type(model_name: str) -> str:


def execute():
critic_config = tempfile.NamedTemporaryFile("w", suffix=".yaml", delete=False)
critic_config.write("lr: 1e-5\n")
critic_config.close()

ckpt_args = f"--hf-checkpoint /path/to/model/checkpoints/{MODEL_NAME} "

wandb_args = (
Expand Down Expand Up @@ -69,7 +73,6 @@ def execute():
"--eps-clip 4e-4 "
"--num-critic-only-steps 1 "
"--normalize-advantages "
"--critic-lr 1e-5 "
)

optimizer_args = (
Expand All @@ -79,7 +82,6 @@ def execute():
"--weight-decay 0.1 "
"--adam-beta1 0.9 "
"--adam-beta2 0.98 "

"--optimizer-cpu-offload "
"--overlap-cpu-optimizer-d2h-h2d "
"--use-precision-aware-optimizer "
Expand Down Expand Up @@ -122,9 +124,7 @@ def execute():

misc_args = (
"--actor-num-nodes 1 "
"--actor-num-gpus-per-node 4 "
"--critic-num-nodes 1 "
"--critic-num-gpus-per-node 4 "
"--actor-num-gpus-per-node 8 "
"--rollout-num-gpus 8 "
"--no-gradient-accumulation-fusion "
"--use-flash-attn "
Expand All @@ -138,6 +138,7 @@ def execute():
exit()

train_args = (
f"--critic-config-path {critic_config.name} "
f"{ckpt_args} "
f"{rollout_args} "
f"{optimizer_args} "
Expand Down
85 changes: 29 additions & 56 deletions slime/backends/megatron_utils/actor.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
import logging
import os
import random
import socket
from argparse import Namespace
from contextlib import nullcontext

Expand All @@ -10,14 +9,13 @@
import torch
import torch.distributed as dist
from megatron.core import mpu
from ray.actor import ActorHandle
from torch_memory_saver import torch_memory_saver
from transformers import AutoConfig, AutoTokenizer

from slime.ray.train_actor import TrainRayActor
from slime.utils import train_dump_utils
from slime.utils.data import process_rollout_data
from slime.utils.distributed_utils import get_gloo_group, init_process_group
from slime.utils.distributed_utils import get_gloo_group
from slime.utils.logging_utils import init_tracking
from slime.utils.memory_utils import clear_memory, print_memory
from slime.utils.misc import Box
Expand All @@ -30,7 +28,7 @@
from ...utils.tensor_backper import TensorBackuper
from .checkpoint import load_checkpoint
from .cp_utils import slice_log_prob_with_cp, slice_with_cp
from .data import DataIterator, get_data_iterator, log_perf_data, log_rollout_data, sync_actor_critic_data
from .data import DataIterator, get_data_iterator, log_perf_data, log_rollout_data
from .initialize import init, is_megatron_main_rank
from .loss import compute_advantages_and_returns, get_log_probs_and_entropy, get_values
from .model import forward_only, initialize_model_and_optimizer, save, train
Expand Down Expand Up @@ -83,12 +81,6 @@ def init(
logger.info(f"Set torch_memory_saver.memory_margin_bytes to {x}")
torch_memory_saver.memory_margin_bytes = x

if role == "critic":
self.args.load = self.args.critic_load
self.args.save = self.args.critic_save
self.args.lr = self.args.critic_lr
self.args.lr_warmup_iters = self.args.critic_lr_warmup_iters

(self.model, self.optimizer, self.opt_param_scheduler, loaded_rollout_id) = initialize_model_and_optimizer(
args, role
)
Expand Down Expand Up @@ -360,9 +352,9 @@ def compute_log_prob(
store_prefix=store_prefix,
)

def train(self, rollout_id: int, rollout_data_ref: Box) -> None:
def train(self, rollout_id: int, rollout_data_ref: Box, external_data=None):
if self.args.debug_rollout_only:
return
return None

if self.args.offload_train:
self.wake_up()
Expand All @@ -371,25 +363,22 @@ def train(self, rollout_id: int, rollout_data_ref: Box) -> None:
rollout_data = self._get_rollout_data(rollout_data_ref)

if self.role == "critic":
return self.train_critic(rollout_id, rollout_data)
result = self.train_critic(rollout_id, rollout_data)
else:
return self.train_actor(rollout_id, rollout_data)
self.train_actor(rollout_id, rollout_data, external_data=external_data)
result = None

def train_critic(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
# Create data iterator for log_probs and train.
if self.args.offload_train:
self.sleep()

return result

def train_critic(self, rollout_id: int, rollout_data: RolloutBatch):
"""Train critic and return CPU values (used as old-values for the next actor train)."""
data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data)
rollout_data.update(
forward_only(
get_values,
self.args,
self.model,
data_iterator,
num_microbatches,
)
)

if rollout_id >= self.args.num_critic_only_steps and not self.args.critic_train_only:
sync_actor_critic_data(self.args, rollout_data, self._actor_critic_groups)
# Compute current critic values (used as old_values for value loss and for actor advantages).
rollout_data.update(forward_only(get_values, self.args, self.model, data_iterator, num_microbatches))

compute_advantages_and_returns(self.args, rollout_data)

Expand All @@ -403,7 +392,13 @@ def train_critic(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
num_microbatches,
)

def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
if mpu.is_pipeline_last_stage() and "values" in rollout_data:
from slime.backends.megatron_utils.data import tensors_to_cpu

return {"values": tensors_to_cpu(rollout_data["values"])}
return {}

def train_actor(self, rollout_id: int, rollout_data: RolloutBatch, external_data=None) -> None:
# Create data iterator for log_probs and train.
data_iterator, num_microbatches = get_data_iterator(self.args, self.model, rollout_data)

Expand Down Expand Up @@ -455,11 +450,12 @@ def train_actor(self, rollout_id: int, rollout_data: RolloutBatch) -> None:
RoutingReplay.clear_all_forward()

if self.args.use_critic:
sync_actor_critic_data(
self.args,
rollout_data,
self._actor_critic_groups,
)
if external_data is not None and mpu.is_pipeline_last_stage():
values = external_data.get("values")
if values is not None:
from slime.backends.megatron_utils.data import tensors_to_gpu

rollout_data["values"] = tensors_to_gpu(values)
if self._active_model_tag != "actor":
self._switch_model("actor")

Expand Down Expand Up @@ -623,26 +619,3 @@ def load_other_checkpoint(self, model_tag: str, path: str) -> None:

self.weights_backuper.backup(model_tag)
self._active_model_tag = model_tag

def connect_actor_critic(
self,
actor_handle: ActorHandle | None = None,
master_address: str | None = None,
master_port: int | None = None,
) -> None:
if self.role == "actor":
master_address = ray.util.get_node_ip_address()
with socket.socket() as sock:
sock.bind(("", 0))
master_port = sock.getsockname()[1]
actor_handle.connect_actor_critic.remote(master_address=master_address, master_port=master_port)

group_name = "actor_critic"
world_size = 2
self._actor_critic_groups = init_process_group(
backend="nccl",
init_method=f"tcp://{master_address}:{master_port}",
world_size=world_size,
rank=0 if self.role == "actor" else 1,
group_name=group_name,
)
6 changes: 2 additions & 4 deletions slime/backends/megatron_utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@

def validate_args(args):
"""Run megatron's own validate_args plus slime-specific megatron validations."""

_megatron_validate_args(args)

# always use varlen
Expand Down Expand Up @@ -116,9 +117,6 @@ def megatron_parse_args(extra_args_provider, skip_hf_validate=False):
_hf_validate_args(args, hf_config)

args.rank = 0
if args.critic_train_only:
args.world_size = args.critic_num_nodes * args.critic_num_gpus_per_node
else:
args.world_size = args.actor_num_nodes * args.actor_num_gpus_per_node
args.world_size = args.actor_num_nodes * args.actor_num_gpus_per_node
args = _set_default_megatron_args(args)
return args
73 changes: 26 additions & 47 deletions slime/backends/megatron_utils/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -610,53 +610,32 @@ def log_perf_data(rollout_id: int, args: Namespace) -> None:
)


def sync_actor_critic_data(
args: Namespace,
rollout_data: RolloutBatch | None = None,
group: dist.ProcessGroup | None = None,
) -> None:
def tensors_to_cpu(tensor_list):
"""Move a list of GPU tensors to CPU for Ray object store transfer.

Args:
tensor_list: List of GPU tensors, or None.

Returns:
List of CPU tensors (detached), or None if input is None.
"""
Broadcast `values` (from critic) and optionally `log_probs`/`ref_log_probs`
(from actor) across PP ranks to align data dependencies.
if tensor_list is None:
return None
return [t.detach().cpu() for t in tensor_list]


- Values are broadcast from src=1.
- Log-probs and ref-log-probs are broadcast from src=0 when KL is used.
Updates `rollout_data` in place with the synchronized tensors.
def tensors_to_gpu(tensor_list, device=None):
"""Move a list of CPU tensors back to GPU.

Args:
tensor_list: List of CPU tensors, or None.
device: Target CUDA device. If None, uses current device.

Returns:
List of GPU tensors, or None if input is None.
"""
log_probs_key = "log_probs" if not args.use_rollout_logprobs else "rollout_log_probs"
values, log_probs, ref_log_probs = map(rollout_data.get, ("values", log_probs_key, "ref_log_probs"))

# return when not the pp last stage
if not values and not log_probs:
return

handles = []

if not values:
values = [torch.empty_like(log_prob) for log_prob in log_probs]
for value in values:
handles.append(dist.broadcast(value, src=1, group=group, async_op=True))

if args.kl_coef != 0 or args.use_kl_loss:
if not log_probs:
log_probs = [torch.empty_like(value) for value in values]
if not ref_log_probs:
ref_log_probs = [torch.empty_like(value) for value in values]
for ref_log_prob, log_prob in zip(ref_log_probs, log_probs, strict=False):
handles.append(dist.broadcast(log_prob, src=0, group=group, async_op=True))
handles.append(dist.broadcast(ref_log_prob, src=0, group=group, async_op=True))

for handle in handles:
handle.wait()

rollout_data.update(
{
k: v
for k, v in {
"values": values,
log_probs_key: log_probs,
"ref_log_probs": ref_log_probs,
}.items()
if v is not None
}
)
if tensor_list is None:
return None
if device is None:
device = torch.cuda.current_device()
return [t.to(device=device, dtype=torch.float32) for t in tensor_list]
12 changes: 11 additions & 1 deletion slime/backends/megatron_utils/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -581,6 +581,10 @@ def compute_advantages_and_returns(args: Namespace, rollout_data: RolloutBatch)
Early returns if both `log_probs` and `values` are None (intermediate
pipeline stages).

If ``args.custom_advantage_function_path`` is set, it is called after KL computation
and must populate ``rollout_data["advantages"]`` and
``rollout_data["returns"]``.

Args:
args: Configuration specifying estimator type, KL coefficient,
normalization settings, and other hyperparameters.
Expand Down Expand Up @@ -615,8 +619,14 @@ def compute_advantages_and_returns(args: Namespace, rollout_data: RolloutBatch)
)
for i in range(len(log_probs))
]
rollout_data["kl"] = kl

if args.custom_advantage_function_path is not None:
custom_adv_fn = load_function(args.custom_advantage_function_path)
custom_adv_fn(args, rollout_data)
advantages, returns = rollout_data["advantages"], rollout_data["returns"]

Comment on lines +627 to 628

Copilot AI Apr 24, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

compute_advantages_and_returns assumes a custom advantage function will always populate rollout_data['advantages'] and rollout_data['returns'], but if it doesn't, the code will fail later with a KeyError/unclear error. Consider validating after custom_adv_fn(...) returns and raising a clear ValueError if required keys are missing or have the wrong type/shape.

Suggested change
advantages, returns = rollout_data["advantages"], rollout_data["returns"]
missing_keys = [
key for key in ("advantages", "returns") if key not in rollout_data
]
if missing_keys:
raise ValueError(
f"Custom advantage function {args.custom_advantage_fn!r} must populate "
f"rollout_data with keys 'advantages' and 'returns'; missing keys: "
f"{missing_keys}."
)
advantages = rollout_data["advantages"]
returns = rollout_data["returns"]
if not isinstance(advantages, list) or not isinstance(returns, list):
raise ValueError(
f"Custom advantage function {args.custom_advantage_fn!r} must populate "
f"rollout_data['advantages'] and rollout_data['returns'] as lists of "
f"torch.Tensor objects, got {type(advantages).__name__} and "
f"{type(returns).__name__}."
)
expected_num_samples = len(kl)
if len(advantages) != expected_num_samples or len(returns) != expected_num_samples:
raise ValueError(
f"Custom advantage function {args.custom_advantage_fn!r} returned "
f"{len(advantages)} advantages and {len(returns)} returns, expected "
f"{expected_num_samples} of each."
)
for i, (advantage, ret, k) in enumerate(
zip(advantages, returns, kl, strict=False)
):
if not isinstance(advantage, torch.Tensor) or not isinstance(ret, torch.Tensor):
raise ValueError(
f"Custom advantage function {args.custom_advantage_fn!r} must return "
f"lists of torch.Tensor objects; sample {i} has types "
f"{type(advantage).__name__} and {type(ret).__name__}."
)
if advantage.shape != k.shape or ret.shape != k.shape:
raise ValueError(
f"Custom advantage function {args.custom_advantage_fn!r} returned "
f"incompatible tensor shapes for sample {i}: advantages shape "
f"{tuple(advantage.shape)}, returns shape {tuple(ret.shape)}, "
f"expected {tuple(k.shape)}."
)

Copilot uses AI. Check for mistakes.
if args.advantage_estimator in ["grpo", "gspo"]:
elif args.advantage_estimator in ["grpo", "gspo"]:
rewards = torch.tensor(rewards, dtype=torch.float32, device=kl[0].device)
returns = get_grpo_returns(rewards, kl)
# TODO: is the copy necessary?
Expand Down
Loading
Loading