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

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

View File

@@ -0,0 +1,50 @@
# Rollout Buffer
## Overview
Rollout Buffer is an independent component for asynchronous agent trajectory generation, with the main function of using the LLM OpenAI Server launched by slime training to generate agent trajectories.
### Workflow
```
slime Training Process ←─── HTTP API ───→ Rollout Buffer
↓ ↓
LLM Server ←─────── HTTP Requests ─────── Agent Framework
↓ ↓
Model Response ──────────────────────→ Trajectory Generation
```
For each different Agent task, there should be a corresponding independent Generator class, responsible for generating trajectories for that type of task. Rollout Buffer automatically reads and loads different types of Generators.
## Quick Start
### Basic Usage Process
1. **Copy Template**: Copy `base_generator.py` as a template
2. **Modify Task Type**: Change `TASK_TYPE` to your task name (cannot duplicate with other Generators)
3. **Implement Core Function**: Implement the `run_rollout()` function
4. **Optional Customization**: Rewrite five optional functions as needed
Generator files must end with `_generator.py` and be placed in the `generator/` directory:
```
generator/
├── base_generator.py # Math task implementation (default template)
└── your_task_generator.py # Your custom task
```
Each Generator file must define `TASK_TYPE` and `run_rollout()`.
In addition, Rollout Buffer also provides some customizable functions to meet special needs of different tasks. If no custom implementation is provided, the system will use default implementations (located in `slime_plugins/rollout_buffer/default_func.py`).
### Example Script
First, you need to follow [Example: Qwen3-4B Model](../../docs/en/models/qwen3-4B.md) to configure the environment, download data and convert model checkpoints. And then run the following scripts:
```bash
cd slime_plugins/rollout_buffer
bash rollout_buffer_example.sh
# In a different terminal
python buffer.py
```

View File

@@ -0,0 +1,51 @@
# Rollout Buffer
## 概述
Rollout Buffer 是用于辅助纯异步 agent 训练的独立组件,其主要功能是使用 slime 训练启动的 LLM OpenAI Server 进行智能体轨迹的生成。
### 工作流程
```
slime Training Process ←─── HTTP API ───→ Rollout Buffer
↓ ↓
LLM Server ←─────── HTTP Requests ─────── Agent Framework
↓ ↓
Model Response ──────────────────────→ Trajectory Generation
```
对于每一个不同的 Agent 任务,都应该对应一个独立的 Generator 类负责生成该类任务的轨迹。Rollout Buffer 会自动读取并加载不同类型的 Generator。
## 快速开始
### 基本使用流程
1. **复制模板**:将 `base_generator.py` 作为模板进行复制
2. **修改任务类型**:将 `TASK_TYPE` 修改为您的任务名称(不能与其他 Generator 重复)
3. **实现核心函数**:实现 `run_rollout()` 函数
4. **可选定制**:根据需要重写五个可选函数
Generator 文件必须以 `_generator.py` 结尾,并放置在 `generator/` 目录下:
```
generator/
├── base_generator.py # Math 任务实现(默认模板)
└── your_task_generator.py # 您的自定义任务
```
每个 Generator 文件必须定义 `TASK_TYPE``run_rollout()`
此外Rollout Buffer 还提供了一些可自定义的函数来满足不同任务的特殊需求。如果不提供自定义实现,系统将使用默认实现(位于 `slime_plugins/rollout_buffer/default_func.py`)。
### 示例脚本
请仿照 [示例Qwen3-4B 模型](../../docs/zh/models/qwen3-4B.md) 文档中配置好 slime 的运行环境,下载数据,并转换模型 ckpt。之后分别运行
```bash
cd slime_plugins/rollout_buffer
bash rollout_buffer_example.sh
# In a different terminal
python buffer.py
```

View File

