初始化项目,由ModelHub XC社区提供模型

Model: ayh015/myLightningOPD
Source: Original Platform
This commit is contained in:
ModelHub XC
2026-08-27 23:50:14 +08:00
commit d4e0a1af66
368 changed files with 559583 additions and 0 deletions

4
slime/utils/__init__.py Normal file
View File

@@ -0,0 +1,4 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Utility package root for Slime."""

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

1737
slime/utils/arguments.py Normal file

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,39 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
import threading
__all__ = ["get_async_loop", "run"]
# Create a background event loop thread
class AsyncLoopThread:
def __init__(self):
self.loop = asyncio.new_event_loop()
self._thread = threading.Thread(target=self._start_loop, daemon=True)
self._thread.start()
def _start_loop(self):
asyncio.set_event_loop(self.loop)
self.loop.run_forever()
def run(self, coro):
# Schedule a coroutine onto the loop and block until it's done
return asyncio.run_coroutine_threadsafe(coro, self.loop).result()
# Create one global instance
async_loop = None
def get_async_loop():
global async_loop
if async_loop is None:
async_loop = AsyncLoopThread()
return async_loop
def run(coro):
"""Run a coroutine in the background event loop."""
return get_async_loop().run(coro)

View File

@@ -0,0 +1,18 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from functools import wraps
def with_defer(deferred_func):
def decorator(fn):
@wraps(fn)
def wrapper(*args, **kwargs):
try:
return fn(*args, **kwargs)
finally:
deferred_func()
return wrapper
return decorator

342
slime/utils/data.py Normal file
View File

@@ -0,0 +1,342 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import json
import logging
import os
import random
import re
import numpy as np
import pandas as pd
import ray
from slime.utils.types import MultimodalTypes, Sample
from .timer import Timer
__all__ = ["Dataset", "create_dataset"]
logger = logging.getLogger(__name__)
def _read_single_file(path, row_slice=None):
"""Read a single data file (jsonl or parquet)."""
if path.endswith(".jsonl"):
df = pd.read_json(path, lines=True, dtype={"label": str})
elif path.endswith(".parquet"):
df = pd.read_parquet(path, dtype_backend="pyarrow")
else:
raise ValueError(f"Unsupported file format: {path}. Supported formats are .jsonl and .parquet.")
if row_slice is not None:
logger.info(f"read_file path={path} slice {len(df)=} rows into {row_slice=}")
df = df.iloc[row_slice]
for _, row in df.iterrows():
yield row.to_dict()
def _list_data_files(directory):
"""List all supported data files in a directory (recursively)."""
supported_extensions = ('.jsonl', '.parquet')
data_files = []
for root, _, files in os.walk(directory):
for file in sorted(files): # Sort for deterministic order
if file.endswith(supported_extensions):
data_files.append(os.path.join(root, file))
return sorted(data_files) # Sort by full path for deterministic order
def read_file(path):
"""Read data from a file or directory.
Args:
path: Path to a data file (.jsonl or .parquet) or a directory containing data files.
If a directory is provided, all .jsonl and .parquet files in it (and subdirectories)
will be read and concatenated.
Supports row slicing with @[start:end] suffix, e.g., "data.jsonl@[0:1000]"
Yields:
dict: Each row of data as a dictionary.
"""
path, row_slice = _parse_generalized_path(path)
if not os.path.exists(path):
raise FileNotFoundError(f"Prompt dataset path '{path}' does not exist.")
# Handle directory: read all data files inside
if os.path.isdir(path):
data_files = _list_data_files(path)
if not data_files:
raise ValueError(f"No .jsonl or .parquet files found in directory: {path}")
logger.info(f"Found {len(data_files)} data files in directory {path}")
# For directory, row_slice applies to the combined dataset
if row_slice is not None:
# Collect all data first, then apply slice
all_rows = []
for file_path in data_files:
for row in _read_single_file(file_path):
all_rows.append(row)
logger.info(f"read_file directory={path} slice {len(all_rows)=} rows into {row_slice=}")
for row in all_rows[row_slice]:
yield row
else:
# Stream data from each file
for file_path in data_files:
for row in _read_single_file(file_path):
yield row
else:
# Handle single file
for row in _read_single_file(path, row_slice):
yield row
def _parse_generalized_path(s: str):
if (m := re.match(r"^(?P<real_path>.*)@\[(?P<start>-?\d*):(?P<end>-?\d*)\]$", s)) is not None:
path = m.group("real_path")
start = int(x) if (x := m.group("start")) != "" else None
end = int(x) if (x := m.group("end")) != "" else None
return path, slice(start, end)
return s, None
def _should_skip_prompt(formatted_prompt: str, tokenizer, processor, max_length, multimodal_inputs=None):
if max_length is None:
return False
if processor:
processor_output = processor(text=formatted_prompt, **multimodal_inputs)
input_ids = processor_output["input_ids"][0]
else:
input_ids = tokenizer.encode(formatted_prompt, add_special_tokens=False)
return len(input_ids) > max_length
def _build_messages(data: dict, prompt_key: str, as_conversation: bool, multimodal_keys: dict = None):
prompt = data.get(prompt_key)
if isinstance(prompt, str):
if not as_conversation:
return prompt
else:
prompt = [{"role": "user", "content": prompt}]
if multimodal_keys:
assert as_conversation, "as_conversation must be True when multimodal_keys is not None"
# Build mapping: placeholder -> (MultimodalType, content_list)
multimodals = {}
for type_name, data_key in multimodal_keys.items():
mt = MultimodalTypes.get(type_name)
if mt:
multimodals[mt.placeholder] = (mt, list(data.get(data_key)))
pattern = "(" + "|".join(re.escape(p) for p in multimodals.keys()) + ")"
for message in prompt:
if isinstance(message["content"], str):
content_list = []
for segment in re.split(pattern, message["content"]):
if not segment:
continue
if segment in multimodals:
mt, content = multimodals[segment]
content_list.append({"type": mt.name, mt.name: content.pop(0)})
else:
content_list.append({"type": "text", "text": segment})
message["content"] = content_list
elif isinstance(message["content"], list):
# TODO: handle more general cases. where message['content'] is a dict and contains multiple types of content.
# e.g.
# "content": [
# {
# "type": "image",
# "image": "https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg",
# },
# {"type": "text", "text": "Describe this image."},
# ],
logger.warning("message['content'] is a list of dicts, no processing will be done.")
continue
else:
raise ValueError(
f"Unsupported content type: {type(message['content'])}, expected str or list of dicts"
)
return prompt
class Dataset:
def __init__(
self,
path,
tokenizer,
processor,
max_length,
*,
prompt_key="text",
multimodal_keys=None,
label_key=None,
tool_key=None,
metadata_key="metadata",
seed=42,
apply_chat_template=False,
apply_chat_template_kwargs=None,
):
self.origin_samples = []
for data in read_file(path):
metadata = data.get(metadata_key) or {}
prompt = _build_messages(data, prompt_key, apply_chat_template, multimodal_keys)
tools = None
if tool_key is not None and tool_key in data:
tools = data[tool_key]
if isinstance(tools, str):
tools = json.loads(tools)
elif isinstance(tools, np.ndarray):
tools = tools.tolist()
assert isinstance(tools, list), f"tools must be a list, got {type(tools)} instead"
metadata["tools"] = tools
if apply_chat_template:
formatted_prompt = tokenizer.apply_chat_template(
prompt,
tools=tools,
tokenize=False,
add_generation_prompt=True,
**(apply_chat_template_kwargs or {}),
)
else:
formatted_prompt = prompt
if processor:
# temporary solution, will write image utils for slime later
from qwen_vl_utils import process_vision_info
assert isinstance(
prompt, list
), f"prompt must be a list when processor is not None, got {type(prompt)} instead"
images, videos = process_vision_info(prompt)
multimodal_inputs = {"images": images, "videos": videos}
else:
multimodal_inputs = None
# TODO: this is slow.
if _should_skip_prompt(formatted_prompt, tokenizer, processor, max_length, multimodal_inputs):
continue
self.origin_samples.append(
Sample(
prompt=formatted_prompt,
label=data.get(label_key, None) if label_key is not None else None,
metadata=metadata,
multimodal_inputs=multimodal_inputs,
)
)
logger.info(f"Dataset: Loaded {len(self.origin_samples)} samples from {path}")
self.epoch_id = -1
self.seed = seed
self.samples = self.origin_samples
def shuffle(self, new_epoch_id):
if self.epoch_id == new_epoch_id:
return
random.seed(self.seed + new_epoch_id)
permutation = list(range(len(self.samples)))
random.shuffle(permutation)
self.samples = [self.origin_samples[i] for i in permutation]
self.epoch_id = new_epoch_id
def __getitem__(self, idx):
return self.samples[idx]
def __len__(self):
return len(self.samples)
def get_minimum_num_micro_batch_size(total_lengths, max_tokens_per_gpu):
# use first fit to get the number of micro batches
batches = []
for length in total_lengths:
for i in range(len(batches)):
if batches[i] + length <= max_tokens_per_gpu:
batches[i] += length
break
else:
batches.append(length)
return len(batches)
def process_rollout_data(args, rollout_data_ref, dp_rank, dp_size):
assert len(rollout_data_ref) == dp_size
rollout_data = ray.get(rollout_data_ref[dp_rank].inner)
partition = rollout_data.pop("partition")
total_lengths = rollout_data["total_lengths"]
# save the seqlen of the whole rollout batch
Timer().seq_lens = total_lengths
rollout_data["total_lengths"] = [total_lengths[i] for i in partition]
return rollout_data
def create_dataset(
paths,
tokenizer,
processor,
max_length,
*,
prompt_key="text",
multimodal_keys=None,
label_key=None,
tool_key=None,
metadata_key="metadata",
seed=42,
apply_chat_template=False,
apply_chat_template_kwargs=None,
):
"""Factory function to create a Dataset.
Args:
paths: A single path string, or a list with one path from --prompt-data.
Other args are the same as Dataset.
Returns:
Dataset instance.
"""
if isinstance(paths, list):
if len(paths) != 1:
raise ValueError(f"Only single-path datasets are supported, got {len(paths)} paths.")
paths = paths[0]
return Dataset(
path=paths,
tokenizer=tokenizer,
processor=processor,
max_length=max_length,
prompt_key=prompt_key,
multimodal_keys=multimodal_keys,
label_key=label_key,
tool_key=tool_key,
metadata_key=metadata_key,
seed=seed,
apply_chat_template=apply_chat_template,
apply_chat_template_kwargs=apply_chat_template_kwargs,
)

View File

@@ -0,0 +1,3 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

View File

@@ -0,0 +1,76 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Annotated
import torch
import typer
from slime.ray.rollout import compute_metrics_from_samples
from slime.utils.types import Sample
_WHITELIST_KEYS = [
"group_index",
"index",
"prompt",
"response",
"response_length",
"label",
"reward",
"status",
"metadata",
]
def main(
# Deliberately make this name consistent with main training arguments
load_debug_rollout_data: Annotated[str, typer.Option()],
show_metrics: bool = True,
show_samples: bool = True,
category: list[str] = None,
):
if category is None:
category = ["train", "eval"]
for rollout_id, path in _get_rollout_dump_paths(load_debug_rollout_data, category):
print("-" * 80)
print(f"{rollout_id=} {path=}")
print("-" * 80)
pack = torch.load(path)
sample_dicts = pack["samples"]
if show_metrics:
# TODO read these configs from dumps
args = SimpleNamespace(
advantage_estimator="grpo",
reward_key=None,
log_reward_category=None,
)
sample_objects = [Sample.from_dict(s) for s in sample_dicts]
metrics = compute_metrics_from_samples(args, sample_objects)
print("metrics", metrics)
if show_samples:
for sample in sample_dicts:
print(json.dumps({k: v for k, v in sample.items() if k in _WHITELIST_KEYS}))
def _get_rollout_dump_paths(load_debug_rollout_data: str, categories: list[str]):
# may improve later
for rollout_id in range(1000):
for category in categories:
prefix = {
"train": "",
"eval": "eval_",
}[category]
path = Path(load_debug_rollout_data.format(rollout_id=f"{prefix}{rollout_id}"))
if path.exists():
yield rollout_id, path
if __name__ == "__main__":
"""python -m slime.utils.debug_utils.display_debug_rollout_data --load-debug-rollout-data ..."""
typer.run(main)

View File

@@ -0,0 +1,53 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
from typing import Annotated
import ray
import torch
import typer
from slime.utils.misc import load_function
from slime.utils.types import Sample
def _truncate(text, max_len=200):
"""Truncate text and add ellipsis if too long."""
if text is None:
return None
text = str(text).replace("\n", "\\n")
if len(text) > max_len:
return text[:max_len] + "..."
return text
def main(
rollout_data_path: Annotated[str, typer.Option()],
custom_rm_path: Annotated[str, typer.Option()],
):
if not ray.is_initialized():
ray.init()
pack = torch.load(rollout_data_path)
samples = [Sample.from_dict(s) for s in pack["samples"]]
asyncio.run(_main_async(samples=samples, custom_rm_path=custom_rm_path))
async def _main_async(samples, custom_rm_path):
rm_function = load_function(custom_rm_path)
rewards = await asyncio.gather(*[rm_function(None, sample) for sample in samples])
for i, (sample, reward) in enumerate(zip(samples, rewards, strict=True)):
print("-" * 60)
print(f"Sample {i + 1}/{len(samples)}")
print(f" Index: {sample.index}")
print(f" Status: {sample.status}")
print(f" Reward: {reward}")
print(f" Prompt: {_truncate(sample.prompt, 200)}")
print(f" Response: {_truncate(sample.response, 200)}")
print("-" * 60)
if __name__ == "__main__":
typer.run(main)

View File

@@ -0,0 +1,61 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
import json
from typing import Annotated
import typer
from openai import AsyncOpenAI
from slime.utils.data import read_file
# can unify w/ sglang_rollout.py later, e.g. add RM, if needed
def main(
prompt_data: Annotated[str, typer.Option()],
url: Annotated[str, typer.Option()] = "http://localhost:30000/v1",
input_key: Annotated[str, typer.Option()] = "input",
n_samples_per_prompt: Annotated[int, typer.Option()] = 1,
rollout_max_response_len: Annotated[int, typer.Option()] = 1024,
rollout_temperature: Annotated[float, typer.Option()] = 1.0,
rollout_top_p: Annotated[float, typer.Option()] = 1.0,
):
"""
Minimally send prompts to SGLang using OpenAI endpoints with arguments in the same format as main Slime.
Example usage:
python -m slime.utils.debug_utils.send_to_sglang --prompt-data /root/datasets/aime-2024/aime-2024.jsonl --input-key prompt --n-samples-per-prompt 16 --rollout-max-response-len 32768 --rollout-temperature 0.8 --rollout-top-p 0.7
"""
async def _main_async():
tasks = [
asyncio.create_task(_run_one(row, row_index=row_index, repeat_index=repeat_index))
for row_index, row in enumerate(read_file(prompt_data))
for repeat_index in range(n_samples_per_prompt)
]
outputs = await asyncio.gather(*tasks)
for output in outputs:
print(json.dumps(output))
async def _run_one(row, row_index: int, repeat_index: int):
resp = await client.chat.completions.create(
messages=row[input_key],
model="dummy_model",
max_tokens=rollout_max_response_len,
temperature=rollout_temperature,
top_p=rollout_top_p,
)
return dict(
row_index=row_index,
repeat_index=repeat_index,
**row,
response=resp.choices[0].message.content,
)
client = AsyncOpenAI(api_key="dummy_key", base_url=url)
asyncio.run(_main_async())
if __name__ == "__main__":
typer.run(main)

View File

@@ -0,0 +1,157 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from datetime import timedelta
from typing import Any
import torch
import torch.distributed as dist
from torch.distributed.distributed_c10d import (
Backend,
PrefixStore,
Store,
_new_process_group_helper,
_world,
default_pg_timeout,
rendezvous,
)
GLOO_GROUP = None
def init_gloo_group():
"""Initialize Gloo group for distributed communication."""
global GLOO_GROUP
if GLOO_GROUP is None:
GLOO_GROUP = dist.new_group(backend="gloo")
return GLOO_GROUP
def get_gloo_group():
"""Get the Gloo group for distributed communication."""
global GLOO_GROUP
if GLOO_GROUP is None:
raise RuntimeError("Gloo group has not been initialized. Call _init_gloo_group() first.")
return GLOO_GROUP
# Copy from pytorch to allow creating multiple main groups.
# https://github.com/pytorch/pytorch/blob/main/torch/distributed/distributed_c10d.py
def init_process_group(
backend: str | Backend = None,
init_method: str | None = None,
timeout: timedelta | None = None,
world_size: int = -1,
rank: int = -1,
store: Store | None = None,
group_name: str = None,
pg_options: Any | None = None,
):
assert (store is None) or (init_method is None), "Cannot specify both init_method and store."
if store is not None:
assert world_size > 0, "world_size must be positive if using store"
assert rank >= 0, "rank must be non-negative if using store"
elif init_method is None:
init_method = "env://"
if backend:
backend = Backend(backend)
else:
backend = Backend("undefined")
if timeout is None:
timeout = default_pg_timeout
# backward compatible API
if store is None:
rendezvous_iterator = rendezvous(init_method, rank, world_size, timeout=timeout)
store, rank, world_size = next(rendezvous_iterator)
store.set_timeout(timeout)
# Use a PrefixStore to avoid accidental overrides of keys used by
# different systems (e.g. RPC) in case the store is multi-tenant.
store = PrefixStore(group_name, store)
# NOTE: The pg_options parameter was renamed into backend_options in PyTorch 2.6.0
# https://github.com/pytorch/pytorch/commit/a0c7029a75628cd5fa8df83c0de0ea98ee7fd844
# We need to determine the appropriate parameter name based on PyTorch version
pg_options_param_name = "backend_options" if str(torch.__version__) >= "2.6" else "pg_options"
pg, _ = _new_process_group_helper(
world_size,
rank,
[],
backend,
store,
group_name=group_name,
**{pg_options_param_name: pg_options},
timeout=timeout,
)
_world.pg_group_ranks[pg] = {i: i for i in range(world_size)}
return pg
def distributed_masked_whiten(
values: torch.Tensor,
mask: torch.Tensor,
process_group: dist.ProcessGroup | None = None,
shift_mean: bool = True,
epsilon: float = 1e-8,
):
"""
Performs whitening on a tensor using global statistics from all participating GPUs.
It calculates the global mean and variance across all ranks in the default
process group (the WORLD) and uses these global statistics to normalize the
local data on each rank.
Args:
values (torch.Tensor): The local tensor of values to whiten.
mask (torch.Tensor): The local mask corresponding to the values.
process_group: The process group for all_reduce.
If None, uses the default world group.
shift_mean (bool): If True, the output is zero-mean. Defaults to True.
epsilon (float): A small value for numerical stability.
Returns:
torch.Tensor: The locally whitened tensor using global statistics.
"""
# Calculate local intermediate statistics
local_sum = (values * mask).sum()
local_sum_sq = ((values**2) * mask).sum()
local_mask_sum = mask.sum()
stats_tensor = torch.tensor(
[local_sum, local_sum_sq, local_mask_sum],
device=values.device,
dtype=torch.float32,
)
# Aggregate via all_reduce within the DP group
dist.all_reduce(stats_tensor, group=process_group)
# Calculate global stats from aggregated results
global_sum, global_sum_sq, global_mask_sum = stats_tensor
if global_mask_sum.item() == 0:
raise ValueError("The global mask sum across all participating GPUs is zero.")
global_mean = global_sum / global_mask_sum
global_mean_sq = global_sum_sq / global_mask_sum
global_var = global_mean_sq - global_mean**2
# Bessel's correction for unbiased estimate
if global_mask_sum.item() >= 2:
bessel_correction = global_mask_sum / (global_mask_sum - 1)
global_var = global_var * bessel_correction
# Whiten local data using global stats
whitened_values = (values - global_mean) * torch.rsqrt(global_var + epsilon)
if not shift_mean:
whitened_values += global_mean
return whitened_values

208
slime/utils/eval_config.py Normal file
View File

@@ -0,0 +1,208 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass, field
from typing import Any
_MISSING = object()
# TODO: This is ugly, temporarily leave this. We should unify all the config name for dataset, default, and args. (advice from Tom.)
DATASET_RUNTIME_SPECS: dict[str, dict[str, tuple[str, ...]]] = {
"n_samples_per_eval_prompt": {
"dataset_keys": ("n_samples_per_eval_prompt",),
"default_keys": ("n_samples_per_eval_prompt",),
"arg_attrs": ("n_samples_per_eval_prompt", "n_samples_per_prompt"),
},
"temperature": {
"dataset_keys": ("temperature",),
"default_keys": ("temperature",),
"arg_attrs": ("eval_temperature", "rollout_temperature"),
},
"top_p": {
"dataset_keys": ("top_p",),
"default_keys": ("top_p",),
"arg_attrs": ("eval_top_p", "rollout_top_p"),
},
"top_k": {
"dataset_keys": ("top_k",),
"default_keys": ("top_k",),
"arg_attrs": ("eval_top_k", "rollout_top_k"),
},
"max_response_len": {
"dataset_keys": ("max_response_len",),
"default_keys": ("max_response_len",),
"arg_attrs": ("eval_max_response_len", "rollout_max_response_len"),
},
}
DATASET_SAMPLE_SPECS: dict[str, dict[str, tuple[str, ...]]] = {
"input_key": {
"dataset_keys": ("input_key",),
"default_keys": ("input_key",),
"arg_attrs": ("eval_input_key", "input_key"),
},
"label_key": {
"dataset_keys": ("label_key",),
"default_keys": ("label_key",),
"arg_attrs": ("eval_label_key", "label_key"),
},
"tool_key": {
"dataset_keys": ("tool_key",),
"default_keys": ("tool_key",),
"arg_attrs": ("eval_tool_key", "tool_key"),
},
"metadata_key": {
"dataset_keys": ("metadata_key",),
"default_keys": ("metadata_key",),
"arg_attrs": ("metadata_key",),
},
}
def _first_not_missing(*values: Any) -> Any:
for value in values:
if value is not _MISSING:
return value
return _MISSING
def _pick_from_mapping(data: dict[str, Any], key_names: tuple[str, ...] | None) -> Any:
if key_names is None:
return _MISSING
for key_name in key_names:
if key_name in data:
return data[key_name]
return _MISSING
def pick_from_args(args: Any, attrs: tuple[str, ...]) -> Any:
for attr in attrs:
value = getattr(args, attr, None)
if value is not None:
return value
return None
def _ensure_metadata_overrides(value: Any) -> dict[str, Any]:
if value is None:
return {}
if not isinstance(value, dict):
raise TypeError("metadata_overrides must be a mapping.")
return value
@dataclass
class EvalDatasetConfig:
"""Configuration for a single evaluation dataset."""
name: str
path: str
rm_type: str | None = None
# Dataset-specific overrides
input_key: str | None = None
label_key: str | None = None
tool_key: str | None = None
metadata_key: str | None = None
n_samples_per_eval_prompt: int | None = None
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
max_response_len: int | None = None
stop: list[str] | None = None
stop_token_ids: list[int] | None = None
min_new_tokens: int | None = None
metadata_overrides: dict[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
self.metadata_overrides = _ensure_metadata_overrides(self.metadata_overrides)
@property
def cache_key(self) -> tuple[Any, ...]:
"""Return a tuple uniquely identifying dataset config for caching."""
return (
self.name,
self.path,
self.input_key,
self.label_key,
self.tool_key,
self.metadata_key,
)
def inject_metadata(self, sample_metadata: Any) -> dict[str, Any]:
"""Return updated metadata merging overrides."""
if not isinstance(sample_metadata, dict):
metadata = {}
else:
metadata = dict(sample_metadata)
if self.rm_type is not None:
metadata["rm_type"] = self.rm_type
for key, value in self.metadata_overrides.items():
metadata[key] = value
return metadata
def ensure_dataset_list(config: Any) -> list[dict[str, Any]]:
"""
Normalize OmegaConf containers into a list of dicts.
Accepts either a list or dictionary keyed by dataset name.
"""
if config is None:
return []
if isinstance(config, dict):
datasets = []
for name, cfg in config.items():
dataset = dict(cfg or {})
dataset.setdefault("name", name)
datasets.append(dataset)
return datasets
if isinstance(config, (list, tuple)):
datasets = []
for item in config:
dataset = dict(item or {})
if "name" not in dataset:
raise ValueError("Each evaluation dataset entry must include a `name` field.")
datasets.append(dataset)
return datasets
raise TypeError("eval.datasets must be either a list or a mapping.")
def _apply_dataset_field_overrides(
args: Any, dataset_cfg: dict[str, Any], defaults: dict[str, Any], spec_names: dict[str, Any]
) -> None:
for field_name, spec in spec_names.items():
dataset_value = _pick_from_mapping(dataset_cfg, spec["dataset_keys"])
default_value = _pick_from_mapping(defaults, spec["default_keys"])
resolved_value = _first_not_missing(dataset_value, default_value)
if resolved_value is not _MISSING:
dataset_cfg[field_name] = resolved_value
continue
dataset_cfg[field_name] = pick_from_args(args, spec["arg_attrs"])
def build_eval_dataset_configs(
args: Any,
raw_config: Iterable[dict[str, Any]],
defaults: dict[str, Any],
) -> list[EvalDatasetConfig]:
defaults = defaults or {}
datasets: list[EvalDatasetConfig] = []
for cfg in raw_config:
cfg_dict = dict(cfg or {})
combined_specs = {**DATASET_RUNTIME_SPECS, **DATASET_SAMPLE_SPECS}
_apply_dataset_field_overrides(args, cfg_dict, defaults, combined_specs)
dataset = EvalDatasetConfig(**cfg_dict)
datasets.append(dataset)
return datasets

View File

@@ -0,0 +1,3 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

View File

@@ -0,0 +1,275 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""
This file is not for slime framework itself, but as an optional utility to easily launch slime jobs and tests.
"""
import datetime
import json
import os
import random
import time
from dataclasses import dataclass
from pathlib import Path
from slime.utils.misc import exec_command
from slime.utils.typer_utils import dataclass_cli
_ = exec_command, dataclass_cli
repo_base_dir = Path(os.path.abspath(__file__)).resolve().parents[3]
def convert_checkpoint(
model_name,
megatron_model_type,
num_gpus_per_node: int,
multinode: bool = False,
extra_args: str = "",
dir_dst: str = "/root/models",
hf_checkpoint: str | None = None,
):
hf_checkpoint = hf_checkpoint or f"/root/models/{model_name}"
# TODO shall we make it in host-mapped folder and thus can cache it to speedup CI
path_dst = f"{dir_dst}/{model_name}_torch_dist"
if Path(path_dst).exists():
print(f"convert_checkpoint skip {path_dst} since exists")
return
multinode_args = ""
if multinode:
# This variable can be provided via:
# `export SLURM_JOB_HOSTNAMES=$(scontrol show hostnames "$SLURM_JOB_NODELIST")`
print(f"{os.environ.get('SLURM_JOB_HOSTNAMES')=} {os.environ.get('SLURM_NODEID')=}")
job_hostnames = os.environ["SLURM_JOB_HOSTNAMES"].strip().split("\n")
master_addr = job_hostnames[0]
nnodes = len(job_hostnames)
node_rank = int(os.environ["SLURM_NODEID"])
multinode_args = (
f"--master-addr {master_addr} " "--master-port 23456 " f"--nnodes={nnodes} " f"--node-rank {node_rank} "
)
exec_command(
f"source {repo_base_dir}/configs/models/{megatron_model_type}.sh && "
f"PYTHONPATH=/root/Megatron-LM "
f"torchrun "
f"--nproc-per-node {num_gpus_per_node} "
f"{multinode_args}"
f"tools/convert_hf_to_torch_dist.py "
"${MODEL_ARGS[@]} "
f"--hf-checkpoint {hf_checkpoint} "
f"--save {path_dst}"
f"{extra_args}"
)
def rsync_simple(path_src: str, path_dst: str):
exec_command(f"mkdir -p {path_dst} && rsync -a --info=progress2 {path_src}/ {path_dst}")
def hf_download_dataset(full_name: str):
_, partial_name = full_name.split("/")
exec_command(f"hf download --repo-type dataset {full_name} --local-dir /root/datasets/{partial_name}")
def fp8_cast_bf16(path_src, path_dst):
if Path(path_dst).exists():
print(f"fp8_cast_bf16 skip {path_dst} since exists")
return
exec_command(
"python tools/fp8_cast_bf16.py " f"--input-fp8-hf-path {path_src} " f"--output-bf16-hf-path {path_dst} "
)
# This class can be extended by concrete scripts
@dataclass
class ExecuteTrainConfig:
cuda_core_dump: bool = False
num_nodes: int = int(os.environ.get("SLURM_JOB_NUM_NODES", "1"))
extra_env_vars: str = ""
def execute_train(
train_args: str,
num_gpus_per_node: int,
megatron_model_type: str | None,
train_script: str = "train.py",
before_ray_job_submit=None,
extra_env_vars=None,
config: ExecuteTrainConfig | None = None,
rerun: bool = False,
):
if extra_env_vars is None:
extra_env_vars = {}
if config is None:
config = ExecuteTrainConfig()
external_ray = get_bool_env_var("SLIME_SCRIPT_EXTERNAL_RAY")
master_addr = os.environ.get("MASTER_ADDR", "127.0.0.1")
train_backend_fsdp = "--train-backend fsdp" in train_args
assert train_backend_fsdp == (megatron_model_type is None)
if rerun:
exec_command(
"pkill -9 sglang; "
"sleep 3; "
f"{'' if external_ray else 'ray stop --force; '}"
f"{'' if external_ray else 'pkill -9 ray; '}"
# cannot be run in CI, o/w kill the parent script
# TODO: do we really need this kill? (or can we instead kill slime)
# "pkill -9 python; "
"pkill -9 slime; "
"sleep 3; "
f"{'' if external_ray else 'pkill -9 ray; '}"
# "pkill -9 python; "
"pkill -9 slime; "
"pkill -9 redis; "
"true; "
)
if not external_ray:
exec_command(
# will prevent ray from buffering stdout/stderr
f"export PYTHONBUFFERED=16 && "
f"ray start --head --node-ip-address {master_addr} --num-gpus {num_gpus_per_node} --disable-usage-stats"
)
if (f := before_ray_job_submit) is not None:
f()
runtime_env_json = json.dumps(
{
"env_vars": {
"PYTHONPATH": "/root/Megatron-LM/",
# If setting this in FSDP, the computation communication overlapping may have issues
**(
{}
if train_backend_fsdp
else {
"CUDA_DEVICE_MAX_CONNECTIONS": "1",
}
),
"NCCL_NVLS_ENABLE": str(int(check_has_nvlink())),
"no_proxy": f"127.0.0.1,{master_addr}",
# This is needed by megatron / torch distributed in multi-node setup
"MASTER_ADDR": master_addr,
**(
{
"CUDA_ENABLE_COREDUMP_ON_EXCEPTION": "1",
"CUDA_COREDUMP_SHOW_PROGRESS": "1",
"CUDA_COREDUMP_GENERATION_FLAGS": "skip_nonrelocated_elf_images,skip_global_memory,skip_shared_memory,skip_local_memory,skip_constbank_memory",
"CUDA_COREDUMP_FILE": "/root/shared_data/cuda_coredump_%h.%p.%t",
}
if config.cuda_core_dump
else {}
),
**extra_env_vars,
**_parse_extra_env_vars(config.extra_env_vars),
}
}
)
if get_bool_env_var("SLIME_SCRIPT_ENABLE_RAY_SUBMIT", "1"):
cmd_megatron_model_source = (
f'source "{repo_base_dir}/configs/models/{megatron_model_type}.sh" && '
if megatron_model_type is not None
else ""
)
exec_command(
f"export no_proxy=127.0.0.1 && export PYTHONBUFFERED=16 && "
f"{cmd_megatron_model_source}"
f'ray job submit --address="http://127.0.0.1:8265" '
f"--runtime-env-json='{runtime_env_json}' "
f"-- python3 {train_script} "
f"{'${MODEL_ARGS[@]}' if megatron_model_type is not None else ''} "
f"{train_args}"
)
def _parse_extra_env_vars(text: str):
try:
return json.loads(text)
except ValueError:
return {kv[0]: kv[1] for item in text.split(" ") if item.strip() != "" if (kv := item.split("=")) or True}
def check_has_nvlink():
output = exec_command("nvidia-smi topo -m 2>/dev/null | grep -o 'NV[0-9][0-9]*' | wc -l", capture_output=True)
return int(output) > 0
def get_default_wandb_args(test_file: str, run_name_prefix: str | None = None, run_id: str | None = None):
if not os.environ.get("WANDB_API_KEY"):
print("Skip wandb configuration since WANDB_API_KEY is not found")
return ""
test_file = Path(test_file)
test_name = test_file.stem
if len(test_name) < 6:
test_name = f"{test_file.parent.name}_{test_name}"
wandb_run_name = run_id or create_run_id()
if (x := os.environ.get("GITHUB_COMMIT_NAME")) is not None:
wandb_run_name += f"_{x}"
if (x := run_name_prefix) is not None:
wandb_run_name = f"{x}_{wandb_run_name}"
# do not put wandb_api_key value here to avoid leaking to logs explicitly
return (
"--use-wandb "
f"--wandb-project slime-{test_name} "
f"--wandb-group {wandb_run_name} "
f"--wandb-key ${{WANDB_API_KEY}} "
"--disable-wandb-random-suffix "
)
def create_run_id() -> str:
return datetime.datetime.utcnow().strftime("%y%m%d-%H%M%S") + f"-{random.Random().randint(0, 999):03d}"
_warned_bool_env_var_keys = set()
# copied from SGLang
def get_bool_env_var(name: str, default: str = "false") -> bool:
value = os.getenv(name, default)
value = value.lower()
truthy_values = ("true", "1")
falsy_values = ("false", "0")
if (value not in truthy_values) and (value not in falsy_values):
if value not in _warned_bool_env_var_keys:
print(f"get_bool_env_var({name}) see non-understandable value={value} and treat as false")
_warned_bool_env_var_keys.add(value)
return value in truthy_values
def get_env_enable_infinite_run():
return get_bool_env_var("SLIME_TEST_ENABLE_INFINITE_RUN", "false")
def save_to_temp_file(text: str, ext: str):
path = Path(f"/tmp/slime_temp_file_{time.time()}_{random.randrange(0, 10000000)}.{ext}")
path.write_text(text)
print(f"Write the following content to {path=}: {text=}")
return str(path)
NUM_GPUS_OF_HARDWARE = {
"H100": 8,
"GB200": 4,
"GB300": 4,
}
GENERATION_HARDWARE = {
"H100": "Hopper",
"GB200": "Blackwell",
"GB300": "Blackwell",
}

