初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
4
slime/utils/__init__.py
Normal file
4
slime/utils/__init__.py
Normal 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."""
|
||||
BIN
slime/utils/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/arguments.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/arguments.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/async_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/async_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/context_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/context_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/data.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/data.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/distributed_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/distributed_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/eval_config.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/eval_config.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/flops_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/flops_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/fp8_kernel.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/fp8_kernel.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/health_monitor.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/health_monitor.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/http_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/http_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/iter_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/iter_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/logging_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/logging_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/mask_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/mask_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/megatron_bridge_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/megatron_bridge_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/memory_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/memory_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/metric_checker.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/metric_checker.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/metric_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/metric_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/misc.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/misc.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/ppo_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/ppo_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/processing_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/processing_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/profile_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/profile_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/ray_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/ray_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/reloadable_process_group.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/reloadable_process_group.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/routing_replay.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/routing_replay.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/seqlen_balancing.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/seqlen_balancing.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/tensor_backper.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/tensor_backper.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/tensorboard_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/tensorboard_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/timer.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/timer.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/tracking_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/tracking_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/train_dump_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/train_dump_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/train_metric_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/train_metric_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/typer_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/typer_utils.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/types.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/types.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime/utils/__pycache__/wandb_utils.cpython-312.pyc
Normal file
BIN
slime/utils/__pycache__/wandb_utils.cpython-312.pyc
Normal file
Binary file not shown.
1737
slime/utils/arguments.py
Normal file
1737
slime/utils/arguments.py
Normal file
File diff suppressed because it is too large
Load Diff
39
slime/utils/async_utils.py
Normal file
39
slime/utils/async_utils.py
Normal 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)
|
||||
18
slime/utils/context_utils.py
Normal file
18
slime/utils/context_utils.py
Normal 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
342
slime/utils/data.py
Normal 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,
|
||||
)
|
||||
|
||||
3
slime/utils/debug_utils/__init__.py
Normal file
3
slime/utils/debug_utils/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
76
slime/utils/debug_utils/display_debug_rollout_data.py
Normal file
76
slime/utils/debug_utils/display_debug_rollout_data.py
Normal 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)
|
||||
53
slime/utils/debug_utils/replay_reward_fn.py
Normal file
53
slime/utils/debug_utils/replay_reward_fn.py
Normal 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)
|
||||
61
slime/utils/debug_utils/send_to_sglang.py
Normal file
61
slime/utils/debug_utils/send_to_sglang.py
Normal 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)
|
||||
157
slime/utils/distributed_utils.py
Normal file
157
slime/utils/distributed_utils.py
Normal 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
208
slime/utils/eval_config.py
Normal 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
|
||||
3
slime/utils/external_utils/__init__.py
Normal file
3
slime/utils/external_utils/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
BIN
slime/utils/external_utils/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime/utils/external_utils/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
Binary file not shown.
275
slime/utils/external_utils/command_utils.py
Normal file
275
slime/utils/external_utils/command_utils.py
Normal 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
130
slime/utils/flops_utils.py
Normal 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
82
slime/utils/fp8_kernel.py
Normal 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
|
||||
106
slime/utils/health_monitor.py
Normal file
106
slime/utils/health_monitor.py
Normal 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
263
slime/utils/http_utils.py
Normal 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
45
slime/utils/iter_utils.py
Normal 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
|
||||
22
slime/utils/logging_utils.py
Normal file
22
slime/utils/logging_utils.py
Normal 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
185
slime/utils/mask_utils.py
Normal 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
|
||||
25
slime/utils/megatron_bridge_utils.py
Normal file
25
slime/utils/megatron_bridge_utils.py
Normal 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")
|
||||
47
slime/utils/memory_utils.py
Normal file
47
slime/utils/memory_utils.py
Normal 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
|
||||
31
slime/utils/metric_checker.py
Normal file
31
slime/utils/metric_checker.py
Normal 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
121
slime/utils/metric_utils.py
Normal 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
94
slime/utils/misc.py
Normal 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
718
slime/utils/ppo_utils.py
Normal 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
|
||||
37
slime/utils/processing_utils.py
Normal file
37
slime/utils/processing_utils.py
Normal 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")
|
||||
150
slime/utils/profile_utils.py
Normal file
150
slime/utils/profile_utils.py
Normal 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
10
slime/utils/ray_utils.py
Normal 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
|
||||
285
slime/utils/reloadable_process_group.py
Normal file
285
slime/utils/reloadable_process_group.py
Normal 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
|
||||
30
slime/utils/rocm_checkpoint_writer.py
Normal file
30
slime/utils/rocm_checkpoint_writer.py
Normal 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)
|
||||
95
slime/utils/routing_replay.py
Normal file
95
slime/utils/routing_replay.py
Normal 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)
|
||||
186
slime/utils/seqlen_balancing.py
Normal file
186
slime/utils/seqlen_balancing.py
Normal 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
|
||||
118
slime/utils/tensor_backper.py
Normal file
118
slime/utils/tensor_backper.py
Normal 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()
|
||||
62
slime/utils/tensorboard_utils.py
Normal file
62
slime/utils/tensorboard_utils.py
Normal 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
92
slime/utils/timer.py
Normal 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)
|
||||
24
slime/utils/tracking_utils.py
Normal file
24
slime/utils/tracking_utils.py
Normal 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])
|
||||
25
slime/utils/train_dump_utils.py
Normal file
25
slime/utils/train_dump_utils.py
Normal 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,
|
||||
)
|
||||
51
slime/utils/train_metric_utils.py
Normal file
51
slime/utils/train_metric_utils.py
Normal 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")
|
||||
74
slime/utils/typer_utils.py
Normal file
74
slime/utils/typer_utils.py
Normal 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
143
slime/utils/types.py
Normal 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
228
slime/utils/wandb_utils.py
Normal 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
|
||||
Reference in New Issue
Block a user