@@ -0,0 +1,343 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import copy
import glob
import importlib.util
import json
import pathlib
import threading
import time
from typing import Any
import uvicorn
from fastapi import BackgroundTasks, FastAPI, HTTPException, Request
from pydantic import BaseModel
app = FastAPI(title="Rollout Buffer Server", debug=True)
def default_is_valid_group(group_data, min_valid_group_size, task_type):
instance_id, samples = group_data
return len(samples) >= min_valid_group_size
def default_get_group_data_meta_info(temp_data: dict[str, list[dict[str, Any]]]) -> dict[str, Any]:
"""
Default implementation for getting meta information about the temporary data
collected between get_batch calls.
"""
if not temp_data:
return {
"total_samples": 0,
"num_groups": 0,
"avg_group_size": 0,
"avg_reward": 0,
}
meta_info = {"total_samples": 0, "num_groups": len(temp_data)}
all_rewards = []
# Calculate per-group statistics
for _instance_id, samples in temp_data.items():
group_size = len(samples)
group_rewards = [s["reward"] for s in samples] # Calculate group reward standard deviation
meta_info["total_samples"] += group_size
all_rewards.extend(group_rewards)
# Calculate global statistics
meta_info["avg_group_size"] = meta_info["total_samples"] / meta_info["num_groups"]
if all_rewards:
meta_info["avg_reward"] = sum(all_rewards) / len(all_rewards)
else:
meta_info["avg_reward"] = 0
return meta_info
def discover_generators():
"""
Automatically discover generator modules in the generator directory.
Returns a dictionary mapping task_type to module with run_rollout function.
"""
generator_map = {}
generator_dir = pathlib.Path(__file__).parent / "generator"
# Find all files within generator_dir
for file_path in glob.glob(str(generator_dir / "*.py")):
if file_path.endswith("__init__.py"):
continue
try:
# Load the module
spec = importlib.util.spec_from_file_location("generator_module", file_path)
if spec is None or spec.loader is None:
print(f"Warning: Could not load spec for {file_path}")
continue
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
# Check if module has TASK_TYPE constant
if not hasattr(module, "TASK_TYPE"):
print(f"Warning: {file_path} does not define TASK_TYPE constant")
continue
# Check if module has run_rollout function
if not hasattr(module, "run_rollout"):
print(f"Warning: {file_path} does not define run_rollout function")
continue
task_type = module.TASK_TYPE
generator_info = {
"module": module,
"file_path": file_path,
"run_rollout": module.run_rollout,
}
# Check for optional functions and use defaults if not present
for func_name in [
"transform_group",
"is_valid_group",
"get_group_data_meta_info",
]:
generator_info[func_name] = getattr(module, func_name, None)
generator_map[task_type] = generator_info
print(f"Discovered generator: {task_type} -> {file_path}")
except Exception as e:
print(f"Error loading generator from {file_path}: {str(e)}")
continue
return generator_map
@app.middleware("http")
async def set_body_size(request: Request, call_next):
request._body_size_limit = 1_073_741_824 # 1GB
response = await call_next(request)
return response
class BufferResponse(BaseModel):
success: bool
message: str = ""
data: dict[str, Any] | None = None
class BufferQueue:
def __init__(
self,
group_size,
task_type="math",
transform_group_func=None,
is_valid_group_func=None,
get_group_data_meta_info_func=None,
):
self.data = {}
self.temp_data = {}
self.group_timestamps = {}
self.group_size = group_size
self.task_type = task_type
# Set up function handlers with defaults
self.is_valid_group_func = is_valid_group_func or default_is_valid_group
self.get_group_data_meta_info_func = get_group_data_meta_info_func or default_get_group_data_meta_info
self.transform_group_func = transform_group_func or (lambda group, task_type: group)
def append(self, item):
instance_id = item["instance_id"]
current_time = time.time()
# Update timestamp for this group
self.group_timestamps[instance_id] = current_time
if instance_id not in self.temp_data:
self.temp_data[instance_id] = [copy.deepcopy(item)]
else:
self.temp_data[instance_id].append(copy.deepcopy(item))
if instance_id not in self.data:
self.data[instance_id] = [item]
else:
self.data[instance_id].append(item)
def _get_valid_groups_with_timeout(self, del_data=False):
"""Get valid groups including timeout-based groups"""
valid_groups = {}
timed_out_groups = {}
finished_groups = []
for instance_id, group_data in self.data.items():
if self.is_valid_group_func((instance_id, group_data), self.group_size, self.task_type):
valid_groups[instance_id] = group_data
# Remove finished groups and timed out groups with insufficient data
if del_data:
for instance_id in finished_groups:
self.data.pop(instance_id, None)
self.group_timestamps.pop(instance_id, None)
print(f"Removed finished group {instance_id}")
# Combine normal valid groups and timeout groups
all_valid_groups = {**valid_groups, **timed_out_groups}
return all_valid_groups, finished_groups
def get(self):
output = {"data": [], "meta_info": {}}
# Get meta information about temp data before processing
meta_info = self.get_group_data_meta_info_func(self.temp_data)
output["meta_info"] = meta_info
valid_groups, finished_groups = self._get_valid_groups_with_timeout(del_data=True)
output["meta_info"]["finished_groups"] = finished_groups
print(f"meta info: {json.dumps(meta_info, indent=2)}")
valid_groups = list(valid_groups.items())
for instance_id, group in valid_groups:
# First filter individual items
transformed_group = self.transform_group_func((instance_id, group), self.task_type)
output["data"].extend(transformed_group[1])
if instance_id in self.data:
self.data.pop(instance_id)
return output
def __len__(self):
valid_groups, _ = self._get_valid_groups_with_timeout()
num = sum([len(v) for v in valid_groups.values()])
num_of_all_groups = sum([len(v) for v in self.data.values()])
print(f"valid_groups: {len(valid_groups)}, num: {num}, num_of_all_groups: {num_of_all_groups}")
return num
class RolloutBuffer:
def __init__(
self,
group_size=16,
task_type="math",
transform_group_func=None,
is_valid_group_func=None,
get_group_data_meta_info_func=None,
):
self.buffer = BufferQueue(
group_size=group_size,
task_type=task_type,
transform_group_func=transform_group_func,
is_valid_group_func=is_valid_group_func,
get_group_data_meta_info_func=get_group_data_meta_info_func,
)
self.lock = threading.RLock()
self.not_empty = threading.Condition(self.lock)
self.total_written = 0
self.total_read = 0
self.task_type = task_type
def write(self, data):
with self.lock:
self.buffer.append(data)
self.total_written += 1
self.not_empty.notify_all()
return data
def read(self):
with self.not_empty:
if len(self.buffer) == 0:
return {"data": [], "meta_info": {}}
# Don't clear temp_data for regular read operations
result = self.buffer.get()
self.total_read += len(result["data"])
return result
buffer = RolloutBuffer()
@app.post("/buffer/write", response_model=BufferResponse)
async def write_to_buffer(request: Request):
try:
data = await request.json()
item = buffer.write(data)
return BufferResponse(
success=True,
message="Data has been successfully written to buffer",
data={"data": [item], "meta_info": "write to buffer"},
)
except Exception as e:
print(f"Write failed: {str(e)}")
import traceback
traceback.print_exc()
raise HTTPException(status_code=500, detail=f"Write failed: {str(e)}") from e
@app.post("/get_rollout_data", response_model=BufferResponse)
async def get_rollout_data(request: Request):
items = buffer.read()
if not items["data"]:
return BufferResponse(
success=False,
message="No data available to read",
data={"data": [], "meta_info": items["meta_info"]},
)
print(f"return {len(items['data'])} items and save them to local")
buffer.buffer.temp_data = {}
return BufferResponse(
success=True,
message=f"Successfully read {len(items['data'])} items",
data=items,
)
def run_rollout(data: dict):
global buffer
# Auto-discover generators
generator_map = discover_generators()
task_type = data["task_type"]
if task_type not in generator_map:
print(f"Error: No generator found for task_type '{task_type}'")
print(f"Available generators: {list(generator_map.keys())}")
return
generator_info = generator_map[task_type]
print(f"Using generator: {generator_info['file_path']} for task_type: {task_type}")
buffer = RolloutBuffer(
group_size=int(data["num_repeat_per_sample"]),
task_type=task_type,
transform_group_func=generator_info.get("transform_group", None),
is_valid_group_func=generator_info.get("is_valid_group"),
get_group_data_meta_info_func=generator_info.get("get_group_data_meta_info"),
)
# Call the run_rollout function from the appropriate generator module
generator_info["run_rollout"](data)
print(f"Rollout completed successfully for task_type: {task_type}")
@app.post("/start_rollout")
async def start_rollout(request: Request, background: BackgroundTasks):
payload = await request.json()
background.add_task(run_rollout, payload)
return {"message": "Rollout started"}
if __name__ == "__main__":
uvicorn.run(
app,
host="0.0.0.0",
port=8889,
limit_concurrency=1000, # Connection concurrency limit
# limit_max_requests=1000000, # Maximum request limit
timeout_keep_alive=5, # Keep-alive timeout,
)