130
slime/utils/flops_utils.py Normal file
View File

@@ -0,0 +1,130 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
def calculate_embedding_flops(seqlen, hidden_size):
return 2 * seqlen * hidden_size
def calculate_lm_head_flops(seqlen, hidden_size, vocab_size):
return 2 * seqlen * hidden_size * vocab_size
def calculate_qkv_projection_flops(args, seqlen, hidden_size, num_attention_heads, num_query_groups):
if args.q_lora_rank is None:
q_flops = 2 * seqlen * hidden_size * num_attention_heads * args.kv_channels
else:
q_flops = (
2
* seqlen
* args.q_lora_rank
* (args.hidden_size + args.num_attention_heads * (args.qk_head_dim + args.qk_pos_emb_head_dim))
)
if args.kv_lora_rank is None:
kv_flops = 2 * 2 * seqlen * hidden_size * num_query_groups * args.kv_channels
else:
kv_flops = (
2
* seqlen
* (
args.kv_lora_rank
* (args.hidden_size + args.num_attention_heads * (args.qk_head_dim + args.v_head_dim))
+ args.hidden_size * args.qk_pos_emb_head_dim
)
)
return q_flops + kv_flops
def calculate_attention_flops(args, seqlen, num_attention_heads):
# QK^T with causal
if args.qk_pos_emb_head_dim:
flops = 2 * num_attention_heads * seqlen * seqlen * (args.qk_head_dim + args.qk_pos_emb_head_dim) / 2
else:
flops = 2 * num_attention_heads * seqlen * seqlen * args.kv_channels / 2
# A*V
if args.v_head_dim:
flops += num_attention_heads * seqlen * seqlen * args.v_head_dim
else:
flops += num_attention_heads * seqlen * seqlen * args.kv_channels
return flops
def calculate_output_flops(seqlen, hidden_size):
return 2 * seqlen * hidden_size * hidden_size
def calculate_mlp_flops(seqlen, hidden_size, ffn_hidden_size):
return 2 * seqlen * hidden_size * ffn_hidden_size * 3
def calculate_layer_flops(args, seqlen, hidden_size, num_attention_heads, num_query_groups, ffn_hidden_size):
return (
calculate_qkv_projection_flops(args, seqlen, hidden_size, num_attention_heads, num_query_groups)
+ calculate_attention_flops(args, seqlen, num_attention_heads)
+ calculate_output_flops(seqlen, hidden_size)
+ calculate_mlp_flops(seqlen, hidden_size, ffn_hidden_size)
)
def calculate_fwd_flops(
seqlens,
args,
):
hidden_size = args.hidden_size
num_attention_heads = args.num_attention_heads
num_query_groups = args.num_query_groups
vocab_size = args.vocab_size
total_flops = 0
dense_ffn = args.ffn_hidden_size
if args.num_experts is None:
num_dense_layers = args.num_layers
num_moe_layers = 0
else:
shared_expert_ffn = getattr(args, "moe_shared_expert_intermediate_size", None)
if shared_expert_ffn is None:
shared_expert_ffn = 0
moe_ffn = args.moe_ffn_hidden_size * args.moe_router_topk + shared_expert_ffn
if hasattr(args, "moe_layer_freq"):
if isinstance(args.moe_layer_freq, list):
num_dense_layers = sum(1 for freq in args.moe_layer_freq if freq == 0)
num_moe_layers = sum(1 for freq in args.moe_layer_freq if freq > 0)
else:
num_dense_layers = sum(1 for i in range(args.num_layers) if i % args.moe_layer_freq != 0)
num_moe_layers = sum(1 for i in range(args.num_layers) if i % args.moe_layer_freq == 0)
else:
num_dense_layers = 0
num_moe_layers = args.num_layers
for seqlen in seqlens:
if num_dense_layers > 0:
total_flops += (
calculate_layer_flops(
args,
seqlen,
hidden_size,
num_attention_heads,
num_query_groups,
dense_ffn,
)
* num_dense_layers
)
if num_moe_layers > 0:
total_flops += (
calculate_layer_flops(
args,
seqlen,
hidden_size,
num_attention_heads,
num_query_groups,
moe_ffn,
)
* num_moe_layers
)
total_flops += calculate_lm_head_flops(seqlen, hidden_size, vocab_size)
return total_flops

82
slime/utils/fp8_kernel.py Normal file
View File

@@ -0,0 +1,82 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
import triton
import triton.language as tl
fp8_dtype = torch.float8_e4m3fn
fp8_max = torch.finfo(fp8_dtype).max
fp8_min = -fp8_max
def ceil_div(x: int, y: int) -> int:
"""
Perform ceiling division of two integers.
Args:
x: the dividend.
y: the divisor.
Returns:
The result of the ceiling division.
"""
return (x + y - 1) // y
@triton.jit
def _blockwise_cast_to_fp8_triton(
X,
Y,
S,
stride_xm,
stride_xn,
stride_ym,
stride_yn,
stride_sm,
stride_sn,
M,
N,
eps,
fp8_min,
fp8_max,
BLOCK_M: tl.constexpr = 32,
BLOCK_N: tl.constexpr = 128,
):
pid_m = tl.cast(tl.program_id(axis=0), tl.int64)
pid_n = tl.cast(tl.program_id(axis=1), tl.int64)
off_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
off_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
mask_m = off_m < M
mask_n = off_n < N
mask = mask_m[:, None] & mask_n[None, :]
x = tl.load(X + off_m[:, None] * stride_xm + off_n[None, :] * stride_xn, mask=mask, other=0.0).to(tl.float32)
_absmax = tl.maximum(tl.max(tl.abs(x)), eps)
x_s = _absmax / fp8_max
s_inv = 1.0 / x_s
y_q = tl.clamp(x * s_inv, fp8_min, fp8_max).to(Y.dtype.element_ty)
tl.store(Y + off_m[:, None] * stride_ym + off_n[None, :] * stride_yn, y_q, mask=mask)
tl.store(S + pid_m * stride_sm + pid_n * stride_sn, x_s)
def blockwise_cast_to_fp8_triton(x: torch.Tensor, block_size=None) -> tuple[torch.Tensor, torch.Tensor]:
BLOCK_M, BLOCK_N = 128, 128
if block_size:
BLOCK_M, BLOCK_N = block_size[0], block_size[1]
M, N = x.shape
y = torch.empty(M, N, device=x.device, dtype=torch.float8_e4m3fn)
s = torch.empty(ceil_div(M, BLOCK_M), ceil_div(N, BLOCK_N), dtype=torch.float32, device=x.device)
def grid(meta):
return (triton.cdiv(M, meta["BLOCK_M"]), triton.cdiv(N, meta["BLOCK_N"]))
if x.is_contiguous():
kwargs = {"BLOCK_M": BLOCK_M, "BLOCK_N": BLOCK_N, "num_warps": 8, "num_stages": 2}
else:
kwargs = {"BLOCK_M": BLOCK_M, "BLOCK_N": BLOCK_N, "num_warps": 1, "num_stages": 4}
_blockwise_cast_to_fp8_triton[grid](
x, y, s, *x.stride(), *y.stride(), *s.stride(), M, N, 1e-10, fp8_min, fp8_max, **kwargs
)
return y, s

View File

@@ -0,0 +1,106 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import threading
import ray
logger = logging.getLogger(__name__)
class RolloutHealthMonitor:
def __init__(self, rollout_manager, args):
# TODO may remove this dependency after refactoring
self._rollout_manager = rollout_manager
self._thread = None
self._stop_event = None
self._check_interval = args.rollout_health_check_interval
self._check_timeout = args.rollout_health_check_timeout
self._check_first_wait = args.rollout_health_check_first_wait
def start(self) -> bool:
if not self._rollout_manager.rollout_engines:
return False
assert self._thread is None, "Health monitor thread is already running."
logger.info("Starting RolloutHealthMonitor...")
self._stop_event = threading.Event()
self._thread = threading.Thread(
target=self._health_monitor_loop,
name="RolloutHealthMonitor",
daemon=True,
)
self._thread.start()
logger.info("RolloutHealthMonitor started.")
return True
def stop(self) -> None:
if not self._thread:
return
logger.info("Stopping RolloutHealthMonitor...")
assert self._stop_event is not None
self._stop_event.set()
timeout = self._check_timeout + self._check_interval + 5
self._thread.join(timeout=timeout)
if self._thread.is_alive():
logging.warning("Rollout health monitor thread did not terminate within %.1fs", timeout)
else:
logger.info("RolloutHealthMonitor stopped.")
self._thread = None
self._stop_event = None
def _health_monitor_loop(self) -> None:
assert self._stop_event is not None
logger.info(f"Health monitor loop started. Waiting for first wait: {self._check_first_wait}s")
# TODO: need to be waiting for the large moe to be ready. this is hacky.
if self._stop_event.wait(self._check_first_wait):
logger.info("Health monitor stopped during first wait.")
return
while not self._stop_event.is_set():
self._run_health_checks()
if self._stop_event.wait(self._check_interval):
break
def _run_health_checks(self) -> None:
for rollout_engine_id, engine in enumerate(self._rollout_manager.rollout_engines):
if self._stop_event is not None and self._stop_event.is_set():
break
self._check_engine_health(rollout_engine_id, engine)
def _check_engine_health(self, rollout_engine_id, engine) -> None:
if engine is None:
logger.info(f"Skipping health check for engine {rollout_engine_id} (None)")
return
try:
ray.get(engine.health_generate.remote(timeout=self._check_timeout))
except Exception as e:
logger.error(
f"Health check failed for rollout engine {rollout_engine_id} (ray timeout or error). Killing actor. Exception: {e}"
)
self._kill_engine(rollout_engine_id=rollout_engine_id)
def _kill_engine(self, rollout_engine_id: int):
logger.info(f"Killing engine group {rollout_engine_id}...")
for i in range(
rollout_engine_id * self._rollout_manager.nodes_per_engine,
(rollout_engine_id + 1) * self._rollout_manager.nodes_per_engine,
):
engine = self._rollout_manager.all_rollout_engines[i]
if engine:
logger.info(f"Shutting down and killing engine at index {i}")
try:
ray.get(engine.shutdown.remote())
ray.kill(engine)
logger.info(f"Successfully killed engine at index {i}")
except Exception as e:
logger.warning(f"Fail to kill engine at index {i} (e: {e})")
else:
logger.info(f"Engine at index {i} is already None")
self._rollout_manager.all_rollout_engines[i] = None

263
slime/utils/http_utils.py Normal file
View File