View File

@@ -0,0 +1,9 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from .base_generator import BaseGenerator, query_single_turn
__all__ = [
"BaseGenerator",
"query_single_turn",
]

View File

@@ -0,0 +1,354 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import copy
import json
import random
import time
import uuid
from functools import partial
from multiprocessing import Process, Queue
from time import sleep
import requests
from openai import OpenAI
from tqdm import tqdm
from slime.rollout.rm_hub import get_deepscaler_rule_based_reward
TASK_TYPE = "math"
SAMPLING_PARAMS = {
"top_p": 1,
}
def get_rule_based_math_reward(item):
messages = item["messages"]
label = item["label"]
assert messages[-1]["role"] == "assistant", "last message must be assistant, but got {}".format(
messages[-1]["role"]
)
response = messages[-1]["content"]
if response is None or len(response) == 0:
return 0
reward = get_deepscaler_rule_based_reward(response, label)
return reward
def query_single_turn(client, messages, sampling_params, tools=None):
base_payload = {
"messages": messages,
**sampling_params,
"model": "custom",
"stream": False,
"seed": random.randint(1, 10000000),
"tools": tools,
}
text = None
accumulated_tokens = 0
finish_reason = "stop"
for _attempt in range(6):
try:
# Create a fresh payload for each attempt
current_payload = copy.deepcopy(base_payload)
if text is not None:
# Update messages with current progress
current_messages = copy.deepcopy(messages)
current_messages.append({"role": "assistant", "content": text})
current_payload["messages"] = current_messages
# Adjust max_tokens based on accumulated tokens
if "max_tokens" in sampling_params:
current_payload["max_tokens"] = max(0, sampling_params["max_tokens"] - accumulated_tokens)
# Add continue flag for partial rollouts
current_payload["extra_body"] = {"continue_final_message": True}
if current_payload["max_tokens"] == 0:
break
response = client.chat.completions.create(**current_payload)
if len(response.choices) > 0:
finish_reason = response.choices[0].finish_reason
if finish_reason == "abort":
print(
f"query failed, reason: {response.choices[0].finish_reason}, currently generated: {response.usage.completion_tokens}"
)
accumulated_tokens += response.usage.completion_tokens
if text is None:
text = response.choices[0].message.content
else:
text += response.choices[0].message.content
sleep(10)
continue
if text is None:
text = response.choices[0].message.content
elif response.choices[0].message.content is not None:
text += response.choices[0].message.content
break
else:
print(f"Error in query, status code: {response.status_code}")
continue
except Exception as e:
print(f"query failed in single turn, error: {e}")
continue
# Update final messages
if len(messages) > 0 and messages[-1]["role"] == "assistant":
messages = messages[:-1]
messages.append({"role": "assistant", "content": text})
return messages, finish_reason
def worker_process(task_queue, done_queue, rollout_func, reward_func, client, sampling_params):
for line in iter(task_queue.get, "STOP"):
if isinstance(line, str):
item = json.loads(line)
else:
item = line
# try:
messages, finish_reason = rollout_func(client, item["prompt"], sampling_params)
item["uid"] = str(uuid.uuid4())
item["messages"] = messages
reward = reward_func(item)
item["rollout_index"] = 1
item["reward"] = reward
item["extra_info"] = {}
item.update(sampling_params)
item["timestamp"] = str(time.time())
item["round_number"] = len([_ for _ in item["messages"] if _["role"] == "assistant"])
item["finish_reason"] = finish_reason
output_item = {
"uid": item.pop("uid"),
"messages": messages,
"reward": reward,
"instance_id": item.pop("instance_id"),
"extra_info": item,
}
done_queue.put(output_item)
done_queue.put("COMPLETE")
class BaseGenerator:
def __init__(
self,
remote_engine_url,
remote_buffer_url,
num_repeat_per_sample=1,
queue_size=1000000,
num_process=10,
task_type="math",
max_tokens=4096,
num_repeats=10,
skip_instance_ids: list[str] | None = None,
):
self.queue_size = queue_size
self.num_process = num_process
self.remote_engine_url = remote_engine_url
self.remote_buffer_url = remote_buffer_url
self.num_repeat_per_sample = num_repeat_per_sample
self.task_type = task_type
self.max_tokens = max_tokens
self.num_repeats = num_repeats
# Ensure skip_instance_ids is a mutable list (copy to avoid modifying original)
self.skip_instance_ids = list(skip_instance_ids) if skip_instance_ids is not None else None
if self.skip_instance_ids is not None:
print(f"BaseGenerator initialized with {len(self.skip_instance_ids)} instance_ids to skip")
self.skip_instance_ids = self.skip_instance_ids * self.num_repeat_per_sample
if "/v1" in remote_engine_url:
self.client = OpenAI(api_key="test", base_url=remote_engine_url)
else:
remote_engine_url = remote_engine_url.strip("/") + "/v1"
self.client = OpenAI(api_key="test", base_url=remote_engine_url)
def send_data_to_buffer(self, data):
remote_buffer_url = self.remote_buffer_url.rstrip("/") + "/buffer/write"
for _ in range(2):
try:
response = requests.post(remote_buffer_url, json=data)
if response.status_code == 200:
break
else:
print(f"send data to buffer failed, status code: {response.status_code}")
continue
except Exception as e:
print(f"send data to buffer failed, error: {e}")
continue
def run(self, input_file, rollout_func, reward_func):
task_queue, done_queue = Queue(maxsize=self.queue_size), Queue(maxsize=self.queue_size)
def read_data_into_queue():
cnt = 0
items = []
skipped_count = 0
with open(input_file) as f:
for i, line in enumerate(f):
item = json.loads(line)
if "instance_id" not in item:
item["instance_id"] = i
items.append(item)
random.shuffle(items)
for _ in range(self.num_repeats):
for item in items:
for _ in range(self.num_repeat_per_sample):
item_repeat = copy.deepcopy(item)
if "uid" not in item_repeat:
item_repeat["uid"] = str(uuid.uuid4())
# Check if instance_id should be skipped
if self.skip_instance_ids is not None and item_repeat["instance_id"] in self.skip_instance_ids:
print(f"Skipping instance_id: {item_repeat['instance_id']}")
# Remove from skip list to handle potential duplicates in multiple epochs
self.skip_instance_ids.remove(item_repeat["instance_id"])
skipped_count += 1
continue
task_queue.put(item_repeat)
cnt += 1
time.sleep(300)
if skipped_count > 0:
remaining_skip_count = len(self.skip_instance_ids) if self.skip_instance_ids is not None else 0
print(
f"Rollout summary: skipped {skipped_count} instance_ids, {remaining_skip_count} still in skip list"
)
for _ in range(self.num_process):
task_queue.put("STOP")
processes = []
SAMPLING_PARAMS["max_tokens"] = self.max_tokens
for _ in range(self.num_process):
process = Process(
target=partial(worker_process, client=self.client, sampling_params=SAMPLING_PARAMS),
args=(task_queue, done_queue, rollout_func, reward_func),
)
process.start()
processes.append(process)
process = Process(target=read_data_into_queue)
process.start()
progress_bar = tqdm()
num_finished = 0
while num_finished < self.num_process:
item = done_queue.get()
if item == "COMPLETE":
num_finished += 1
else:
assert "reward" in item, f"reward not in item: {item}"
assert "instance_id" in item, f"instance_id not in item: {item}"
self.send_data_to_buffer(item)
progress_bar.update(1)
progress_bar.close()
return "finished"
def entry(self, input_file, rollout_func, reward_func, num_epoch=1):
for _ in range(num_epoch):
self.run(input_file, rollout_func, reward_func)
def run_rollout(data: dict):
print(f"Starting math rollout with data: {data}")
rollout_func = query_single_turn
reward_func = get_rule_based_math_reward
print("Waiting for 10 seconds for buffer server to start")
time.sleep(10)
global SAMPLING_PARAMS
for k, v in data["sampling_params"].items():
SAMPLING_PARAMS[k] = v
print(f"Set {k} to {v}", type(v))
generator = BaseGenerator(
data["remote_engine_url"],
data["remote_buffer_url"],
num_repeat_per_sample=int(data["num_repeat_per_sample"]),
queue_size=1000000,
max_tokens=int(data["sampling_params"]["max_tokens"]),
num_process=int(data.get("num_process", 100)),
task_type=data["task_type"],
skip_instance_ids=data.get("skip_instance_ids", None),
)
generator.entry(data["input_file"], rollout_func, reward_func, int(data.get("num_epoch", 1)))
def normalize_group_data(group, epsilon=1e-8, algo="grpo"):
print(f"Using math-specific normalization for group {group[0]}")
assert algo == "grpo", "Only 'grpo' is supported for now."
instance_id = group[0]
data = group[1]
rewards = [item["reward"] for item in data]
valid_rewards = [r for r in rewards if 1 >= r >= 0]
if set(valid_rewards) == {0}:
normalized_rewards = rewards
else:
mean_reward = sum(valid_rewards) / len(valid_rewards)
std_reward = (sum((r - mean_reward) ** 2 for r in valid_rewards) / len(valid_rewards)) ** 0.5
if std_reward < epsilon:
print(f"[Math Info] Zero variance in group {instance_id}, setting all to 0.")
normalized_rewards = [0.0 if 1 >= r >= 0 else r for r in rewards]
else:
normalized_rewards = [(r - mean_reward) / (std_reward + epsilon) if 1 >= r >= 0 else r for r in rewards]
for i, item in enumerate(data):
item["reward"] = normalized_rewards[i]
item["raw_reward"] = rewards[i]
return (instance_id, data)
def is_valid_group(group, min_valid_group_size, task_type="math"):
# Handle both tuple and list inputs
if isinstance(group, tuple):
instance_id, items = group
else:
items = group
# Count valid items (non-empty responses)
valid_indices = []
for i, item in enumerate(items):
if item["messages"][-1]["content"].strip():
valid_indices.append(i)
group_size = len(items)
valid_count = len(valid_indices)
# A group is finished if it has reached the target size
is_finished = group_size >= min_valid_group_size
is_valid = is_finished and valid_count >= min_valid_group_size
return is_valid

View File

@@ -0,0 +1,310 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
import time
from typing import Any
import aiohttp
import requests
import wandb
from transformers import AutoTokenizer
from slime.utils.async_utils import run
from slime.utils.mask_utils import MultiTurnLossMaskGenerator
from slime.utils.types import Sample
__all__ = ["generate_rollout"]
# Global variables for evaluation
TOKENIZER = None
START_ROLLOUT = True
def select_rollout_data(args, results, need_length):
"""
Select the most recent groups when there are too many samples.
Groups all samples by instance_id, sorts groups by timestamp.
Args:
args: Arguments containing configuration
results: List of rollout data items with timestamps
Returns:
Selected samples from the newest groups based on timestamp cutoff
"""
if not results:
return results
# Group samples by instance_id
groups = {}
for item in results:
assert "instance_id" in item, "instance_id must be in item"
instance_id = item["instance_id"]
if instance_id not in groups:
groups[instance_id] = []
groups[instance_id].append(item)
print(f"📊 Total groups: {len(groups)}, total samples: {len(results)}")
# If we don't have too many samples, return all
assert need_length < len(results), "need_length must be smaller than results length"
# Get timestamp for each group (use the latest timestamp in the group)
def get_group_timestamp(group_items):
timestamps = []
for item in group_items:
if "timestamp" in item:
timestamps.append(float(item["timestamp"]))
elif "extra_info" in item and "timestamp" in item["extra_info"]:
timestamps.append(float(item["extra_info"]["timestamp"]))
return max(timestamps) if timestamps else 0
# Create list of (group_id, timestamp, samples) and sort by timestamp
group_data = []
for group_id, group_items in groups.items():
group_timestamp = get_group_timestamp(group_items)
group_data.append((group_id, group_timestamp, group_items))
# Sort groups by timestamp (newest first)
group_data.sort(key=lambda x: x[1], reverse=True)
selected_groups = group_data[:need_length]
# Flatten selected groups back to sample list
selected_results = []
for _group_id, _timestamp, group_items in selected_groups:
selected_results.append(group_items)
# Statistics for monitoring
if selected_groups:
newest_ts = selected_groups[0][1]
oldest_ts = selected_groups[-1][1]
print(
f"📈 Selected {len(selected_groups)} groups with {len(selected_results)*args.n_samples_per_prompt} samples"
)
print(f"📈 Group timestamp range: {oldest_ts:.2f} to {newest_ts:.2f}")
print(f"📈 Time span: {newest_ts - oldest_ts:.2f} seconds")
return selected_results
def log_raw_info(args, all_meta_info, rollout_id):
final_meta_info = {}
if all_meta_info:
final_meta_info = {
"total_samples": sum(meta["total_samples"] for meta in all_meta_info if "total_samples" in meta)
}
total_samples = final_meta_info["total_samples"]
if total_samples > 0:
weighted_reward_sum = sum(
meta["avg_reward"] * meta["total_samples"]
for meta in all_meta_info
if "avg_reward" in meta and "total_samples" in meta
)
final_meta_info.update(
{
"avg_reward": weighted_reward_sum / total_samples,
}
)
if hasattr(args, "use_wandb") and args.use_wandb:
log_dict = {
"rollout/no_filter/total_samples": final_meta_info["total_samples"],
"rollout/no_filter/avg_reward": final_meta_info["avg_reward"],
}
try:
step = (
rollout_id
if not args.wandb_always_use_train_step
else rollout_id * args.rollout_batch_size * args.n_samples_per_prompt // args.global_batch_size
)
if args.use_wandb:
log_dict["rollout/step"] = step
wandb.log(log_dict)
if args.use_tensorboard:
from slime.utils.tensorboard_utils import _TensorboardAdapter
tb = _TensorboardAdapter(args)
tb.log(data=log_dict, step=step)
print(f"no filter rollout log {rollout_id}: {log_dict}")
except Exception as e:
print(f"Failed to log to wandb: {e}")
print(f"no filter rollout log {rollout_id}: {final_meta_info}")
else:
print(f"no filter rollout log {rollout_id}: {final_meta_info}")
async def get_rollout_data(api_base_url: str) -> tuple[list[dict[str, Any]], dict[str, Any]]:
start_time = time.time()
async with aiohttp.ClientSession() as session:
while True:
async with session.post(
f"{api_base_url}/get_rollout_data", json={}, timeout=aiohttp.ClientTimeout(total=120)
) as response:
response.raise_for_status()
resp_json = await response.json()
if resp_json["success"]:
break
await asyncio.sleep(3)
if time.time() - start_time > 30:
print("rollout data is not ready, have been waiting for 30 seconds")
# Reset start_time to continue waiting or handle timeout differently
start_time = time.time() # Or raise an exception, or return empty list
data = resp_json["data"]
meta_info = {}
if isinstance(data, list):
if "data" in data[0]:
data = [item["data"] for item in data]
elif isinstance(data, dict):
if "data" in data:
meta_info = data["meta_info"]
data = data["data"]
print(f"Meta info: {meta_info}")
required_keys = {"uid", "instance_id", "messages", "reward", "extra_info"}
for item in data:
if not required_keys.issubset(item.keys()):
raise ValueError(f"Missing required keys in response item: {item}")
return data, meta_info
def start_rollout(api_base_url: str, args, metadata):
url = f"{api_base_url}/start_rollout"
print(f"metadata: {metadata}")
finished_groups_instance_id_list = [item for sublist in metadata.values() for item in sublist]
payload = {
"num_process": str(getattr(args, "rollout_num_process", 100)),
"num_epoch": str(args.num_epoch or 3),
"remote_engine_url": f"http://{args.sglang_router_ip}:{args.sglang_router_port}",
"remote_buffer_url": args.rollout_buffer_url,
"task_type": args.rollout_task_type,
"input_file": args.prompt_data,
"num_repeat_per_sample": str(args.n_samples_per_prompt),
"max_tokens": str(args.rollout_max_response_len),
"sampling_params": {
"max_tokens": args.rollout_max_response_len,
"temperature": args.rollout_temperature,
"top_p": args.rollout_top_p,
},
"tokenizer_path": args.hf_checkpoint,
"skip_instance_ids": finished_groups_instance_id_list,
}
print("start rollout with payload: ", payload)
while True:
try:
resp = requests.post(url, json=payload, timeout=10)
resp.raise_for_status()
data = resp.json()
print(f"[start_rollout] Success: {data}")
return data
except Exception as e:
print(f"[start_rollout] Failed to send rollout config: {e}")
async def generate_rollout_async(args, rollout_id: int, data_buffer, evaluation: bool = False) -> dict[str, Any]:
global START_ROLLOUT
if evaluation:
raise NotImplementedError("Evaluation rollout is not implemented")
if START_ROLLOUT:
metadata = data_buffer.get_metadata()
start_inform = start_rollout(args.rollout_buffer_url, args, metadata)
print(f"start rollout with payload: {start_inform}")
print(f"start rollout id: {rollout_id}")
START_ROLLOUT = False
data_number_to_fetch = args.rollout_batch_size * args.n_samples_per_prompt - data_buffer.get_buffer_length()
if data_number_to_fetch <= 0:
print(
f"❕buffer length: {data_buffer.get_buffer_length()}, buffer has enough data, return {args.rollout_batch_size} prompts"
)
return data_buffer.get_samples(args.rollout_batch_size)
assert (
data_number_to_fetch % args.n_samples_per_prompt == 0
), "data_number_to_fetch must be a multiple of n_samples_per_prompt"
print(f"INFO: buffer length: {data_buffer.get_buffer_length()}, data_number_to_fetch: {data_number_to_fetch}")
base_url = args.rollout_buffer_url
tokenizer = AutoTokenizer.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
retry_times = 0
results = []
all_meta_info = []
if args.fetch_trajectory_retry_times == -1:
print(
"⚠️ [get_rollout_data] Fetch trajectory retry times set to -1, will retry indefinitely until sufficient data is collected"
)
while args.fetch_trajectory_retry_times == -1 or retry_times < args.fetch_trajectory_retry_times:
try:
while len(results) < data_number_to_fetch:
time.sleep(5)
data, meta_info = await get_rollout_data(api_base_url=base_url)
results.extend(data)
if meta_info:
all_meta_info.append(meta_info)
print(f"get rollout data with length: {len(results)}")
break
except Exception as err:
print(f"[get_rollout_data] Failed to get rollout data: {err}, retry times: {retry_times}")
retry_times += 1
log_raw_info(args, all_meta_info, rollout_id)
# Apply group-based data selection if there are too many samples
results = select_rollout_data(args, results, data_number_to_fetch // args.n_samples_per_prompt)
if len(all_meta_info) > 0 and "finished_groups" in all_meta_info[0]:
finished_groups_instance_id_list = []
for item in all_meta_info:
finished_groups_instance_id_list.extend(item["finished_groups"])
data_buffer.update_metadata({str(rollout_id): finished_groups_instance_id_list})
print("finally get rollout data with length: ", len(results))
sample_results = []
for _i, group_record in enumerate(results):
group_results = []
for record in group_record:
oai_messages = record["messages"]
mask_generator = MultiTurnLossMaskGenerator(tokenizer, tokenizer_type=args.loss_mask_type)
token_ids, loss_mask = mask_generator.get_loss_mask(oai_messages)
response_length = mask_generator.get_response_lengths([loss_mask])[0]
loss_mask = loss_mask[-response_length:]
group_results.append(
Sample(
index=record["instance_id"],
prompt=record["uid"],
tokens=token_ids,
response_length=response_length,
reward=record["reward"],
status=(
Sample.Status.COMPLETED
if "finish_reason" not in record["extra_info"]
or record["extra_info"]["finish_reason"] != "length"
else Sample.Status.TRUNCATED
),
loss_mask=loss_mask,
metadata={**record["extra_info"]},
)
)
sample_results.append(group_results)
data_buffer.add_samples(sample_results)
final_return_results = data_buffer.get_samples(args.rollout_batch_size) # type: ignore
return final_return_results
def generate_rollout(args, rollout_id, data_buffer, evaluation=False):
"""Generate rollout for both training and evaluation."""
return run(generate_rollout_async(args, rollout_id, data_buffer, evaluation))

View File

@@ -0,0 +1,137 @@
#!/bin/bash
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
# for rerun the task
pkill -9 sglang
sleep 3
ray stop --force
pkill -9 ray
pkill -9 python
sleep 3
pkill -9 ray
pkill -9 python
set -ex
export PYTHONBUFFERED=16
# DeepSeek-R1-Distill-Qwen-7B
MODEL_ARGS=(
--swiglu
--num-layers 28
--hidden-size 3584
--ffn-hidden-size 18944
--num-attention-heads 28
--group-query-attention
--num-query-groups 4
--max-position-embeddings 131072
--seq-length 4096
--use-rotary-position-embeddings
--disable-bias-linear
--add-qkv-bias
--normalization "RMSNorm"
--norm-epsilon 1e-06
--rotary-base 10000
--vocab-size 152064
--accumulate-allreduce-grads-in-fp32
--attention-softmax-in-fp32
--attention-backend flash
--moe-token-dispatcher-type alltoall
--untie-embeddings-and-output-weights
--attention-dropout 0.0
--hidden-dropout 0.0
)
CKPT_ARGS=(
--hf-checkpoint /root/DeepSeek-R1-Distill-Qwen-7B
--ref-load /root/DeepSeek-R1-Distill-Qwen-7B_torch_dist
--save-interval 100
--save /root/DeepSeek-R1-Distill-Qwen-7B_slime
)
ROLLOUT_ARGS=(
--rollout-function-path slime_plugins.rollout_buffer.rollout_buffer_example.generate_rollout
--rm-type deepscaler
--prompt-data /root/dapo-math-17k/dapo-math-17k.jsonl
--input-key prompt
--label-key label
--num-rollout 3000
--rollout-batch-size 128
--rollout-max-response-len 8192
--rollout-temperature 0.8
--rollout-shuffle
--n-samples-per-prompt 8
--global-batch-size 1024
--micro-batch-size 8
--ref-micro-batch-size 8
--use-dynamic-batch-size
--max-tokens-per-gpu 9216
--balance-data
)
DISTRIBUTED_ARGS=(
--tensor-model-parallel-size 2
--pipeline-model-parallel-size 1
--context-parallel-size 1
--sequence-parallel
)
PERF_ARGS=(
--recompute-granularity full
--recompute-method uniform
--recompute-num-layers 1
)
GRPO_ARGS=(
--advantage-estimator grpo
--use-kl-loss
--kl-loss-coef 0.001
--kl-loss-type low_var_kl
--entropy-coef 0.00
)
OPTIMIZER_ARGS=(
--lr 1e-6
--lr-decay-style constant
--weight-decay 0.1
--adam-beta1 0.9
--adam-beta2 0.98
)
WANDB_ARGS=(
# --use-wandb
)
# launch the master node of ray in container
export MASTER_ADDR=${MASTER_ADDR:-"127.0.0.1"}
ray start --head --node-ip-address ${MASTER_ADDR} --num-gpus 8 --disable-usage-stats
ray job submit --address="http://127.0.0.1:8265" \
--runtime-env-json='{
"env_vars": {
"PYTHONPATH": "/root/Megatron-LM/",
"CUDA_DEVICE_MAX_CONNECTIONS": "1",
"NCCL_CUMEM_ENABLE": "0"
}
}' \
-- python3 train_async.py \
--actor-num-nodes 1 \
--actor-num-gpus-per-node 4 \
--rollout-num-gpus 4 \
--rollout-num-gpus-per-engine 1 \
${MODEL_ARGS[@]} \
${CKPT_ARGS[@]} \
${ROLLOUT_ARGS[@]} \
${OPTIMIZER_ARGS[@]} \
${GRPO_ARGS[@]} \
${DISTRIBUTED_ARGS[@]} \
${WANDB_ARGS[@]} \
${PERF_ARGS[@]} \
--rollout-buffer-url http://${MASTER_ADDR}:8889 \
--keep-old-actor \
--disable-rewards-normalization \
--loss-mask-type distill_qwen \
--log-passrate