@@ -0,0 +1,263 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
import ipaddress
import json
import logging
import multiprocessing
import os
import random
import socket
import httpx
logger = logging.getLogger(__name__)
SLIME_HOST_IP_ENV = "SLIME_HOST_IP"
def find_available_port(base_port: int):
port = base_port + random.randint(100, 1000)
while True:
if is_port_available(port):
return port
if port < 60000:
port += 42
else:
port -= 43
def is_port_available(port):
"""Return whether a port is available."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
try:
s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
s.bind(("", port))
s.listen(1)
return True
except OSError:
return False
except OverflowError:
return False
def get_host_info():
hostname = socket.gethostname()
if env_overwrite_local_ip := os.getenv(SLIME_HOST_IP_ENV, None):
return hostname, env_overwrite_local_ip
# try DNS
try:
return hostname, socket.gethostbyname(hostname)
except socket.gaierror:
pass
# try IPv4
try:
with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as udp_sock:
udp_sock.connect(("8.8.8.8", 80)) # Google DNS
return hostname, udp_sock.getsockname()[0]
except OSError:
pass
# try IPv6
try:
with socket.socket(socket.AF_INET6, socket.SOCK_DGRAM) as s6:
s6.connect(("2001:4860:4860::8888", 80))
return hostname, s6.getsockname()[0]
except OSError:
pass
# hostname -I
try:
local_ip = os.popen("hostname -I | awk '{print $1}'").read().strip()
return hostname, local_ip or "::1"
except Exception:
return hostname, "::1"
def _wrap_ipv6(host):
"""Wrap IPv6 address in [] if needed."""
try:
ipaddress.IPv6Address(host.strip("[]"))
return f"[{host.strip('[]')}]"
except ipaddress.AddressValueError:
return host
def run_router(args):
try:
from sglang_router.launch_router import launch_router
router = launch_router(args)
if router is None:
return 1
return 0
except Exception as e:
logger.info(e)
return 1
def terminate_process(process: multiprocessing.Process, timeout: float = 1.0) -> None:
"""Terminate a process gracefully, with forced kill as fallback.
Args:
process: The process to terminate
timeout: Seconds to wait for graceful termination before forcing kill
"""
if not process.is_alive():
return
process.terminate()
process.join(timeout=timeout)
if process.is_alive():
process.kill()
process.join()
_http_client: httpx.AsyncClient | None = None
_client_concurrency: int = 0
# Optional Ray-based distributed POST dispatch
_distributed_post_enabled: bool = False
_post_actors: list[object] = []
_post_actor_idx: int = 0
def _next_actor():
global _post_actor_idx
if not _post_actors:
return None
actor = _post_actors[_post_actor_idx % len(_post_actors)]
_post_actor_idx = (_post_actor_idx + 1) % len(_post_actors)
return actor
async def _post(client, url, payload, max_retries=60):
retry_count = 0
while retry_count < max_retries:
try:
response = await client.post(url, json=payload or {})
response.raise_for_status()
try:
output = response.json()
except json.JSONDecodeError:
output = response.text
except Exception as e:
retry_count += 1
if isinstance(e, httpx.HTTPStatusError):
response_text = e.response.text
else:
response_text = None
logger.info(
f"Error: {e}, retrying... (attempt {retry_count}/{max_retries}, url={url}, response={response_text})"
)
if retry_count >= max_retries:
logger.info(f"Max retries ({max_retries}) reached, failing... (url={url})")
raise e
await asyncio.sleep(1)
continue
break
return output
def init_http_client(args):
"""Initialize HTTP client and optionally enable distributed POST via Ray."""
global _http_client, _client_concurrency, _distributed_post_enabled
if not args.rollout_num_gpus:
return
_client_concurrency = args.sglang_server_concurrency * args.rollout_num_gpus // args.rollout_num_gpus_per_engine
if _http_client is None:
_http_client = httpx.AsyncClient(
limits=httpx.Limits(max_connections=_client_concurrency),
timeout=httpx.Timeout(None),
)
# Optionally initialize distributed POST via Ray without changing interfaces
if args.use_distributed_post:
_init_ray_distributed_post(args)
_distributed_post_enabled = True
def _init_ray_distributed_post(args):
"""Initialize one or more Ray async actors per node for HTTP POST.
Uses NodeAffinitySchedulingStrategy to place actors on distinct nodes.
Controlled by SLIME_HTTP_POST_ACTORS_PER_NODE.
"""
global _post_actors
if _post_actors:
return # Already initialized
import ray
from ray.util.scheduling_strategies import NodeAffinitySchedulingStrategy
# Discover alive nodes
nodes = [n for n in ray.nodes() if n.get("Alive")]
if not nodes:
raise RuntimeError("No alive Ray nodes to place HTTP POST actors.")
# Define the async actor
@ray.remote
class _HttpPosterActor:
def __init__(self, concurrency: int):
# Lazy creation to this actor's event loop
self._client = httpx.AsyncClient(
limits=httpx.Limits(max_connections=max(1, concurrency)),
timeout=httpx.Timeout(None),
)
async def do_post(self, url, payload, max_retries=60):
return await _post(self._client, url, payload, max_retries)
# Create actors per node
created = []
# Distribute client concurrency across actors (at least 1 per actor)
per_actor_conc = (_client_concurrency + len(nodes)) // len(nodes)
for node in nodes:
node_id = node["NodeID"]
scheduling = NodeAffinitySchedulingStrategy(node_id=node_id, soft=False)
for _ in range(args.num_gpus_per_node):
actor = _HttpPosterActor.options(
name=None,
lifetime="detached",
scheduling_strategy=scheduling,
max_concurrency=per_actor_conc,
# Use tiny CPU to schedule
num_cpus=0.001,
).remote(per_actor_conc)
created.append(actor)
_post_actors = created
async def post(url, payload, max_retries=60):
# If distributed mode is enabled and actors exist, dispatch via Ray.
if _distributed_post_enabled and _post_actors:
try:
import ray
actor = _next_actor()
if actor is not None:
# Use a thread to avoid blocking the event loop on ray.get
obj_ref = actor.do_post.remote(url, payload, max_retries)
return await asyncio.to_thread(ray.get, obj_ref)
except Exception as e:
logger.info(f"[http_utils] Distributed POST failed, falling back to local: {e} (url={url})")
# fall through to local
return await _post(_http_client, url, payload, max_retries)
async def get(url):
response = await _http_client.get(url)
response.raise_for_status()
output = response.json()
return output

45
slime/utils/iter_utils.py Normal file
View File

@@ -0,0 +1,45 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from collections import defaultdict
from collections.abc import Callable, Iterable
from typing import Any
import torch
# details: https://stackoverflow.com/questions/773/how-do-i-use-itertools-groupby
def group_by(iterable, key=None):
"""Similar to itertools.groupby, but do not require iterable to be sorted"""
ret = defaultdict(list)
for item in iterable:
ret[key(item) if key is not None else item].append(item)
return dict(ret)
# TODO fsdp can also use this
def chunk_named_params_by_size(named_params: Iterable[tuple[str, torch.Tensor]], chunk_size: int):
return _chunk_by_size(
named_params,
compute_size=lambda named_weight: named_weight[1].nbytes,
chunk_size=chunk_size,
)
def _chunk_by_size(objects: Iterable[Any], compute_size: Callable[[Any], int], chunk_size: int):
bucket: list[Any] = []
bucket_size = 0
for obj in objects:
obj_size = compute_size(obj)
if bucket and (bucket_size + obj_size) >= chunk_size:
yield bucket
bucket = []
bucket_size = 0
bucket.append(obj)
bucket_size += obj_size
if bucket:
yield bucket

View File

@@ -0,0 +1,22 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
_LOGGER_CONFIGURED = False
# ref: SGLang
def configure_logger(prefix: str = ""):
global _LOGGER_CONFIGURED
if _LOGGER_CONFIGURED:
return
_LOGGER_CONFIGURED = True
logging.basicConfig(
level=logging.INFO,
format=f"[%(asctime)s{prefix}] %(filename)s:%(lineno)d - %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
force=True,
)

185
slime/utils/mask_utils.py Normal file
View File

@@ -0,0 +1,185 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from transformers import AutoTokenizer
def get_response_lengths(loss_masks: list[list[int]]) -> list[int]:
return [mask.count(1) if 1 in mask else 0 for mask in loss_masks]
class MultiTurnLossMaskGenerator:
def __init__(self, tokenizer: AutoTokenizer, tokenizer_type: str = "qwen"):
self.tokenizer = tokenizer
self.system_message_length, self.gen_token_length = self.get_system_message_length()
self.tokenizer_type = tokenizer_type
def get_response_lengths(self, loss_masks: list[list[int]]) -> list[int]:
return get_response_lengths(loss_masks)
def find_all_sublist_indices(self, main_list, sublist):
sublist_len = len(sublist)
indices = []
for i in range(len(main_list) - sublist_len + 1):
if main_list[i : i + sublist_len] == sublist:
indices.append(i)
return indices
def get_system_message_length(self) -> tuple[int, int]:
test_string = "FOR TESTING ONLY"
test_messages = [
{"role": "user", "content": test_string},
{"role": "user", "content": test_string},
]
raw_token_ids = self.tokenizer(test_string, add_special_tokens=False)["input_ids"]
chat_template_token = self.tokenizer.apply_chat_template(
test_messages, add_special_tokens=False, tokenize=False
)
chat_template_token_ids = self.tokenizer(chat_template_token, add_special_tokens=False)["input_ids"]
idx_1, idx_2 = self.find_all_sublist_indices(chat_template_token_ids, raw_token_ids)
end_interval = len(chat_template_token_ids) - len(raw_token_ids) - idx_2
gen_token_length = len(
self.tokenizer.apply_chat_template(
test_messages, add_special_tokens=False, tokenize=True, add_generation_prompt=True
)
) - len(chat_template_token_ids)
system_message_length = idx_1 - ((idx_2 - idx_1) - end_interval - len(raw_token_ids))
return system_message_length, gen_token_length
def gen_multi_turn_loss_mask_qwen(
self, messages: list[dict], tools: list[dict] = None
) -> tuple[list[int], list[int]]:
all_loss_masks = []
all_token_ids = []
for i, message in enumerate(messages):
if i == 0:
message_ids = self.tokenizer.apply_chat_template([message], tokenize=True, tools=tools)
else:
message_ids = self.tokenizer.apply_chat_template([message], tokenize=True)
if message["role"] != "system" and i > 0:
message_ids = message_ids[self.system_message_length :]
if message["role"] == "assistant":
loss_mask = [0] * self.gen_token_length + [1] * (len(message_ids) - self.gen_token_length)
else:
loss_mask = [0] * len(message_ids)
if message.get("step_loss_mask", 1) != 1:
loss_mask = [0] * len(message_ids)
all_loss_masks.extend(loss_mask)
all_token_ids.extend(message_ids)
return all_token_ids, all_loss_masks
def gen_multi_turn_loss_mask_qwen3(
self, messages: list[dict], tools: list[dict] = None
) -> tuple[list[int], list[int]]:
all_loss_masks = []
all_token_ids = []
prefix_message = {"role": "user", "content": "FOR CALCULATING LOSS MASK ONLY"}
prefix_token_ids = self.tokenizer.apply_chat_template([prefix_message], tokenize=True)
for i, message in enumerate(messages):
if i == 0:
tailed_message_ids = self.tokenizer.apply_chat_template(
[message, prefix_message], tokenize=True, tools=tools
)
message_ids = tailed_message_ids[: -len(prefix_token_ids)]
else:
prefixed_message_ids = self.tokenizer.apply_chat_template([prefix_message, message], tokenize=True)
message_ids = prefixed_message_ids[len(prefix_token_ids) :]
if message["role"] != "system" and i > 0:
message_ids = message_ids[self.system_message_length :]
if message["role"] == "assistant":
loss_mask = [0] * self.gen_token_length + [1] * (len(message_ids) - self.gen_token_length)
else:
loss_mask = [0] * len(message_ids)
if message.get("step_loss_mask", 1) != 1:
loss_mask = [0] * len(message_ids)
all_loss_masks.extend(loss_mask)
all_token_ids.extend(message_ids)
return all_token_ids, all_loss_masks
def gen_multi_turn_loss_mask_distill_qwen(
self, messages: list[dict], tools: list[dict] = None
) -> tuple[list[int], list[int]]:
prompt = self.tokenizer.apply_chat_template(
messages[:1], tokenize=False, add_generation_prompt=True, tools=tools
)
response = messages[-1]["content"]
prompt_tokens = self.tokenizer(prompt, add_special_tokens=False)["input_ids"]
response_tokens = self.tokenizer(response, add_special_tokens=False)["input_ids"]
response_length = len(response_tokens)
token_ids = prompt_tokens + response_tokens
loss_mask = [0] * len(prompt_tokens) + [1] * response_length
if messages[-1].get("step_loss_mask", 1) != 1:
loss_mask = [0] * len(token_ids)
return token_ids, loss_mask
def get_loss_mask(self, messages: list[dict], tools: list[dict] = None) -> tuple[list[int], list[int]]:
if self.tokenizer_type == "qwen":
if "<Assistant>" in self.tokenizer.get_added_vocab():
return self.gen_multi_turn_loss_mask_distill_qwen(messages, tools)
return self.gen_multi_turn_loss_mask_qwen(messages, tools)
elif self.tokenizer_type == "qwen3":
return self.gen_multi_turn_loss_mask_qwen3(messages, tools)
elif self.tokenizer_type == "distill_qwen":
return self.gen_multi_turn_loss_mask_distill_qwen(messages, tools)
else:
raise ValueError(f"Unsupported tokenizer type: {self.tokenizer_type}")
def get_loss_mask_with_multimodal_alignment(
self, messages: list[dict], input_ids: list[int], tools: list[dict] = None
) -> tuple[list[int], list[int]]:
text = []
for msg in messages:
if isinstance(msg.get("content"), list):
text_parts = []
for item in msg["content"]:
if isinstance(item, dict) and item.get("type") == "text":
text_parts.append(item.get("text", ""))
elif isinstance(item, str):
text_parts.append(item)
text.append({"role": msg["role"], "content": " ".join(text_parts)})
else:
text.append(msg)
_, loss_mask_text = self.get_loss_mask(text, tools=tools)
diff = len(input_ids) - len(loss_mask_text)
assert diff >= 0, (
f"input_ids (length={len(input_ids)}) is shorter than text loss_mask (length={len(loss_mask_text)}) "
f"Please check if processor and tokenizer tokenization are consistent."
)
loss_mask = [0] * diff + loss_mask_text
return input_ids, loss_mask
def get_text_from_loss_mask(self, token_ids: list[int], loss_masks: list[int]) -> list[str]:
selected_texts = []
current_tokens = []
for idx, mask in enumerate(loss_masks):
if mask == 1:
current_tokens.append(token_ids[idx])
elif current_tokens:
selected_texts.append(self.tokenizer.decode(current_tokens))
current_tokens = []
if current_tokens:
selected_texts.append(self.tokenizer.decode(current_tokens))
return selected_texts

View File

@@ -0,0 +1,25 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from contextlib import contextmanager
try:
from megatron.core.utils import unwrap_model
except ImportError:
unwrap_model = None
@contextmanager
def patch_megatron_model(model):
unwrapped_model = unwrap_model(model)[0]
model_config = unwrapped_model.config
attribute_was_added = False
if not hasattr(model_config, "share_embeddings_and_output_weights"):
model_config.share_embeddings_and_output_weights = unwrapped_model.share_embeddings_and_output_weights
attribute_was_added = True
try:
yield
finally:
if attribute_was_added:
delattr(model_config, "share_embeddings_and_output_weights")

View File

@@ -0,0 +1,47 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import gc
import logging
import torch
import torch.distributed as dist
logger = logging.getLogger(__name__)
def clear_memory(clear_host_memory: bool = False):
torch.cuda.synchronize()
gc.collect()
torch.cuda.empty_cache()
if clear_host_memory:
torch._C._host_emptyCache()
def available_memory():
device = torch.cuda.current_device()
free, total = torch.cuda.mem_get_info(device)
return {
"gpu": str(device),
"total_GB": _byte_to_gb(total),
"free_GB": _byte_to_gb(free),
"used_GB": _byte_to_gb(total - free),
"allocated_GB": _byte_to_gb(torch.cuda.memory_allocated(device)),
"reserved_GB": _byte_to_gb(torch.cuda.memory_reserved(device)),
}
def _byte_to_gb(n: int):
return round(n / (1024**3), 2)
def print_memory(msg, clear_before_print: bool = False):
if clear_before_print:
clear_memory()
memory_info = available_memory()
# Need to print for all ranks, b/c different rank can have different behaviors
logger.info(
f"[Rank {dist.get_rank()}] Memory-Usage {msg}{' (cleared before print)' if clear_before_print else ''}: {memory_info}"
)
return memory_info

View File

@@ -0,0 +1,31 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
logger = logging.getLogger(__name__)
class MetricChecker:
@staticmethod
def maybe_create(args):
if args.ci_test and (args.ci_metric_checker_key is not None):
return MetricChecker(args)
return None
def __init__(self, args):
self.args = args
self._exists_check_success = False
def on_eval(self, metrics: dict[str, float]):
actual_value = metrics.get(self.args.ci_metric_checker_key)
assert actual_value is not None, f"{metrics=} {self.args.ci_metric_checker_key=}"
check_success = actual_value >= self.args.ci_metric_checker_threshold
logger.info(f"[MetricChecker] {check_success=} {actual_value=} {self.args.ci_metric_checker_threshold=}")
self._exists_check_success |= check_success
def dispose(self):
assert self._exists_check_success, "[MetricChecker] accuracy check failed"
logger.info("[MetricChecker] pass dispose check")

121
slime/utils/metric_utils.py Normal file
View File

@@ -0,0 +1,121 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any, Literal
import numpy as np
def dict_add_prefix(d: dict[str, Any], prefix: str) -> dict[str, Any]:
return {f"{prefix}{k}": v for k, v in d.items()}
def compute_pass_rate(
flat_rewards: list[float],
group_size: int,
num_groups: int | None = None,
):
if group_size == 1:
return {}
if num_groups is None:
num_groups = len(flat_rewards) // group_size
pass_rate_name_list = [2**i for i in range(int(math.log2(group_size)) + 1)]
assert len(flat_rewards) == num_groups * group_size, f"{len(flat_rewards)=} {num_groups=} {group_size=}"
rewards_of_group = np.array(flat_rewards).reshape(num_groups, group_size)
log_dict = {}
for k in pass_rate_name_list:
num_correct = np.sum(rewards_of_group == 1, axis=1)
num_samples = np.full(num_groups, group_size)
pass_k_estimates = _estimate_pass_at_k(num_samples, num_correct, k)
pass_k = np.mean(pass_k_estimates)
log_dict[f"pass@{k}"] = pass_k
return log_dict
def _estimate_pass_at_k(num_samples, num_correct, k):
"""
Estimates pass@k of each problem and returns them in an array.
"""
def estimator(n, c, k):
"""
Calculates 1 - comb(n - c, k) / comb(n, k).
"""
if n - c < k:
return 1.0
return 1.0 - np.prod(1.0 - k / np.arange(n - c + 1, n + 1))
return np.array([estimator(int(n), int(c), k) for n, c in zip(num_samples, num_correct, strict=False)])
def compute_statistics(values: list[float]) -> dict[str, float]:
values = np.array(values)
return {
"mean": np.mean(values).item(),
"median": np.median(values).item(),
}
def compression_ratio(
data: str | bytes,
*,
encoding: str = "utf-8",
algorithm: Literal["zlib", "gzip", "bz2", "lzma"] = "zlib",
level: int = 9,
) -> tuple[float, float]:
if isinstance(data, str):
raw = data.encode(encoding)
else:
raw = data
original = len(raw)
if original == 0:
return float("inf"), 0.0
if algorithm == "zlib":
import zlib
compressed = zlib.compress(raw, level)
elif algorithm == "gzip":
import gzip
compressed = gzip.compress(raw, compresslevel=level)
elif algorithm == "bz2":
import bz2
compressed = bz2.compress(raw, compresslevel=level)
elif algorithm == "lzma":
import lzma
compressed = lzma.compress(raw, preset=level)
else:
raise ValueError(f"Unsupported algorithm: {algorithm}")
comp_len = len(compressed)
if comp_len == 0:
return float("inf"), 100.0
ratio = original / comp_len
savings_pct = 100.0 * (1.0 - comp_len / original)
return ratio, savings_pct
def has_repetition(text: str = None):
if len(text) > 10000 and compression_ratio(text[-10000:])[0] > 10:
return True
else:
return False
def compute_rollout_step(args, rollout_id):
if args.wandb_always_use_train_step:
return rollout_id * args.rollout_batch_size * args.n_samples_per_prompt // args.global_batch_size
return rollout_id

94
slime/utils/misc.py Normal file
View File

@@ -0,0 +1,94 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import importlib
import subprocess
import ray
from slime.utils.http_utils import is_port_available
def load_function(path):
"""
Load a function from a module.
:param path: The path to the function, e.g. "module.submodule.function".
:return: The function object.
"""
module_path, _, attr = path.rpartition(".")
module = importlib.import_module(module_path)
return getattr(module, attr)
class SingletonMeta(type):
"""
A metaclass for creating singleton classes.
"""
_instances = {}
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
instance = super().__call__(*args, **kwargs)
cls._instances[cls] = instance
return cls._instances[cls]
def exec_command(cmd: str, capture_output: bool = False) -> str | None:
print(f"EXEC: {cmd}", flush=True)
try:
result = subprocess.run(
["bash", "-c", cmd],
shell=False,
check=True,
capture_output=capture_output,
**(dict(text=True) if capture_output else {}),
)
except subprocess.CalledProcessError as e:
if capture_output:
print(f"{e.stdout=} {e.stderr=}")
raise
if capture_output:
print(f"Captured stdout={result.stdout} stderr={result.stderr}")
return result.stdout
def get_current_node_ip():
address = ray._private.services.get_node_ip_address()
# strip ipv6 address
address = address.strip("[]")
return address
def get_free_port(start_port=10000, consecutive=1):
# find the port where port, port + 1, port + 2, ... port + consecutive - 1 are all available
port = start_port
while not all(is_port_available(port + i) for i in range(consecutive)):
port += 1
return port
def should_run_periodic_action(
rollout_id: int,
interval: int | None,
num_rollout_per_epoch: int | None = None,
num_rollout: int | None = None,
) -> bool:
"""
Return True when a periodic action (eval/save/checkpoint) should run.
Args:
rollout_id: The current rollout index (0-based).
interval: Desired cadence; disables checks when None.
num_rollout_per_epoch: Optional epoch boundary to treat as a trigger.
"""
if interval is None:
return False
if num_rollout is not None and rollout_id == num_rollout - 1:
return True
step = rollout_id + 1
return (step % interval == 0) or (num_rollout_per_epoch is not None and step % num_rollout_per_epoch == 0)

718
slime/utils/ppo_utils.py Normal file
View File

@@ -0,0 +1,718 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/OpenRLHF/OpenRLHF/blob/10c733694ed9fbb78a0a2ff6a05efc7401584d46/openrlhf/models/utils.py
# and https://github.com/OpenRLHF/OpenRLHF/blob/10c733694ed9fbb78a0a2ff6a05efc7401584d46/openrlhf/trainer/ppo_utils/experience_maker.py
from argparse import Namespace
import torch
import torch.distributed as dist
import torch.nn.functional as F
@torch.compile(dynamic=True)
def compute_approx_kl(
log_probs: torch.Tensor,
log_probs_base: torch.Tensor,
kl_loss_type: str,
importance_ratio: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Compute the approximate KL divergence between two distributions.
Schulman blog: http://joschu.net/blog/kl-approx.html
Args:
log_probs: Log probabilities of the new distribution.
log_probs_base: Log probabilities of the base distribution.
kl_loss_type: Type of KL estimator (k1, k2, k3, low_var_kl).
importance_ratio: Optional IS ratio (π_θ/π_old) for unbiased KL estimation.
"""
log_ratio = log_probs.float() - log_probs_base.float()
if kl_loss_type == "k1":
kl = log_ratio
elif kl_loss_type == "k2":
kl = log_ratio**2 / 2.0
elif kl_loss_type in ["k3", "low_var_kl"]:
# The non negative kl approximation in
# http://joschu.net/blog/kl-approx.html
# Besides non negative, it is also unbiased and have lower variance.
log_ratio = -log_ratio
kl = log_ratio.exp() - 1 - log_ratio
else:
raise ValueError(f"Unknown kl_loss_type: {kl_loss_type}")
# Apply IS ratio for unbiased KL estimation (DeepSeek-V3.2)
if importance_ratio is not None:
kl = importance_ratio * kl
# Clamp only for low_var_kl for numerical stability
if kl_loss_type == "low_var_kl":
kl = torch.clamp(kl, min=-10, max=10)
return kl
def compute_opsm_mask(
args: Namespace,
full_log_probs: list[torch.Tensor],
full_old_log_probs: list[torch.Tensor],
advantages: list[torch.Tensor],
loss_masks: list[torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
"""Compute Off-Policy Sequence Masking (OPSM) mask.
Args:
args: Configuration containing `opsm_delta` threshold.
full_log_probs: Current policy log-probs per sample.
full_old_log_probs: Old policy log-probs per sample.
advantages: Advantage values per sample.
loss_masks: Loss masks per sample.
Returns:
Tuple of `(opsm_mask, opsm_clipfrac)` where `opsm_mask` is a
concatenated tensor of per-token masks and
`opsm_clipfrac` is the count of masked sequences.
"""
opsm_mask_list = []
device = advantages[0].device
opsm_clipfrac = torch.tensor(0.0, device=device)
for full_log_prob, full_old_log_prob, advantage, loss_mask in zip(
full_log_probs, full_old_log_probs, advantages, loss_masks, strict=False
):
# Calculate sequence-level KL
seq_kl = ((full_old_log_prob - full_log_prob) * loss_mask).sum() / torch.clamp_min(loss_mask.sum(), 1)
# Create mask: 0 if (advantage < 0 and seq_kl > delta), else 1
mask = ((advantage < 0) & (seq_kl > args.opsm_delta)).float()
opsm_clipfrac += mask.sum() / torch.clamp_min(loss_mask.sum(), 1)
opsm_mask_list.append(1 - mask)
opsm_mask = torch.cat(opsm_mask_list, dim=0)
return opsm_mask, opsm_clipfrac
def compute_gspo_kl(
full_log_probs: list[torch.Tensor],
full_old_log_probs: list[torch.Tensor],
local_log_probs: list[torch.Tensor],
loss_masks: list[torch.Tensor],
) -> torch.Tensor:
"""Compute GSPO-style per-sequence KL divergence.
Args:
full_log_probs: Current policy log-probs per sample (full or CP-local).
full_old_log_probs: Old policy log-probs per sample (full or CP-local).
local_log_probs: Local (CP-local) log-probs for expansion shape reference.
loss_masks: Loss masks per sample.
Returns:
Concatenated tensor of per-token KL values where each token in a
sequence has the same KL value (the sequence-level KL).
"""
# Compute sequence-level KL and expand to per-token
ppo_kl = [
((old_logprob - log_prob) * loss_mask).sum() / torch.clamp_min(loss_mask.sum(), 1)
for log_prob, old_logprob, loss_mask in zip(full_log_probs, full_old_log_probs, loss_masks, strict=False)
]
ppo_kl = [kl.expand_as(log_prob) for kl, log_prob in zip(ppo_kl, local_log_probs, strict=False)]
ppo_kl = torch.cat(ppo_kl, dim=0)
return ppo_kl
@torch.compile(dynamic=True)
def compute_policy_loss(
ppo_kl: torch.Tensor,
advantages: torch.Tensor,
eps_clip: float,
eps_clip_high: float,
eps_clip_c: float | None = None,
):
ratio = (-ppo_kl).exp()
pg_losses1 = -ratio * advantages
pg_losses2 = -ratio.clamp(1 - eps_clip, 1 + eps_clip_high) * advantages
clip_pg_losses1 = torch.maximum(pg_losses1, pg_losses2)
clipfrac = torch.gt(pg_losses2, pg_losses1).float()
if eps_clip_c is not None:
assert (
eps_clip_c > 1.0
), f"The lower bound of the clip_ratio_c for dual-clip PPO should be greater than 1.0, but get the value: {eps_clip_c}."
pg_losses3 = -eps_clip_c * advantages
clip_pg_losses2 = torch.min(pg_losses3, clip_pg_losses1)
pg_losses = torch.where(advantages < 0, clip_pg_losses2, clip_pg_losses1)
else:
pg_losses = clip_pg_losses1
return pg_losses, clipfrac
def compute_log_probs(logits: torch.Tensor, tokens: torch.Tensor, process_group: dist.ProcessGroup | None):
from megatron.core.fusions.fused_cross_entropy import fused_vocab_parallel_cross_entropy
# convert to [seq_len, batch_size, vocab_size] as expected by fused_vocab_parallel_cross_entropy
logits = logits.unsqueeze(1)
tokens = tokens.unsqueeze(1)
return -fused_vocab_parallel_cross_entropy(logits, tokens, process_group)
# from https://github.com/volcengine/verl/blob/0bdf7f469854815177e73dcfe9e420836c952e6e/verl/utils/megatron/tensor_parallel.py#L99
class _VocabParallelEntropy(torch.autograd.Function):
@staticmethod
def forward(ctx, vocab_parallel_logits: torch.Tensor, process_group: dist.ProcessGroup) -> torch.Tensor:
@torch.compile(dynamic=True)
def mul_reduce(a, b):
return (a * b).sum(dim=-1, keepdim=True)
logits_max = vocab_parallel_logits.max(dim=-1, keepdim=True).values
dist.all_reduce(logits_max, op=dist.ReduceOp.MAX, group=process_group)
normalized_vocab_parallel_logits = vocab_parallel_logits - logits_max
normalized_exp_logits = normalized_vocab_parallel_logits.exp_()
normalized_sum_exp_logits = normalized_exp_logits.sum(dim=-1, keepdim=True)
dist.all_reduce(normalized_sum_exp_logits, group=process_group)
softmax_logits = normalized_exp_logits.div_(normalized_sum_exp_logits)
sum_softmax_times_logits = mul_reduce(softmax_logits, vocab_parallel_logits)
dist.all_reduce(sum_softmax_times_logits, group=process_group)
entropy = logits_max + normalized_sum_exp_logits.log() - sum_softmax_times_logits
ctx.save_for_backward(vocab_parallel_logits, softmax_logits, sum_softmax_times_logits)
return entropy.squeeze(dim=-1)
@staticmethod
def backward(ctx, grad_output: torch.Tensor) -> torch.Tensor:
vocab_parallel_logits, softmax_logits, sum_softmax_times_logits = ctx.saved_tensors
# reuse softmax_logits as grad
vocab_parallel_logits.sub_(sum_softmax_times_logits)
softmax_logits.mul_(vocab_parallel_logits)
softmax_logits.mul_(grad_output.unsqueeze(dim=-1))
# recover vocab_parallel_logits
vocab_parallel_logits.add_(sum_softmax_times_logits)
softmax_logits.mul_(-1)
return softmax_logits, None
def compute_entropy_from_logits(logits: torch.Tensor, process_group) -> torch.Tensor:
return _VocabParallelEntropy.apply(logits, process_group)
def get_grpo_returns(
rewards: torch.Tensor,
kl: list[torch.Tensor],
):
returns = []
for i in range(len(rewards)):
returns.append(torch.ones_like(kl[i]) * rewards[i])
return returns
def get_reinforce_plus_plus_returns(
rewards: torch.Tensor,
kl: list[torch.Tensor],
loss_masks: list[torch.Tensor],
response_lengths: list[int],
total_lengths: list[int],
kl_coef: float,
gamma: float,
) -> list[torch.Tensor]:
"""
Calculates discounted returns for REINFORCE++ (https://arxiv.org/pdf/2501.03262)
Args:
rewards (Tensor): A tensor of scalar rewards for each sequence.
kl (List[Tensor]): List of per-token KL divergence tensors for sequence chunks.
loss_masks (List[Tensor]): List of response-only loss masks for each full sequence.
response_lengths (List[int]): The full length of each response sequence.
total_lengths (List[int]): The full length of each sequence (prompt + response).
kl_coef (float): Coefficient for the KL penalty.
gamma (float): The discount factor.
Returns:
List[torch.Tensor]: A list of return (G_t) tensors for the
local sequence chunks owned by the current GPU rank.
"""
from megatron.core import mpu
cp_size = mpu.get_context_parallel_world_size()
final_returns_chunks = []
for i in range(len(rewards)):
local_kl_chunk = kl[i]
total_len, response_len = total_lengths[i], response_lengths[i]
if cp_size > 1:
# Step 1,2:Gather all chunks and token_offsets from all ranks and reconstruct the full response tensor by splitting and placing each part
from slime.backends.megatron_utils.cp_utils import all_gather_with_cp
full_kl_response = all_gather_with_cp(local_kl_chunk, total_len, response_len)
else:
full_kl_response = local_kl_chunk
# Step 3: Compute returns on full response kl tensor.
token_level_rewards = -kl_coef * full_kl_response
full_mask = loss_masks[i]
assert full_mask.sum().item() > 0, f"Sequence at index {i} is fully masked."
last_idx = full_mask.nonzero(as_tuple=True)[0][-1]
token_level_rewards[last_idx] += rewards[i]
returns_for_seq = torch.zeros_like(token_level_rewards)
running_return = 0.0
for t in reversed(range(token_level_rewards.size(0))):
# G_t = r_t + gamma * G_{t+1}
running_return = token_level_rewards[t] + gamma * running_return
returns_for_seq[t] = running_return
# Step 4: Pick up the results corresponding to our local chunk's parts.
if cp_size > 1:
from slime.backends.megatron_utils.cp_utils import slice_log_prob_with_cp
local_returns_chunk = slice_log_prob_with_cp(returns_for_seq, total_len, response_len)
else:
local_returns_chunk = returns_for_seq
final_returns_chunks.append(local_returns_chunk)
return final_returns_chunks
def get_reinforce_plus_plus_baseline_advantages(
rewards: torch.Tensor,
kl: list[torch.Tensor],
loss_masks: list[torch.Tensor],
kl_coef: float,
) -> list[torch.Tensor]:
"""
Calculates the unwhitened advantages for the REINFORCE++-baseline algorithm.
Broadcasting the scalar (reward - group_baseline) to each token.
Args:
rewards (Tensor): A tensor of scalar rewards, where the group-wise
baseline has already been subtracted.
kl (list[Tensor]): A list of per-token KL divergence tensors. Used to
get the shape for broadcasting.
loss_masks (list[Tensor]): A list of per-token loss masks.
kl_coef (float): Coefficient for the KL penalty.
Returns:
list[Tensor]: A list of tensors containing the unwhitened advantages.
"""
# Broadcast to get unwhitened advantages
unwhitened_advantages = [
torch.ones_like(kl_tensor) * reward_val - kl_coef * kl_tensor
for kl_tensor, reward_val in zip(kl, rewards, strict=False)
]
return unwhitened_advantages
def get_advantages_and_returns(
total_len: int,
response_len: int,
values: torch.Tensor,
rewards: torch.Tensor,
gamma: float,
lambd: float,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Function that computes advantages and returns from rewards and values.
Calculated as in the original PPO paper: https://arxiv.org/abs/1707.06347
Note that rewards may include a KL divergence loss term.
Advantages looks like this:
Adv1 = R1 + γ * λ * R2 + γ^2 * λ^2 * R3 + ...
- V1 + γ * (1 - λ) V2 + γ^2 * λ * (1 - λ) V3 + ...
Returns looks like this:
Ret1 = R1 + γ * λ * R2 + γ^2 * λ^2 * R3 + ...
+ γ * (1 - λ) V2 + γ^2 * λ * (1 - λ) V3 + ...
Input:
- values: Tensor of shape (response_size,)
- rewards: Tensor of shape (response_size,)
Output:
- advantages: Tensor of shape (response_size,)
- returns: Tensor of shape (response_size,)
"""
from megatron.core import mpu
cp_size = mpu.get_context_parallel_world_size()
if cp_size > 1:
from slime.backends.megatron_utils.cp_utils import all_gather_with_cp
full_rewards = all_gather_with_cp(rewards, total_len, response_len)
full_values = all_gather_with_cp(values, total_len, response_len)
else:
full_rewards = rewards
full_values = values
lastgaelam = 0
advantages_reversed = []
for t in reversed(range(response_len)):
nextvalues = full_values[t + 1] if t < response_len - 1 else 0.0
delta = full_rewards[t] + gamma * nextvalues - full_values[t]
lastgaelam = delta + gamma * lambd * lastgaelam
advantages_reversed.append(lastgaelam)
full_advantages = torch.tensor(advantages_reversed[::-1], dtype=full_values.dtype, device=full_values.device)
full_returns = full_advantages + full_values
if cp_size > 1:
from slime.backends.megatron_utils.cp_utils import slice_log_prob_with_cp
advantages = slice_log_prob_with_cp(full_advantages, total_len, response_len)
returns = slice_log_prob_with_cp(full_returns, total_len, response_len)
else:
advantages = full_advantages
returns = full_returns
return advantages.detach(), returns
def get_advantages_and_returns_batch(
total_lengths,
response_lengths,
values_list,
rewards_list,
gamma,
lambd,
chunked: bool = True,
):
"""
Batched GAE with CP support.
Input:
total_lengths: list[int], each sample's total_len
response_lengths: list[int], each sample's response_len
values_list: list[Tensor], each shape = [resp_len_i]
rewards_list: list[Tensor], same shape
Output:
advantages_list: list[Tensor], each shape = [resp_len_i]
returns_list: list[Tensor], same shape
"""
from megatron.core import mpu
with torch.no_grad():
B = len(response_lengths)
assert B == len(values_list)
assert B == len(rewards_list)
cp_size = mpu.get_context_parallel_world_size()
device = values_list[0].device
dtype = values_list[0].dtype
if cp_size > 1:
from slime.backends.megatron_utils.cp_utils import all_gather_with_cp
full_values_list = []
full_rewards_list = []
for total_len, resp_len, v, r in zip(
total_lengths, response_lengths, values_list, rewards_list, strict=False
):
full_v = all_gather_with_cp(v, total_len, resp_len)
full_r = all_gather_with_cp(r, total_len, resp_len)
full_values_list.append(full_v)
full_rewards_list.append(full_r)
# full_values_list[i].shape = [total_len_i]
else:
full_values_list = values_list
full_rewards_list = rewards_list
# pad to max_len for batched GAE
max_len = max(response_lengths)
full_values = torch.zeros(B, max_len, device=device, dtype=dtype)
full_rewards = torch.zeros(B, max_len, device=device, dtype=dtype)
for i in range(B):
L = response_lengths[i]
full_values[i, :L] = full_values_list[i][:L]
full_rewards[i, :L] = full_rewards_list[i][:L]
if not chunked:
full_advantages, full_returns = vanilla_gae(
rewards=full_rewards,
values=full_values,
gamma=gamma,
lambd=lambd,
)
else:
full_advantages, full_returns = chunked_gae(
rewards=full_rewards,
values=full_values,
gamma=gamma,
lambd=lambd,
)
advantages_list = []
returns_list = []
if cp_size > 1:
from slime.backends.megatron_utils.cp_utils import slice_log_prob_with_cp
for total_len, resp_len, adv_row, ret_row in zip(
total_lengths,
response_lengths,
full_advantages,
full_returns,
strict=False,
):
adv_full = adv_row # shape = [resp_len_i padded to max_len]
ret_full = ret_row
adv_sliced = slice_log_prob_with_cp(adv_full[:resp_len], total_len, resp_len)
ret_sliced = slice_log_prob_with_cp(ret_full[:resp_len], total_len, resp_len)
advantages_list.append(adv_sliced)
returns_list.append(ret_sliced)
else:
for i in range(B):
L = response_lengths[i]
advantages_list.append(full_advantages[i, :L])
returns_list.append(full_returns[i, :L])
return advantages_list, returns_list
def vanilla_gae(
rewards: torch.Tensor,
values: torch.Tensor,
gamma: float,
lambd: float,
):
B, T = rewards.shape
device = rewards.device
dtype = rewards.dtype
lastgaelam = torch.zeros(B, device=device, dtype=dtype)
adv_rev = []
for t in reversed(range(T)):
next_value = values[:, t + 1] if t < T - 1 else 0.0
delta = rewards[:, t] + gamma * next_value - values[:, t]
lastgaelam = delta + gamma * lambd * lastgaelam
adv_rev.append(lastgaelam)
full_advantages = torch.stack(adv_rev[::-1], dim=1) # [B, max_len]
full_returns = full_advantages + values # [B, max_len]
return full_advantages, full_returns
def chunked_gae(
rewards: torch.Tensor,
values: torch.Tensor,
gamma: float,
lambd: float,
chunk_size: int = 128,
):
"""
Compute Generalized Advantage Estimation (GAE) using a FlashLinearAttention-
inspired algorithm: parallel prefix scan within chunks and recurrent state
propagation across chunks.
This reduces the sequential dependency length from O(T) to O(T / chunk_size),
while keeping chunk computations fully parallelizable (O(C^2) per chunk).
Args:
rewards (Tensor): [B, T] reward sequence.
values (Tensor): [B, T] value predictions. The next-value of the final
step is assumed to be zero (standard PPO convention).
gamma (float): discount factor.
lam (float): GAE lambda.
chunk_size (int): sequence chunk length for parallel scan.
Returns:
advantages (Tensor): [B, T] computed advantages.
returns (Tensor): [B, T] advantages + values.
"""
# -------------------------------------------------------------------------
# Validate inputs
# -------------------------------------------------------------------------
assert rewards.ndim == 2 and values.ndim == 2
B, T = rewards.shape
assert values.shape == (B, T)
device = rewards.device
dtype = rewards.dtype
# -------------------------------------------------------------------------
# Build δ_t = r_t + γ * V_{t+1} - V_t with V_{T} = 0
# -------------------------------------------------------------------------
next_values = torch.cat(
[values[:, 1:], torch.zeros(B, 1, device=device, dtype=dtype)],
dim=1,
)
deltas = rewards + gamma * next_values - values
# Reformulate backward GAE as a forward scan on the reversed sequence:
# S[i] = Δ[i] + w * S[i - 1], w = γλ
w = gamma * lambd
deltas_rev = torch.flip(deltas, dims=[1]) # [B, T]
# -------------------------------------------------------------------------
# Pad to a multiple of chunk_size
# -------------------------------------------------------------------------
if T % chunk_size != 0:
pad = chunk_size - (T % chunk_size)
deltas_rev = F.pad(deltas_rev, (0, pad))
else:
pad = 0
B, T_pad = deltas_rev.shape
n_chunks = T_pad // chunk_size
deltas_chunks = deltas_rev.view(B, n_chunks, chunk_size)
# -------------------------------------------------------------------------
# Construct the intra-chunk parallel scan kernel M
#
# For a chunk Δ[0..C-1], we want:
# S_local[t] = sum_{k=0..t} w^(t-k) * Δ[k]
#
# This is implemented as:
# S_local = Δ @ M
#
# where:
# M[i, j] = w^(j - i) if j >= i
# 0 otherwise
# -------------------------------------------------------------------------
idx = torch.arange(chunk_size, device=device)
row = idx[:, None]
col = idx[None, :]
diff = col - row
M = torch.zeros(chunk_size, chunk_size, device=device, dtype=dtype)
mask = diff >= 0
if w == 0.0:
M[mask & (diff == 0)] = 1.0
else:
M[mask] = w ** diff[mask].to(dtype)
# pow_vec[t] = w^(t+1), used to inject the recurrent state s_prev
if w == 0.0:
pow_vec = torch.zeros(chunk_size, device=device, dtype=dtype)
else:
pow_vec = w ** torch.arange(1, chunk_size + 1, device=device, dtype=dtype)
# -------------------------------------------------------------------------
# Parallel compute local chunk results (assuming initial state = 0)
# -------------------------------------------------------------------------
deltas_flat = deltas_chunks.reshape(B * n_chunks, chunk_size)
S_local_flat = deltas_flat @ M
S_local_chunks = S_local_flat.view(B, n_chunks, chunk_size)
# Effective length of each chunk (the last chunk may be padded)
lengths = [chunk_size] * n_chunks
if pad > 0:
lengths[-1] = chunk_size - pad
# -------------------------------------------------------------------------
# Recurrent propagation between chunks
#
# Each chunk contributes:
# S_global[t] = S_local[t] + w^(t+1) * s_prev
#
# And updates:
# s_prev = S_global[last_t]
# -------------------------------------------------------------------------
S_rev = deltas_rev.new_zeros(B, T_pad)
s_prev = torch.zeros(B, device=device, dtype=dtype)
for c in range(n_chunks):
Lc = lengths[c]
start = c * chunk_size
end = start + Lc
S_local = S_local_chunks[:, c, :Lc]
S_global = S_local + s_prev.unsqueeze(1) * pow_vec[:Lc]
S_rev[:, start:end] = S_global
s_prev = S_global[:, -1] # state for next chunk
# Remove padding and flip back to original time order
if pad > 0:
S_rev = S_rev[:, :T]
advantages = torch.flip(S_rev, dims=[1])
returns = advantages + values
return advantages, returns
def calculate_log_probs_and_entropy(logits, tokens, tp_group, with_entropy: bool = False, chunk_size: int = -1):
logits = logits.contiguous()
# TODO: not sure why we need to clone the logits here.
# Without the clone, the backward will trigger inplace edit error.
# It seems that the function with tp will modify the logits inplace.
entropy = None
if logits.size(0) != 0:
if chunk_size > 0:
num_chunks = (logits.size(0) - 1) // chunk_size + 1
tokens_chunks = tokens.chunk(num_chunks, dim=0)
logits_chunks = logits.chunk(num_chunks, dim=0)
log_probs = []
for tokens_chunk, logits_chunk in zip(tokens_chunks, logits_chunks, strict=True):
log_prob = compute_log_probs(logits_chunk.clone(), tokens_chunk, tp_group)
log_probs.append(log_prob)
log_prob = torch.cat(log_probs, dim=0)
if with_entropy:
entropys = []
for _, logits_chunk in zip(tokens_chunks, logits_chunks, strict=True):
entropy = compute_entropy_from_logits(logits_chunk.clone(), tp_group)
entropys.append(entropy)
entropy = torch.cat(entropys, dim=0)
else:
log_prob = compute_log_probs(logits.clone(), tokens, tp_group)
if with_entropy:
entropy = compute_entropy_from_logits(logits.clone(), tp_group)
else:
log_prob = logits.new_zeros((0,))
if with_entropy:
entropy = logits.new_zeros((0,))
return log_prob, entropy
def vanilla_tis_function(
args,
*,
pg_loss: torch.Tensor,
train_log_probs: list[torch.Tensor],
rollout_log_probs: list[torch.Tensor],
loss_masks: list[torch.Tensor],
**kwargs,
) -> tuple[torch.Tensor, list[torch.Tensor], dict[str, torch.Tensor]]:
"""Apply TIS off-policy correction using importance sampling.
Parameters:
args: Arguments containing TIS settings.
pg_loss: Policy gradient loss tensor of shape [total_seq_len - 1].
train_log_probs: List of tensors containing training log-probabilities
for each sequence.
rollout_log_probs: List of tensors containing rollout log-probabilities
for each sequence.
loss_masks: List of tensors containing loss masks for each sequence.
"""
rollout_log_probs = torch.cat(rollout_log_probs, dim=0)
old_log_probs = torch.cat(train_log_probs, dim=0)
tis = torch.exp(old_log_probs - rollout_log_probs)
tis_abs = (tis - 1).abs()
tis_clip_low = args.tis_clip_low if args.tis_clip_low is not None else 0.1
tis_clip_high = args.tis_clip if args.tis_clip is not None else 2.0
tis_weights = torch.clamp(tis, min=tis_clip_low, max=tis_clip_high)
tis_clipfrac = (tis_weights != tis).float()
metrics = {
"tis": tis.clone().detach(),
"tis_clipfrac": tis_clipfrac.clone().detach(),
"tis_abs": tis_abs.clone().detach(),
}
pg_loss = pg_loss * tis_weights
return pg_loss, loss_masks, metrics

View File

@@ -0,0 +1,37 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import base64
import io
import logging
from transformers import AutoProcessor, AutoTokenizer, PreTrainedTokenizerBase, ProcessorMixin
logger = logging.getLogger(__name__)
def load_tokenizer(name_or_path: str, **kwargs):
return AutoTokenizer.from_pretrained(name_or_path, **kwargs)
def load_processor(name_or_path: str, **kwargs):
try:
proc = AutoProcessor.from_pretrained(name_or_path, **kwargs)
except (OSError, ValueError) as e:
logger.warning(f"Failed to load processor from {name_or_path}: {e}")
proc = None
# If HF returned a tokenizer, discard it.
if isinstance(proc, PreTrainedTokenizerBase) or not isinstance(proc, ProcessorMixin):
proc = None
return proc
def encode_image_for_rollout_engine(image) -> str:
"""Load an image from path, ensure RGB, encode as PNG base64 string."""
buffer = io.BytesIO()
if image.mode != "RGB":
image = image.convert("RGB")
image.save(buffer, format="PNG")
return base64.b64encode(buffer.getvalue()).decode("utf-8")

View File

@@ -0,0 +1,150 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import time
import traceback
from pathlib import Path
import torch
from slime.utils.memory_utils import print_memory
logger = logging.getLogger(__name__)
class TrainProfiler:
def __init__(self, args):
self.args = args
self._torch_profiler_overall = None
self._memory_profiler_overall = None
if args.use_pytorch_profiler and ("train_overall" in args.profile_target):
self._torch_profiler_overall = _create_torch_profiler(args, name="train_overall")
if args.record_memory_history and ("train_overall" in args.profile_target):
self._memory_profiler_overall = _BaseMemoryProfiler.create(args)
self._memory_profiler_overall.start()
def on_init_end(self):
if self._torch_profiler_overall is not None:
self._torch_profiler_overall.start()
def step(self, rollout_id: int):
if self._torch_profiler_overall is not None:
self._torch_profiler_overall.step()
if (
self._memory_profiler_overall is not None
and ((s := self.args.memory_snapshot_num_steps) is not None)
and (rollout_id == s - 1)
):
self._memory_profiler_overall.stop()
def iterate_train_actor(self, iterator):
return _profile_simple_loop(iterator, self.args, name="train_actor")
def iterate_train_log_probs(self, iterator):
return _profile_simple_loop(iterator, self.args, name="train_log_probs")
def _profile_simple_loop(iterator, args, name):
if not (args.use_pytorch_profiler and (name in args.profile_target)):
yield from iterator
return
torch_profiler = _create_torch_profiler(args, name=name)
torch_profiler.start()
for item in iterator:
yield item
torch_profiler.step()
def _create_torch_profiler(args, name):
return torch.profiler.profile(
schedule=torch.profiler.schedule(
# TODO the train_actor and train_log_probs ones may need to have different args to control step
wait=max(args.profile_step_start - 1, 0),
warmup=1 if args.profile_step_start > 0 else 0,
active=args.profile_step_end - args.profile_step_start,
repeat=1,
),
on_trace_ready=torch.profiler.tensorboard_trace_handler(
args.tensorboard_dir,
worker_name=f"{name}_rank_{torch.distributed.get_rank()}",
use_gzip=True,
),
record_shapes=True,
with_stack=True,
profile_memory=True,
with_flops=True,
)
class _BaseMemoryProfiler:
@staticmethod
def create(args):
c = {
"torch": _TorchMemoryProfiler,
"memray": _MemrayMemoryProfiler,
}[args.memory_recorder]
return c(args)
def __init__(self, args):
self._path_dump = (
Path(args.memory_snapshot_dir)
/ f"memory_snapshot_time{time.time()}_rank{torch.distributed.get_rank()}_{args.memory_snapshot_path}"
)
def start(self):
raise NotImplementedError
def stop(self):
raise NotImplementedError
class _TorchMemoryProfiler(_BaseMemoryProfiler):
def start(self):
logger.info("Attach OOM dump memory history.")
torch.cuda.memory._record_memory_history(
max_entries=1000000,
# record stack information for the trace events
# trace_alloc_record_context=True,
stacks="all",
)
def oom_observer(device, alloc, device_alloc, device_free):
logger.info(
f"Observe OOM, will dump snapshot to {self._path_dump}. ({device=} {alloc=} {device_alloc=} {device_free=}; stacktrace is as follows)"
)
traceback.print_stack()
torch.cuda.memory._dump_snapshot(self._path_dump)
print_memory("when oom")
torch._C._cuda_attach_out_of_memory_observer(oom_observer)
def stop(self):
logger.info(f"Dump memory snapshot to: {self._path_dump}")
torch.cuda.memory._dump_snapshot(self._path_dump)
torch.cuda.memory._record_memory_history(enabled=None)
class _MemrayMemoryProfiler(_BaseMemoryProfiler):
def __init__(self, args):
super().__init__(args)
assert args.memory_snapshot_num_steps is not None, "In memray, must provide --memory-snapshot-num-steps"
def start(self):
logger.info("Memray tracker started.")
import memray
self._tracker = memray.Tracker(
file_name=self._path_dump,
native_traces=True,
)
self._tracker.__enter__()
def stop(self):
logger.info(f"Memray tracker stopped and dump snapshot to: {self._path_dump}")
self._tracker.__exit__(None, None, None)

10
slime/utils/ray_utils.py Normal file
View File

@@ -0,0 +1,10 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
class Box:
def __init__(self, inner):
self._inner = inner
@property
def inner(self):
return self._inner

View File

@@ -0,0 +1,285 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import os
from contextlib import contextmanager
import torch
import torch.distributed as dist
from slime.utils.memory_utils import print_memory
logger = logging.getLogger(__name__)
old_new_group_dict = {}
def monkey_patch_torch_dist():
pid = os.getpid()
if pid in old_new_group_dict:
assert dist.old_new_group == old_new_group_dict[pid]
return
logger.info("Applying monkey patch to torch.distributed")
old_new_group = dist.new_group
old_new_group_dict[pid] = old_new_group
dist.old_new_group = old_new_group
def new_group(*args, **kwargs):
group = old_new_group(*args, **kwargs)
# skip none nccl group.
if len(args) >= 3 and args[2] == "gloo" or "backend" in kwargs and kwargs["backend"] == "gloo":
return group
# Get ranks from arguments
if len(args) >= 1 and args[0] is not None:
ranks = args[0]
elif "ranks" in kwargs and kwargs["ranks"] is not None:
ranks = kwargs["ranks"]
else:
# If no ranks specified, use all ranks in world
ranks = list(range(dist.get_world_size()))
if len(ranks) == 1:
return group
group = ReloadableProcessGroup(group, ranks)
return group
dist.new_group = new_group
def get_new_function(func):
def new_function(*args, **kwargs):
args = tuple([arg.group if isinstance(arg, ReloadableProcessGroup) else arg for arg in args])
kwargs = {k: (v.group if isinstance(v, ReloadableProcessGroup) else v) for k, v in kwargs.items()}
with _wrap_low_level_call():
return func(*args, **kwargs)
return new_function
dist.get_rank = get_new_function(dist.get_rank)
dist.get_world_size = get_new_function(dist.get_world_size)
dist.get_backend = get_new_function(dist.get_backend)
dist.get_global_rank = get_new_function(dist.get_global_rank)
dist.get_group_rank = get_new_function(dist.get_group_rank)
dist.get_process_group_ranks = get_new_function(dist.get_process_group_ranks)
dist.all_reduce = get_new_function(dist.all_reduce)
dist.all_gather = get_new_function(dist.all_gather)
dist.all_gather_into_tensor = get_new_function(dist.all_gather_into_tensor)
dist.all_gather_object = get_new_function(dist.all_gather_object)
dist.all_to_all = get_new_function(dist.all_to_all)
dist.all_to_all_single = get_new_function(dist.all_to_all_single)
dist.broadcast = get_new_function(dist.broadcast)
dist.reduce = get_new_function(dist.reduce)
dist.reduce_scatter = get_new_function(dist.reduce_scatter)
dist.reduce_scatter_tensor = get_new_function(dist.reduce_scatter_tensor)
dist.scatter = get_new_function(dist.scatter)
dist.gather = get_new_function(dist.gather)
dist.barrier = get_new_function(dist.barrier)
dist.send = get_new_function(dist.send)
dist.recv = get_new_function(dist.recv)
dist._coalescing_manager = get_new_function(dist._coalescing_manager)
# p2p
old_isend = dist.isend
old_irecv = dist.irecv
dist.isend = get_new_function(dist.isend)
dist.irecv = get_new_function(dist.irecv)
def get_new_p2pop_function(func):
def new_function(*args, **kwargs):
def convert(arg):
if isinstance(arg, ReloadableProcessGroup):
return arg.group
elif arg == dist.isend:
arg = old_isend
elif arg == dist.irecv:
arg = old_irecv
return arg
args = (convert(arg) for arg in args)
kwargs = {k: convert(v) for k, v in kwargs.items()}
return func(*args, **kwargs)
return new_function
dist.P2POp.__new__ = get_new_p2pop_function(dist.P2POp.__new__)
dist.P2POp.__init__ = get_new_p2pop_function(dist.P2POp.__init__)
class ReloadableProcessGroup(torch.distributed.ProcessGroup):
GROUPS = {}
def __init__(self, group, ranks):
super().__init__(
rank=dist.get_rank(group),
size=dist.get_world_size(group),
)
self.group = group
self.group_info = {
"ranks": ranks,
}
pid = os.getpid()
if pid not in ReloadableProcessGroup.GROUPS:
ReloadableProcessGroup.GROUPS[pid] = []
ReloadableProcessGroup.GROUPS[pid].append(self)
def __getattr__(self, name):
return getattr(self.group, name)
@staticmethod
def destroy_process_groups():
pid = os.getpid()
for reloadable_group in ReloadableProcessGroup.GROUPS.get(pid, []):
if reloadable_group.group is None:
continue
try:
dist.destroy_process_group(reloadable_group.group)
except ValueError as e:
logger.warning(
f"Process group already invalid/destroyed; skipping cleanup. Exception: {e}",
exc_info=True,
)
del reloadable_group.group
reloadable_group.group = None
@staticmethod
def reload_process_groups():
pid = os.getpid()
reloadable_groups = ReloadableProcessGroup.GROUPS.get(pid, [])
logger.info(f"Reloading {len(reloadable_groups)} process groups in pid {pid}")
old_new_group = old_new_group_dict.get(pid)
for reloadable_group in reloadable_groups:
if reloadable_group.group is not None:
continue
group = old_new_group(ranks=reloadable_group.group_info["ranks"], backend="nccl")
reloadable_group.group = group
def rank(self) -> int:
return self.group.rank()
def size(self) -> int:
return self.group.size()
def name(self) -> str:
return self.group.name()
def shutdown(self) -> None:
if self.group is not None:
self.group.shutdown()
def abort(self) -> None:
if self.group is not None:
self.group.abort()
def _fwd(self, method, *args, **kwargs):
inner = self.group
if inner is None:
raise RuntimeError("ReloadableProcessGroup: inner PG is None, call reload() first.")
with _wrap_low_level_call():
return getattr(inner, method)(*args, **kwargs)
def barrier(self, *a, **kw):
return self._fwd("barrier", *a, **kw)
def broadcast(self, *a, **kw):
return self._fwd("broadcast", *a, **kw)
def allreduce(self, *a, **kw):
return self._fwd("allreduce", *a, **kw)
def allreduce_coalesced(self, *a, **kw):
return self._fwd("allreduce_coalesced", *a, **kw)
def reduce(self, *a, **kw):
return self._fwd("reduce", *a, **kw)
def allgather(self, *a, **kw):
return self._fwd("allgather", *a, **kw)
def _allgather_base(self, *a, **kw):
return self._fwd("_allgather_base", *a, **kw)
def allgather_coalesced(self, *a, **kw):
return self._fwd("allgather_coalesced", *a, **kw)
def allgather_into_tensor_coalesced(self, *a, **kw):
return self._fwd("allgather_into_tensor_coalesced", *a, **kw)
def gather(self, *a, **kw):
return self._fwd("gather", *a, **kw)
def scatter(self, *a, **kw):
return self._fwd("scatter", *a, **kw)
def reduce_scatter(self, *a, **kw):
return self._fwd("reduce_scatter", *a, **kw)
def _reduce_scatter_base(self, *a, **kw):
return self._fwd("_reduce_scatter_base", *a, **kw)
def reduce_scatter_tensor_coalesced(self, *a, **kw):
return self._fwd("reduce_scatter_tensor_coalesced", *a, **kw)
def alltoall_base(self, *a, **kw):
return self._fwd("alltoall_base", *a, **kw)
def alltoall(self, *a, **kw):
return self._fwd("alltoall", *a, **kw)
def send(self, *a, **kw):
return self._fwd("send", *a, **kw)
def recv(self, *a, **kw):
return self._fwd("recv", *a, **kw)
def recv_anysource(self, *a, **kw):
return self._fwd("recv_anysource", *a, **kw)
def _start_coalescing(self, *a, **kw):
return self._fwd("_start_coalescing", *a, **kw)
def _end_coalescing(self, *a, **kw):
return self._fwd("_end_coalescing", *a, **kw)
def _get_backend_name(self):
return self._fwd("_get_backend_name")
def _get_backend(self, *a, **kw):
return self._fwd("_get_backend", *a, **kw)
def _set_default_backend(self, *a, **kw):
return self._fwd("_set_default_backend", *a, **kw)
@property
def bound_device_id(self):
return self.group.bound_device_id
@bound_device_id.setter
def bound_device_id(self, dev):
self.group.bound_device_id = dev
def destroy_process_groups():
"""Destroy all reloadable process groups."""
ReloadableProcessGroup.destroy_process_groups()
def reload_process_groups():
"""Reload all reloadable process groups."""
ReloadableProcessGroup.reload_process_groups()
@contextmanager
def _wrap_low_level_call():
try:
yield
except Exception as e:
mem_info = print_memory("after torch distributed error")
e.add_note(f"{mem_info=}")
raise

View File

@@ -0,0 +1,30 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
from megatron.core.dist_checkpointing.strategies.filesystem_async import FileSystemWriterAsync
class ROCmFileSystemWriterAsync(FileSystemWriterAsync):
"""
FileSystemWriterAsync wrapper for ROCm compatibility.
On ROCm/HIP, using non_blocking=True causes tensors to be stored in pinned memory,
which triggers segmentation faults when forking subprocesses afterward.
"""
@staticmethod
def preload_tensors(*args, **kwargs):
# Change argument non_blocking to False on HIP platform
# The tensors will be stored in pinned memory if non_blocking=True
# Currently on the ROCm platform, forking a subprocess afterward
# with pinned_memory=True will trigger segmentation fault
if torch.version.hip:
print("HIP/ROCm detected: setting non_blocking=False in preload_tensors")
if "non_blocking" in kwargs:
kwargs["non_blocking"] = False
elif len(args) > 1 and isinstance(args[-1], bool):
# non_blocking is typically the last argument
args = args[:-1] + (False,)
return FileSystemWriterAsync.preload_tensors(*args, **kwargs)

View File

@@ -0,0 +1,95 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import os
import torch
ROUTING_REPLAY = None
def set_routing_replay(replay):
global ROUTING_REPLAY
ROUTING_REPLAY = replay
class RoutingReplay:
all_routing_replays = []
def __init__(self):
self.forward_index = 0
self.backward_index = 0
self.top_indices_list = []
RoutingReplay.all_routing_replays.append(self)
def record(self, top_indices):
# offload top_indices to CPU pinned memory
buf = torch.empty_like(top_indices, device="cpu", pin_memory=True)
buf.copy_(top_indices)
self.top_indices_list.append(buf)
def pop_forward(self):
top_indices = self.top_indices_list[self.forward_index]
self.forward_index += 1
return top_indices.to(torch.cuda.current_device())
def pop_backward(self):
top_indices = self.top_indices_list[self.backward_index]
self.backward_index += 1
return top_indices.to(torch.cuda.current_device())
def clear(self):
self.forward_index = 0
self.backward_index = 0
self.top_indices_list = []
def clear_forward(self):
self.forward_index = 0
@staticmethod
def clear_all():
for replay in RoutingReplay.all_routing_replays:
replay.clear()
@staticmethod
def clear_all_forward():
for replay in RoutingReplay.all_routing_replays:
replay.clear_forward()
def get_routing_replay_compute_topk(old_compute_topk):
def compute_topk(scores, topk, num_groups=None, group_topk=None):
if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1":
routing_replay_stage = os.environ["ROUTING_REPLAY_STAGE"]
if routing_replay_stage == "fallthrough":
return old_compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)
if routing_replay_stage == "record":
probs, top_indices = old_compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)
ROUTING_REPLAY.record(top_indices)
elif routing_replay_stage == "replay_forward":
top_indices = ROUTING_REPLAY.pop_forward()
assert (
top_indices.shape[0] == scores.shape[0] and top_indices.shape[1] == topk
), f"[{torch.distributed.get_rank()}] top_indices shape {top_indices.shape} does not match scores shape {scores.shape} and topk {topk}"
probs = scores.gather(1, top_indices)
elif routing_replay_stage == "replay_backward":
top_indices = ROUTING_REPLAY.pop_backward()
assert (
top_indices.shape[0] == scores.shape[0] and top_indices.shape[1] == topk
), f"top_indices shape {top_indices.shape} does not match scores shape {scores.shape} and topk {topk}"
probs = scores.gather(1, top_indices)
return probs, top_indices
else:
return old_compute_topk(scores, topk, num_groups=num_groups, group_topk=group_topk)
return compute_topk
def register_routing_replay(module):
if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1":
module.routing_replay = RoutingReplay()
def pre_forward_hook(*args, **kwargs):
set_routing_replay(module.routing_replay)
module.register_forward_pre_hook(pre_forward_hook)

View File

@@ -0,0 +1,186 @@
# Copied from https://github.com/volcengine/verl/blob/468adf22c43b744348051fccd7a5d830c6c3c36a/verl/utils/seqlen_balancing.py
# Copyright 2024 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import copy
import heapq
def karmarkar_karp(seqlen_list: list[int], k_partitions: int, equal_size: bool):
# see: https://en.wikipedia.org/wiki/Largest_differencing_method
class Set:
def __init__(self) -> None:
self.sum = 0
self.items = []
def add(self, idx: int, val: int):
self.items.append((idx, val))
self.sum += val
def merge(self, other):
for idx, val in other.items:
self.items.append((idx, val))
self.sum += val
def __lt__(self, other):
if self.sum != other.sum:
return self.sum < other.sum
if len(self.items) != len(other.items):
return len(self.items) < len(other.items)
return self.items < other.items
class State:
def __init__(self, items: list[tuple[int, int]], k: int) -> None:
self.k = k
# sets should always be decreasing order
self.sets = [Set() for _ in range(k)]
assert len(items) in [1, k], f"{len(items)} not in [1, {k}]"
for i, (idx, seqlen) in enumerate(items):
self.sets[i].add(idx=idx, val=seqlen)
self.sets = sorted(self.sets, reverse=True)
def get_partitions(self):
partitions = []
for i in range(len(self.sets)):
cur_partition = []
for idx, _ in self.sets[i].items:
cur_partition.append(idx)
partitions.append(cur_partition)
return partitions
def merge(self, other):
for i in range(self.k):
self.sets[i].merge(other.sets[self.k - 1 - i])
self.sets = sorted(self.sets, reverse=True)
@property
def spread(self) -> int:
return self.sets[0].sum - self.sets[-1].sum
def __lt__(self, other):
# least heap, let the state with largest spread to be popped first,
# if the spread is the same, let the state who has the largest set
# to be popped first.
if self.spread != other.spread:
return self.spread > other.spread
return self.sets[0] > other.sets[0]
def __repr__(self) -> str:
repr_str = "["
for i in range(self.k):
if i > 0:
repr_str += ","
repr_str += "{"
for j, (_, seqlen) in enumerate(self.sets[i].items):
if j > 0:
repr_str += ","
repr_str += str(seqlen)
repr_str += "}"
repr_str += "]"
return repr_str
sorted_seqlen_list = sorted([(seqlen, i) for i, seqlen in enumerate(seqlen_list)])
states_pq = []
if equal_size:
assert len(seqlen_list) % k_partitions == 0, f"{len(seqlen_list)} % {k_partitions} != 0"
for offset in range(0, len(sorted_seqlen_list), k_partitions):
items = []
for i in range(k_partitions):
seqlen, idx = sorted_seqlen_list[offset + i]
items.append((idx, seqlen))
heapq.heappush(states_pq, State(items=items, k=k_partitions))
else:
for seqlen, idx in sorted_seqlen_list:
heapq.heappush(states_pq, State(items=[(idx, seqlen)], k=k_partitions))
while len(states_pq) > 1:
state0 = heapq.heappop(states_pq)
state1 = heapq.heappop(states_pq)
# merge states
state0.merge(state1)
heapq.heappush(states_pq, state0)
final_state = states_pq[0]
partitions = final_state.get_partitions()
if equal_size:
for _i, partition in enumerate(partitions):
assert len(partition) * k_partitions == len(
seqlen_list
), f"{len(partition)} * {k_partitions} != {len(seqlen_list)}"
return partitions
def greedy_partition(seqlen_list: list[int], k_partitions: int, equal_size: bool):
bias = sum(seqlen_list) + 1 if equal_size else 0
sorted_seqlen = [(seqlen + bias, i) for i, seqlen in enumerate(seqlen_list)]
partitions = [[] for _ in range(k_partitions)]
partition_sums = [0 for _ in range(k_partitions)]
for seqlen, i in sorted_seqlen:
min_idx = None
for j in range(k_partitions):
if min_idx is None or partition_sums[j] < partition_sums[min_idx]:
min_idx = j
partitions[min_idx].append(i)
partition_sums[min_idx] += seqlen
if equal_size:
for _i, partition in enumerate(partitions):
assert len(partition) * k_partitions == len(
seqlen_list
), f"{len(partition)} * {k_partitions} != {len(seqlen_list)}"
return partitions
def get_seqlen_balanced_partitions(seqlen_list: list[int], k_partitions: int, equal_size: bool):
"""get order of seq lengths to make partitions balanced, this is
used in balacing sum of seqlength across dp ranks and microbatches
Parameters:
seqlen_list (List[int]):
seq lengths of each items
k_partitions (int):
resulting number of partitions
equal_size (bool):
if True, number of items in each partitions must be equal.
if False, only consider balancing the sum, each partition can have
variable number of items
Returns:
partitions (List[List[int]]):
return k_partitions list containing the index of items.
"""
assert len(seqlen_list) >= k_partitions, f"number of items:[{len(seqlen_list)}] < k_partitions:[{k_partitions}]"
def _check_and_sort_partitions(partitions):
assert len(partitions) == k_partitions, f"{len(partitions)} != {k_partitions}"
seen_idx = set()
sorted_partitions = [None] * k_partitions
for _i, partition in enumerate(partitions):
assert len(partition) > 0, f"the {_i}-th partition is empty"
for idx in partition:
seen_idx.add(idx)
sorted_partitions[_i] = sorted(partition)
assert seen_idx == set(range(len(seqlen_list)))
return sorted_partitions
partitions = karmarkar_karp(seqlen_list=seqlen_list, k_partitions=k_partitions, equal_size=equal_size)
return _check_and_sort_partitions(partitions)
def get_reverse_idx(idx_map):
reverse_idx_map = copy.deepcopy(idx_map)
for i, idx in enumerate(idx_map):
reverse_idx_map[idx] = i
return reverse_idx_map

View File

@@ -0,0 +1,118 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from collections import defaultdict
from collections.abc import Callable, Iterable
import torch
_SourceGetter = Callable[[], Iterable[tuple[str, torch.Tensor]]]
class TensorBackuper(ABC):
@staticmethod
def create(source_getter, single_tag):
if single_tag is None:
return _TensorBackuperNormal(source_getter=source_getter)
else:
return _TensorBackuperNoop(source_getter=source_getter, single_tag=single_tag)
def __init__(self, source_getter: _SourceGetter):
self._source_getter = source_getter
@property
@abstractmethod
def backup_tags(self):
raise NotImplementedError
@abstractmethod
def get(self, tag: str):
raise NotImplementedError
@abstractmethod
def backup(self, tag: str):
raise NotImplementedError
def copy(self, *, src_tag: str, dst_tag: str):
raise NotImplementedError
@abstractmethod
def restore(self, tag: str):
raise NotImplementedError
class _TensorBackuperNormal(TensorBackuper):
def __init__(self, source_getter):
super().__init__(source_getter=source_getter)
self._backups: dict[str, dict[str, torch.Tensor]] = defaultdict(dict)
@property
def backup_tags(self):
return list(self._backups)
def get(self, tag: str):
return self._backups[tag]
@torch.no_grad()
def backup(self, tag: str) -> None:
backup_dict = self._backups[tag]
for name, param in self._source_getter():
if name not in backup_dict:
backup_dict[name] = torch.empty_like(param, device=torch.device("cpu"), pin_memory=True)
backup_dict[name].copy_(param.detach(), non_blocking=True)
torch.cuda.synchronize()
@torch.no_grad()
def copy(self, *, src_tag: str, dst_tag: str):
for name in self._backups[dst_tag]:
self._backups[dst_tag][name].copy_(self._backups[src_tag][name])
@torch.no_grad()
def restore(self, tag: str) -> None:
backup_dict = self._backups[tag]
for name, param in self._source_getter():
assert name in backup_dict
param.copy_(backup_dict[name], non_blocking=True)
torch.cuda.synchronize()
class _TensorBackuperNoop(TensorBackuper):
def __init__(self, source_getter, single_tag):
super().__init__(source_getter=source_getter)
self._single_tag = single_tag
# Sanity check for safety
self._backup_hash_dict = None
@property
def backup_tags(self):
return [self._single_tag]
def get(self, tag: str):
ans = dict(self._source_getter())
ans = {k: v.detach() for k, v in ans.items()}
assert _compute_hash_dict(ans) == self._backup_hash_dict
return ans
def backup(self, tag: str) -> None:
assert tag == self._single_tag
self._backup_hash_dict = _compute_hash_dict(dict(self._source_getter()))
torch.cuda.synchronize()
def restore(self, tag: str) -> None:
assert tag == self._single_tag
assert _compute_hash_dict(dict(self._source_getter())) == self._backup_hash_dict
torch.cuda.synchronize()
def _compute_hash_dict(tensors: dict[str, torch.Tensor]):
return {k: _compute_hash_tensor(v) for k, v in tensors.items()}
def _compute_hash_tensor(x: torch.Tensor):
# Not a real/good hash, but pretty fast
x = x.contiguous()
x = x.view(-1)
x = x.view(torch.uint32)
x = x.sum()
return x.item()

View File

@@ -0,0 +1,62 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import datetime
import logging
import os
from slime.utils.misc import SingletonMeta
try:
from torch.utils.tensorboard import SummaryWriter
except ImportError:
SummaryWriter = None
logger = logging.getLogger(__name__)
class _TensorboardAdapter(metaclass=SingletonMeta):
_writer = None
"""
# Usage example: This will return the same instance every rank
# tb = _TensorboardAdapter(args) # Initialize on first call
# tb.log({"Loss": 0.1}, step=1)
# In other files:
# from tensorboard_utils import _TensorboardAdapter
# tb = _TensorboardAdapter(args) # No parameters needed to get existing instance
# tb.log({"Accuracy": 0.9}, step=1)
"""
def __init__(self, args):
assert args.use_tensorboard, f"{args.use_tensorboard=}"
tb_project_name = args.tb_project_name
tb_experiment_name = args.tb_experiment_name
if tb_project_name is not None or os.environ.get("TENSORBOARD_DIR", None):
if tb_project_name is not None and tb_experiment_name is None:
tb_experiment_name = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
self._initialize(tb_project_name, tb_experiment_name)
else:
raise ValueError("tb_project_name and tb_experiment_name, or TENSORBOARD_DIR are required")
def _initialize(self, tb_project_name, tb_experiment_name):
"""Actual initialization logic"""
# Get tensorboard directory from environment variable or use default path
tensorboard_dir = os.environ.get("TENSORBOARD_DIR", f"tensorboard_log/{tb_project_name}/{tb_experiment_name}")
os.makedirs(tensorboard_dir, exist_ok=True)
logger.info(f"Saving tensorboard log to {tensorboard_dir}.")
self._writer = SummaryWriter(tensorboard_dir)
def log(self, data, step):
"""Log data to tensorboard
Args:
data (dict): Dictionary containing metric names and values
step (int): Current step/epoch number
"""
for key in data:
self._writer.add_scalar(key, data[key], step)
def finish(self):
"""Close the tensorboard writer"""
self._writer.close()

92
slime/utils/timer.py Normal file
View File

@@ -0,0 +1,92 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
from contextlib import contextmanager
from functools import wraps
from time import time
import torch.distributed
from .misc import SingletonMeta
__all__ = ["Timer", "timer"]
logger = logging.getLogger(__name__)
class Timer(metaclass=SingletonMeta):
def __init__(self):
self.timers = {}
self.start_time = {}
def start(self, name):
assert name not in self.start_time, f"Timer {name} already started."
self.start_time[name] = time()
if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:
logger.info(f"Timer {name} start")
def end(self, name):
assert name in self.start_time, f"Timer {name} not started."
elapsed_time = time() - self.start_time[name]
self.add(name, elapsed_time)
del self.start_time[name]
if torch.distributed.is_initialized() and torch.distributed.get_rank() == 0:
logger.info(f"Timer {name} end (elapsed: {elapsed_time:.1f}s)")
def reset(self, name=None):
if name is None:
self.timers = {}
elif name in self.timers:
del self.timers[name]
def add(self, name, elapsed_time):
self.timers[name] = self.timers.get(name, 0) + elapsed_time
def log_dict(self):
return self.timers
@contextmanager
def context(self, name):
self.start(name)
try:
yield
finally:
self.end(name)
def timer(name_or_func):
"""
Can be used either as a decorator or a context manager:
@timer
def func():
...
or
with timer("block_name"):
...
"""
# When used as a context manager
if isinstance(name_or_func, str):
name = name_or_func
return Timer().context(name)
func = name_or_func
@wraps(func)
def wrapper(*args, **kwargs):
with Timer().context(func.__name__):
return func(*args, **kwargs)
return wrapper
@contextmanager
def inverse_timer(name):
Timer().end(name)
try:
yield
finally:
Timer().start(name)

View File

@@ -0,0 +1,24 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import wandb
from slime.utils.tensorboard_utils import _TensorboardAdapter
from . import wandb_utils
def init_tracking(args, primary: bool = True, **kwargs):
if primary:
wandb_utils.init_wandb_primary(args, **kwargs)
else:
wandb_utils.init_wandb_secondary(args, **kwargs)
# TODO further refactor, e.g. put TensorBoard init to the "init" part
def log(args, metrics, step_key: str):
if args.use_wandb:
wandb.log(metrics)
if args.use_tensorboard:
metrics_except_step = {k: v for k, v in metrics.items() if k != step_key}
_TensorboardAdapter(args).log(data=metrics_except_step, step=metrics[step_key])

View File

@@ -0,0 +1,25 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
from pathlib import Path
import torch
logger = logging.getLogger(__name__)
def save_debug_train_data(args, *, rollout_id, rollout_data):
if (path_template := args.save_debug_train_data) is not None:
rank = torch.distributed.get_rank()
path = Path(path_template.format(rollout_id=rollout_id, rank=rank))
logger.info(f"Save debug train data to {path}")
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(
dict(
rollout_id=rollout_id,
rank=rank,
rollout_data=rollout_data,
),
path,
)

View File

@@ -0,0 +1,51 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
from argparse import Namespace
from collections.abc import Callable
from copy import deepcopy
from slime.utils import tracking_utils
from slime.utils.metric_utils import compute_rollout_step
from slime.utils.timer import Timer
logger = logging.getLogger(__name__)
def log_perf_data_raw(
rollout_id: int, args: Namespace, is_primary_rank: bool, compute_total_fwd_flops: Callable
) -> None:
timer_instance = Timer()
log_dict_raw = deepcopy(timer_instance.log_dict())
timer_instance.reset()
if not is_primary_rank:
return
log_dict = {f"perf/{key}_time": val for key, val in log_dict_raw.items()}
if ("perf/actor_train_time" in log_dict) and (compute_total_fwd_flops is not None):
total_fwd_flops = compute_total_fwd_flops(seq_lens=timer_instance.seq_lens)
if "perf/log_probs_time" in log_dict:
log_dict["perf/log_probs_tflops"] = total_fwd_flops / log_dict["perf/log_probs_time"]
if "perf/ref_log_probs_time" in log_dict:
log_dict["perf/ref_log_probs_tflops"] = total_fwd_flops / log_dict["perf/ref_log_probs_time"]
if log_dict["perf/actor_train_time"] > 0:
log_dict["perf/actor_train_tflops"] = 3 * total_fwd_flops / log_dict["perf/actor_train_time"]
log_dict["perf/actor_train_tok_per_s"] = sum(timer_instance.seq_lens) / log_dict["perf/actor_train_time"]
if "perf/train_wait_time" in log_dict and "perf/train_time" in log_dict:
total_time = log_dict["perf/train_wait_time"] + log_dict["perf/train_time"]
if total_time > 0:
log_dict["perf/step_time"] = total_time
log_dict["perf/wait_time_ratio"] = log_dict["perf/train_wait_time"] / total_time
logger.info(f"perf {rollout_id}: {log_dict}")
step = compute_rollout_step(args, rollout_id)
log_dict["rollout/step"] = step
tracking_utils.log(args, log_dict, step_key="rollout/step")

View File

@@ -0,0 +1,74 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import dataclasses
import inspect
from typing import Annotated
import typer
def dataclass_cli(func, env_var_prefix: str = "SLIME_SCRIPT_"):
"""Modified from https://github.com/fastapi/typer/issues/154#issuecomment-1544876144"""
# The dataclass type is the first argument of the function.
sig = inspect.signature(func)
param = list(sig.parameters.values())[0]
dataclass_cls = param.annotation
assert dataclasses.is_dataclass(dataclass_cls)
# To construct the signature, we remove the first argument (self)
# from the dataclass __init__ signature.
signature = inspect.signature(dataclass_cls.__init__)
old_parameters = list(signature.parameters.values())
if len(old_parameters) > 0 and old_parameters[0].name == "self":
del old_parameters[0]
new_parameters = []
for param in old_parameters:
env_var_name = f"{env_var_prefix}{param.name.upper()}"
new_annotation = Annotated[param.annotation, typer.Option(envvar=env_var_name)]
new_parameters.append(param.replace(annotation=new_annotation))
def wrapped(**kwargs):
data = dataclass_cls(**kwargs)
print(f"Execute command with args: {data}")
return func(data)
wrapped.__signature__ = signature.replace(parameters=new_parameters)
wrapped.__doc__ = func.__doc__
wrapped.__name__ = func.__name__
wrapped.__qualname__ = func.__qualname__
return wrapped
# unit test
if __name__ == "__main__":
from typer.testing import CliRunner
@dataclasses.dataclass
class DemoArgs:
name: str
count: int = 1
app = typer.Typer()
@app.command()
@dataclass_cli
def main(args: DemoArgs):
print(f"{args.name}|{args.count}")
runner = CliRunner()
res1 = runner.invoke(app, [], env={"SLIME_SCRIPT_NAME": "EnvName", "SLIME_SCRIPT_COUNT": "10"})
print(f"{res1.stdout=}")
assert res1.exit_code == 0
assert "EnvName|10" in res1.stdout.strip()
res2 = runner.invoke(app, ["--count", "999"], env={"SLIME_SCRIPT_NAME": "EnvName"})
print(f"{res2.stdout=}")
assert res2.exit_code == 0
assert "EnvName|999" in res2.stdout.strip()
print("✅ All Tests Passed!")

143
slime/utils/types.py Normal file
View File

@@ -0,0 +1,143 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from enum import Enum
from typing import Any
import torch
@dataclass
class Sample:
"""The sample generated"""
group_index: int | None = None
index: int | None = None
# prompt - can be:
# - str: raw text prompt
# - list[dict[str, str]]: chat messages format
prompt: str | list[dict[str, str]] = ""
tokens: list[int] = field(default_factory=list)
multimodal_inputs: dict[str, Any] = None # raw multimodal data, e.g. images, videos, etc.
multimodal_train_inputs: dict[str, Any] = None # processed multimodal data, e.g. pixel_values, etc.
# response
response: str = ""
response_length: int = 0
label: str | None = None
reward: float | dict[str, Any] | None = None
loss_mask: list[int] | None = None
weight_versions: list[str] = field(default_factory=list)
rollout_log_probs: list[float] | None = None # Log probabilities from rollout engine
rollout_routed_experts: list[list[int]] | None = None # Routed experts from rollout engine
remove_sample: bool = False
class Status(Enum):
PENDING = "pending"
COMPLETED = "completed"
TRUNCATED = "truncated"
ABORTED = "aborted"
# Indicates a recoverable or non-critical failure during generation (e.g., tool call failure,
# external API error, parsing error). Unlike ABORTED, FAILED samples may still contain partial
# valid output and can be retried or handled gracefully.
FAILED = "failed"
status: Status = Status.PENDING
metadata: dict = field(default_factory=dict)
# metadata used during training, e.g., what loss to use for this sample.
train_metadata: dict | None = None
class SpecInfo:
spec_accept_token_num: int = 0
spec_draft_token_num: int = 0
spec_verify_ct: int = 0
spec_accept_rate: float = 0.0
spec_accept_length: float = 0.0
def add(self, meta_info: dict, response_length: int):
self.spec_accept_token_num += meta_info["spec_accept_token_num"]
self.spec_draft_token_num += meta_info["spec_draft_token_num"]
self.spec_verify_ct += meta_info["spec_verify_ct"]
if self.spec_draft_token_num > 0:
# Notice: this does not iclude the bonus token generated by verify step.
self.spec_accept_rate = self.spec_accept_token_num / self.spec_draft_token_num
# self.spec_accept_rate = meta_info["spec_accept_rate"] #
if self.spec_verify_ct > 0:
self.spec_accept_length = response_length / self.spec_verify_ct
def to_dict(self):
return {
"spec_accept_token_num": self.spec_accept_token_num,
"spec_draft_token_num": self.spec_draft_token_num,
"spec_verify_ct": self.spec_verify_ct,
"spec_accept_rate": self.spec_accept_rate,
"spec_accept_length": self.spec_accept_length,
}
@staticmethod
def from_dict(data: dict):
info = Sample.SpecInfo()
info.spec_accept_token_num = data.get("spec_accept_token_num", 0)
info.spec_draft_token_num = data.get("spec_draft_token_num", 0)
info.spec_verify_ct = data.get("spec_verify_ct", 0)
info.spec_accept_rate = data.get("spec_accept_rate", 0.0)
info.spec_accept_length = data.get("spec_accept_length", 0.0)
return info
spec_info: SpecInfo = field(default_factory=SpecInfo)
def to_dict(self):
value = self.__dict__.copy()
value["status"] = self.status.value
value["spec_info"] = self.spec_info.to_dict()
return value
@staticmethod
def from_dict(data: dict):
data["status"] = Sample.Status(data["status"])
data["spec_info"] = Sample.SpecInfo.from_dict(data.get("spec_info", {}))
return Sample(**data)
def get_reward_value(self, args) -> float:
return self.reward if not args.reward_key else self.reward[args.reward_key]
@property
def effective_response_length(self):
return sum(self.loss_mask) if self.loss_mask is not None else self.response_length
@dataclass(frozen=True)
class ParamInfo:
name: str
dtype: torch.dtype
shape: torch.Size
attrs: dict
size: int
src_rank: int
# A dict-based batch produced along the rollout -> training path
# In Megatron backend, several fields are converted to torch.Tensor lists on GPU
# before being consumed by data iterators (see megatron_utils.actor._get_rollout_data).
RolloutBatch = dict[str, list[torch.Tensor] | list[int] | list[float] | list[str]]
@dataclass
class MultimodalType:
name: str # Type identifier used in message content (e.g., "image")
placeholder: str # Placeholder token in conversation messages (e.g., "<image>")
class MultimodalTypes:
IMAGE = MultimodalType(name="image", placeholder="<image>")
VIDEO = MultimodalType(name="video", placeholder="<video>")
AUDIO = MultimodalType(name="audio", placeholder="<audio>")
@classmethod
def all(cls) -> list[MultimodalType]:
return [cls.IMAGE, cls.VIDEO, cls.AUDIO]
@classmethod
def get(cls, name: str) -> MultimodalType | None:
return next((m for m in cls.all() if m.name == name), None)

228
slime/utils/wandb_utils.py Normal file
View File

@@ -0,0 +1,228 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import logging
import os
from copy import deepcopy
import wandb
logger = logging.getLogger(__name__)
def _is_offline_mode(args) -> bool:
"""Detect whether W&B should run in offline mode.
Priority order:
1) args.wandb_mode if provided
2) WANDB_MODE environment variable
"""
if args.wandb_mode:
return args.wandb_mode == "offline"
return os.environ.get("WANDB_MODE") == "offline"
def init_wandb_primary(args):
if not args.use_wandb:
args.wandb_run_id = None
return
# Set W&B mode if specified (overrides WANDB_MODE env var)
if args.wandb_mode:
os.environ["WANDB_MODE"] = args.wandb_mode
if args.wandb_mode == "offline":
logger.info("W&B offline mode enabled. Data will be saved locally.")
elif args.wandb_mode == "disabled":
logger.info("W&B disabled mode enabled. No data will be logged.")
elif args.wandb_mode == "online":
logger.info("W&B online mode enabled. Data will be uploaded to cloud.")
offline = _is_offline_mode(args)
# Only perform explicit login when NOT offline
if (not offline) and args.wandb_key is not None:
wandb.login(key=args.wandb_key, host=args.wandb_host)
# Check if we should resume a previous run
# Priority: 1) wandb_resume_run_id from args, 2) wandb_run_id from args, 3) wandb_run_id from checkpoint file
resume_run_id = getattr(args, "wandb_resume_run_id", None) or getattr(args, "wandb_run_id", None)
if not resume_run_id:
resume_run_id = _load_wandb_run_id_from_checkpoint(args)
if resume_run_id:
# Resume an existing run
logger.info(f"Resuming W&B run with id: {resume_run_id}")
init_kwargs = {
"id": resume_run_id,
"entity": args.wandb_team,
"project": args.wandb_project,
"resume": "must", # Fail if run doesn't exist
"config": _compute_config_for_logging(args),
}
# Configure settings based on offline/online mode
if offline:
init_kwargs["settings"] = wandb.Settings(mode="offline")
else:
init_kwargs["settings"] = wandb.Settings(mode="shared", x_primary=True)
else:
# Create a new run
# add random 6 length string with characters
if args.wandb_random_suffix:
suffix = "_" + wandb.util.generate_id()
max_base_len = 128 - len(suffix)
group = args.wandb_group[:max_base_len] + suffix
run_name = f"{group}-RANK_{args.rank}"
else:
group = args.wandb_group
run_name = args.wandb_group
# Prepare wandb init parameters
init_kwargs = {
"entity": args.wandb_team,
"project": args.wandb_project,
"group": group,
"name": run_name,
"config": _compute_config_for_logging(args),
}
# Configure settings based on offline/online mode
if offline:
init_kwargs["settings"] = wandb.Settings(mode="offline")
else:
init_kwargs["settings"] = wandb.Settings(mode="shared", x_primary=True)
# Add custom directory if specified
if args.wandb_dir:
# Ensure directory exists to avoid backend crashes
os.makedirs(args.wandb_dir, exist_ok=True)
init_kwargs["dir"] = args.wandb_dir
logger.info(f"W&B logs will be stored in: {args.wandb_dir}")
wandb.init(**init_kwargs)
_init_wandb_common()
args.wandb_run_id = wandb.run.id
_save_wandb_run_id_to_checkpoint(args)
if resume_run_id:
logger.info(f"Successfully resumed W&B run: {wandb.run.url}")
def _load_wandb_run_id_from_checkpoint(args):
load_dir = getattr(args, "load", None)
if not load_dir:
return None
path = os.path.join(load_dir, "wandb_run_id.txt")
if os.path.exists(path):
with open(path, "r") as f:
run_id = f.read().strip()
if run_id:
logger.info(f"Loaded wandb_run_id from {path}: {run_id}")
return run_id
return None
def _save_wandb_run_id_to_checkpoint(args):
save_dir = getattr(args, "save", None)
if not save_dir:
return
os.makedirs(save_dir, exist_ok=True)
path = os.path.join(save_dir, "wandb_run_id.txt")
with open(path, "w") as f:
f.write(args.wandb_run_id)
logger.info(f"Saved wandb_run_id to {path}")
def _compute_config_for_logging(args):
output = deepcopy(args.__dict__)
whitelist_env_vars = [
"SLURM_JOB_ID",
# We may insert more default values here, and may also allow users to configure a whitelist
]
output["env_vars"] = {k: v for k, v in os.environ.items() if k in whitelist_env_vars}
return output
# https://docs.wandb.ai/guides/track/log/distributed-training/#track-all-processes-to-a-single-run
def init_wandb_secondary(args, router_addr=None):
wandb_run_id = getattr(args, "wandb_run_id", None)
if wandb_run_id is None:
return
# Set W&B mode if specified (same as primary)
if args.wandb_mode:
os.environ["WANDB_MODE"] = args.wandb_mode
offline = _is_offline_mode(args)
if (not offline) and args.wandb_key is not None:
wandb.login(key=args.wandb_key, host=args.wandb_host)
# Configure settings based on offline/online mode
if offline:
settings_kwargs = dict(mode="offline")
else:
settings_kwargs = dict(
mode="shared",
x_primary=False,
x_update_finish_state=False,
)
if args.sglang_enable_metrics and router_addr is not None:
logger.info(f"Forward SGLang metrics at {router_addr} to WandB.")
settings_kwargs |= dict(
x_stats_open_metrics_endpoints={
"sgl_engine": f"{router_addr}/engine_metrics",
},
x_stats_open_metrics_filters={
"sgl_engine.*": {},
},
)
init_kwargs = {
"id": wandb_run_id,
"entity": args.wandb_team,
"project": args.wandb_project,
"config": args.__dict__,
"resume": "allow",
"reinit": True,
"settings": wandb.Settings(**settings_kwargs),
}
# Add custom directory if specified
if args.wandb_dir:
os.makedirs(args.wandb_dir, exist_ok=True)
init_kwargs["dir"] = args.wandb_dir
wandb.init(**init_kwargs)
_init_wandb_common()
def _init_wandb_common():
wandb.define_metric("train/step")
wandb.define_metric("train/*", step_metric="train/step")
wandb.define_metric("rollout/step")
wandb.define_metric("rollout/*", step_metric="rollout/step")
wandb.define_metric("multi_turn/*", step_metric="rollout/step")
wandb.define_metric("passrate/*", step_metric="rollout/step")
wandb.define_metric("eval/step")
wandb.define_metric("eval/*", step_metric="eval/step")
wandb.define_metric("perf/*", step_metric="rollout/step")
def get_wandb_offline_dir(args):
"""Get the directory where offline W&B data is stored."""
if _is_offline_mode(args):
if args and hasattr(args, "wandb_dir") and args.wandb_dir:
# Use custom directory if specified
return args.wandb_dir
else:
# Default offline directory is ~/wandb/offline-run-<timestamp>
# This will be created automatically by wandb
return os.path.expanduser("~/wandb")
return None