初始化项目,由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

39
slime_plugins/PKG-INFO Normal file
View File

@@ -0,0 +1,39 @@
Metadata-Version: 2.4
Name: slime
Version: 0.1.0
Author: slime Team
Classifier: Programming Language :: Python :: 3.10
Classifier: Programming Language :: Python :: 3.11
Classifier: Programming Language :: Python :: 3.12
Classifier: Environment :: GPU :: NVIDIA CUDA
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Classifier: Topic :: System :: Distributed Computing
Requires-Python: >=3.10
License-File: LICENSE
Requires-Dist: accelerate
Requires-Dist: blobfile
Requires-Dist: datasets
Requires-Dist: httpx[http2]
Requires-Dist: mcp[cli]
Requires-Dist: megatron-bridge @ git+https://github.com/fzyzcjy/Megatron-Bridge.git@dev_rl
Requires-Dist: memray
Requires-Dist: nvidia-modelopt[torch]>=0.37.0
Requires-Dist: omegaconf
Requires-Dist: pillow
Requires-Dist: pylatexenc
Requires-Dist: pyyaml
Requires-Dist: ray[default]
Requires-Dist: ring_flash_attn
Requires-Dist: sglang-router>=0.2.3
Requires-Dist: tensorboard
Requires-Dist: transformers
Requires-Dist: wandb
Requires-Dist: liger_kernel
Provides-Extra: fsdp
Requires-Dist: torch>=2.0; extra == "fsdp"
Dynamic: author
Dynamic: classifier
Dynamic: license-file
Dynamic: provides-extra
Dynamic: requires-dist
Dynamic: requires-python

140
slime_plugins/SOURCES.txt Normal file
View File

@@ -0,0 +1,140 @@
LICENSE
README.md
pyproject.toml
setup.py
slime/__init__.py
slime.egg-info/PKG-INFO
slime.egg-info/SOURCES.txt
slime.egg-info/dependency_links.txt
slime.egg-info/requires.txt
slime.egg-info/top_level.txt
slime/backends/__init__.py
slime/backends/fsdp_utils/__init__.py
slime/backends/fsdp_utils/actor.py
slime/backends/fsdp_utils/arguments.py
slime/backends/fsdp_utils/checkpoint.py
slime/backends/fsdp_utils/data_packing.py
slime/backends/fsdp_utils/lr_scheduler.py
slime/backends/fsdp_utils/update_weight_utils.py
slime/backends/fsdp_utils/kernels/__init__.py
slime/backends/fsdp_utils/kernels/fused_experts.py
slime/backends/fsdp_utils/kernels/fused_moe_triton_backward_kernels.py
slime/backends/fsdp_utils/models/__init__.py
slime/backends/fsdp_utils/models/qwen3_moe.py
slime/backends/fsdp_utils/models/qwen3_moe_hf.py
slime/backends/megatron_utils/__init__.py
slime/backends/megatron_utils/actor.py
slime/backends/megatron_utils/arguments.py
slime/backends/megatron_utils/checkpoint.py
slime/backends/megatron_utils/ci_utils.py
slime/backends/megatron_utils/cp_utils.py
slime/backends/megatron_utils/data.py
slime/backends/megatron_utils/initialize.py
slime/backends/megatron_utils/loss.py
slime/backends/megatron_utils/misc_utils.py
slime/backends/megatron_utils/model.py
slime/backends/megatron_utils/model_provider.py
slime/backends/megatron_utils/sglang.py
slime/backends/megatron_utils/megatron_to_hf/__init__.py
slime/backends/megatron_utils/megatron_to_hf/deepseekv3.py
slime/backends/megatron_utils/megatron_to_hf/glm4.py
slime/backends/megatron_utils/megatron_to_hf/glm4moe.py
slime/backends/megatron_utils/megatron_to_hf/llama.py
slime/backends/megatron_utils/megatron_to_hf/mimo.py
slime/backends/megatron_utils/megatron_to_hf/qwen2.py
slime/backends/megatron_utils/megatron_to_hf/qwen3_next.py
slime/backends/megatron_utils/megatron_to_hf/qwen3moe.py
slime/backends/megatron_utils/megatron_to_hf/processors/__init__.py
slime/backends/megatron_utils/megatron_to_hf/processors/padding_remover.py
slime/backends/megatron_utils/megatron_to_hf/processors/quantizer.py
slime/backends/megatron_utils/update_weight/__init__.py
slime/backends/megatron_utils/update_weight/common.py
slime/backends/megatron_utils/update_weight/hf_weight_iterator_base.py
slime/backends/megatron_utils/update_weight/hf_weight_iterator_bridge.py
slime/backends/megatron_utils/update_weight/hf_weight_iterator_direct.py
slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py
slime/backends/megatron_utils/update_weight/update_weight_from_tensor.py
slime/backends/sglang_utils/__init__.py
slime/backends/sglang_utils/arguments.py
slime/backends/sglang_utils/sglang_engine.py
slime/ray/__init__.py
slime/ray/actor_group.py
slime/ray/placement_group.py
slime/ray/ray_actor.py
slime/ray/rollout.py
slime/ray/train_actor.py
slime/ray/utils.py
slime/rollout/__init__.py
slime/rollout/base_types.py
slime/rollout/data_source.py
slime/rollout/on_policy_distillation.py
slime/rollout/sglang_rollout.py
slime/rollout/sleep_rollout.py
slime/rollout/filter_hub/__init__.py
slime/rollout/filter_hub/base_types.py
slime/rollout/filter_hub/dynamic_sampling_filters.py
slime/rollout/rm_hub/__init__.py
slime/rollout/rm_hub/deepscaler.py
slime/rollout/rm_hub/f1.py
slime/rollout/rm_hub/gpqa.py
slime/rollout/rm_hub/ifbench.py
slime/rollout/rm_hub/math_dapo_utils.py
slime/rollout/rm_hub/math_utils.py
slime/router/__init__.py
slime/router/router.py
slime/router/middleware_hub/__init__.py
slime/router/middleware_hub/radix_tree.py
slime/router/middleware_hub/radix_tree_middleware.py
slime/utils/__init__.py
slime/utils/arguments.py
slime/utils/async_utils.py
slime/utils/context_utils.py
slime/utils/data.py
slime/utils/distributed_utils.py
slime/utils/eval_config.py
slime/utils/flops_utils.py
slime/utils/fp8_kernel.py
slime/utils/health_monitor.py
slime/utils/http_utils.py
slime/utils/iter_utils.py
slime/utils/logging_utils.py
slime/utils/mask_utils.py
slime/utils/megatron_bridge_utils.py
slime/utils/memory_utils.py
slime/utils/metric_checker.py
slime/utils/metric_utils.py
slime/utils/misc.py
slime/utils/ppo_utils.py
slime/utils/processing_utils.py
slime/utils/profile_utils.py
slime/utils/ray_utils.py
slime/utils/reloadable_process_group.py
slime/utils/rocm_checkpoint_writer.py
slime/utils/routing_replay.py
slime/utils/seqlen_balancing.py
slime/utils/tensor_backper.py
slime/utils/tensorboard_utils.py
slime/utils/timer.py
slime/utils/tracking_utils.py
slime/utils/train_dump_utils.py
slime/utils/train_metric_utils.py
slime/utils/typer_utils.py
slime/utils/types.py
slime/utils/wandb_utils.py
slime/utils/debug_utils/__init__.py
slime/utils/debug_utils/display_debug_rollout_data.py
slime/utils/debug_utils/replay_reward_fn.py
slime/utils/debug_utils/send_to_sglang.py
slime/utils/external_utils/__init__.py
slime/utils/external_utils/command_utils.py
slime_plugins/__init__.py
slime_plugins/mbridge/__init__.py
slime_plugins/mbridge/glm4.py
slime_plugins/mbridge/glm4moe.py
slime_plugins/mbridge/mimo.py
slime_plugins/mbridge/qwen3_next.py
slime_plugins/megatron_bridge/__init__.py
slime_plugins/models/__init__.py
slime_plugins/models/glm4.py
slime_plugins/models/hf_attention.py
slime_plugins/models/qwen3_next.py

View File

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

Binary file not shown.

View File

@@ -0,0 +1 @@

View File

@@ -0,0 +1,9 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from .glm4 import GLM4Bridge
from .glm4moe import GLM4MoEBridge
from .mimo import MimoBridge
from .qwen3_next import Qwen3NextBridge
__all__ = ["GLM4Bridge", "GLM4MoEBridge", "Qwen3NextBridge", "MimoBridge"]

Binary file not shown.

Binary file not shown.

View File

@@ -0,0 +1,112 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec
from mbridge.core import LLMBridge, register_model
@register_model("glm4")
class GLM4Bridge(LLMBridge):
"""
Bridge implementation for Qwen2 models.
This class extends LLMBridge to provide specific configurations and
optimizations for Qwen2 models, handling the conversion between
Hugging Face Qwen2 format and Megatron-Core.
"""
_DIRECT_MAPPING = {
"embedding.word_embeddings.weight": "model.embed_tokens.weight",
"decoder.final_layernorm.weight": "model.norm.weight",
"output_layer.weight": "lm_head.weight",
}
_ATTENTION_MAPPING = {
"self_attention.linear_proj.weight": ["model.layers.{layer_number}.self_attn.o_proj.weight"],
"self_attention.linear_qkv.layer_norm_weight": ["model.layers.{layer_number}.input_layernorm.weight"],
"self_attention.q_layernorm.weight": ["model.layers.{layer_number}.self_attn.q_norm.weight"],
"self_attention.k_layernorm.weight": ["model.layers.{layer_number}.self_attn.k_norm.weight"],
"self_attention.linear_qkv.weight": [
"model.layers.{layer_number}.self_attn.q_proj.weight",
"model.layers.{layer_number}.self_attn.k_proj.weight",
"model.layers.{layer_number}.self_attn.v_proj.weight",
],
"self_attention.linear_qkv.bias": [
"model.layers.{layer_number}.self_attn.q_proj.bias",
"model.layers.{layer_number}.self_attn.k_proj.bias",
"model.layers.{layer_number}.self_attn.v_proj.bias",
],
}
_MLP_MAPPING = {
"mlp.linear_fc1.weight": [
"model.layers.{layer_number}.mlp.gate_up_proj.weight",
],
"mlp.linear_fc1.layer_norm_weight": ["model.layers.{layer_number}.post_attention_layernorm.weight"],
"mlp.linear_fc2.weight": ["model.layers.{layer_number}.mlp.down_proj.weight"],
}
def _build_config(self):
"""
Build the configuration for Qwen2 models.
Configures Qwen2-specific parameters such as QKV bias settings and
layer normalization options.
Returns:
TransformerConfig: Configuration object for Qwen2 models
"""
return self._build_base_config(
# qwen2
add_qkv_bias=True,
qk_layernorm=False,
post_mlp_layernorm=True,
post_self_attn_layernorm=True,
rotary_interleaved=True,
)
def _get_transformer_layer_spec(self):
"""
Gets the transformer layer specification.
Creates and returns a specification for the transformer layers based on
the current configuration.
Returns:
TransformerLayerSpec: Specification for transformer layers
Raises:
AssertionError: If normalization is not RMSNorm
"""
transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(
post_self_attn_layernorm=True,
post_mlp_layernorm=True,
)
return transformer_layer_spec
def _weight_name_mapping_mcore_to_hf(self, mcore_weights_name: str) -> list[str]:
"""
Map MCore weight names to Hugging Face weight names.
Args:
mcore_weights_name: MCore weight name
Returns:
list: Corresponding Hugging Face weight names
"""
assert "_extra_state" not in mcore_weights_name, "extra_state should not be loaded"
if mcore_weights_name in self._DIRECT_MAPPING:
return [self._DIRECT_MAPPING[mcore_weights_name]]
if "post_self_attn_layernorm" in mcore_weights_name:
layer_number = mcore_weights_name.split(".")[2]
return [f"model.layers.{layer_number}.post_self_attn_layernorm.weight"]
elif "post_mlp_layernorm" in mcore_weights_name:
layer_number = mcore_weights_name.split(".")[2]
return [f"model.layers.{layer_number}.post_mlp_layernorm.weight"]
elif "self_attention" in mcore_weights_name:
return self._weight_name_mapping_attention(mcore_weights_name)
elif "mlp" in mcore_weights_name:
return self._weight_name_mapping_mlp(mcore_weights_name)
else:
raise NotImplementedError(f"Unsupported parameter name: {mcore_weights_name}")

View File

@@ -0,0 +1,125 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import re
from mbridge.core import register_model
from mbridge.models import Qwen2Bridge, Qwen2MoEBridge
@register_model("glm4_moe")
class GLM4MoEBridge(Qwen2MoEBridge):
"""
Bridge implementation for Qwen2 models.
This class extends LLMBridge to provide specific configurations and
optimizations for Qwen2 models, handling the conversion between
Hugging Face Qwen2 format and Megatron-Core.
"""
_MLP_MAPPING = {
**(Qwen2MoEBridge._MLP_MAPPING),
**(Qwen2Bridge._MLP_MAPPING),
"mlp.router.expert_bias": ["model.layers.{layer_number}.mlp.gate.e_score_correction_bias"],
"shared_experts.linear_fc1.weight": [
"model.layers.{layer_number}.mlp.shared_experts.gate_proj.weight",
"model.layers.{layer_number}.mlp.shared_experts.up_proj.weight",
],
"shared_experts.linear_fc2.weight": ["model.layers.{layer_number}.mlp.shared_experts.down_proj.weight"],
}
_MTP_MAPPING = {
"enorm.weight": ["model.layers.{layer_number}.enorm.weight"],
"hnorm.weight": ["model.layers.{layer_number}.hnorm.weight"],
"eh_proj.weight": ["model.layers.{layer_number}.eh_proj.weight"],
"final_layernorm.weight": ["model.layers.{layer_number}.shared_head.norm.weight"],
}
def _weight_name_mapping_mtp(self, name: str, num_layers: int) -> str:
convert_names = []
for keyword, mapping_names in self._MTP_MAPPING.items():
if keyword in name:
convert_names.extend([x.format(layer_number=num_layers) for x in mapping_names])
break
elif "mlp" in name:
mtp_layer_index = int(re.findall(r"mtp\.layers\.(\d+)\.", name)[0])
name_ = re.sub(
r"^mtp\.layers.\d+.transformer_layer", f"model.layers.{num_layers+mtp_layer_index}", name
)
convert_names = self._weight_name_mapping_mlp(name_)
break
elif "self_attention" in name:
mtp_layer_index = int(re.findall(r"mtp\.layers.(\d+)\.", name)[0])
name_ = re.sub(
r"^mtp\.layers.\d+.transformer_layer", f"model.layers.{num_layers+mtp_layer_index}", name
)
convert_names = self._weight_name_mapping_attention(name_)
break
if len(convert_names) == 0:
raise NotImplementedError(f"Unsupported parameter name: {name}")
return convert_names
def _weight_name_mapping_mcore_to_hf(self, mcore_weights_name: str) -> list[str]:
"""
Map MCore weight names to Hugging Face weight names.
Args:
mcore_weights_name: MCore weight name
Returns:
list: Corresponding Hugging Face weight names
"""
assert "_extra_state" not in mcore_weights_name, "extra_state should not be loaded"
direct_name_mapping = {
"embedding.word_embeddings.weight": "model.embed_tokens.weight",
"decoder.final_layernorm.weight": "model.norm.weight",
"output_layer.weight": "lm_head.weight",
}
if mcore_weights_name in direct_name_mapping:
return [direct_name_mapping[mcore_weights_name]]
if "mtp" in mcore_weights_name: # first check mtp
return self._weight_name_mapping_mtp(mcore_weights_name, self.config.num_layers)
elif "self_attention" in mcore_weights_name:
return self._weight_name_mapping_attention(mcore_weights_name)
elif "mlp" in mcore_weights_name:
return self._weight_name_mapping_mlp(mcore_weights_name)
else:
raise NotImplementedError(f"Unsupported parameter name: {mcore_weights_name}")
def _build_config(self):
"""
Build the configuration for Qwen2 models.
Configures Qwen2-specific parameters such as QKV bias settings and
layer normalization options.
Returns:
TransformerConfig: Configuration object for Qwen2 models
"""
return self._build_base_config(
use_cpu_initialization=False,
# MoE specific
moe_ffn_hidden_size=self.hf_config.moe_intermediate_size,
moe_router_bias_update_rate=0.001,
moe_router_topk=self.hf_config.num_experts_per_tok,
num_moe_experts=self.hf_config.n_routed_experts,
# moe_router_load_balancing_type="aux_loss",
moe_router_load_balancing_type="none", # default None for RL
moe_grouped_gemm=True,
moe_router_score_function="sigmoid",
moe_router_enable_expert_bias=True,
moe_router_pre_softmax=True,
# Other optimizations
persist_layer_norm=True,
bias_activation_fusion=True,
bias_dropout_fusion=True,
# GLM specific
qk_layernorm=self.hf_config.use_qk_norm,
add_qkv_bias=True,
add_bias_linear=False,
# post_mlp_layernorm=True,
# post_self_attn_layernorm=True,
rotary_interleaved=True,
)

View File

@@ -0,0 +1,123 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec
from mbridge.core import register_model
from mbridge.models import Qwen2Bridge
@register_model("mimo")
class MimoBridge(Qwen2Bridge):
"""
Bridge implementation for Mimo models.
This class extends Qwen2Bridge to provide specific configurations and
optimizations for Mimo models, handling the conversion between
Hugging Face Mimo format and Megatron-Core.
MiMo adds MTP (Multi-Token Prediction) layers on top of Qwen2 architecture.
"""
def _build_config(self):
"""Override to add MTP configuration."""
hf_config = self.hf_config
# Add MTP configuration if present
mtp_args = {}
if "num_nextn_predict_layers" in hf_config:
mtp_args["mtp_num_layers"] = hf_config.num_nextn_predict_layers
return self._build_base_config(
add_qkv_bias=True,
qk_layernorm=False,
**mtp_args,
)
def _get_gptmodel_args(self) -> dict:
"""Override to add MTP block spec if needed."""
ret = super()._get_gptmodel_args()
# Add MTP block spec if MTP layers are present
if self.config.mtp_num_layers is not None:
transformer_layer_spec = self.config
mtp_block_spec = get_gpt_mtp_block_spec(self.config, transformer_layer_spec, use_transformer_engine=True)
ret["mtp_block_spec"] = mtp_block_spec
return ret
def _weight_name_mapping_mcore_to_hf(self, mcore_weights_name: str) -> list[str]:
"""Override to handle MTP layer mappings."""
# Check if this is an MTP layer weight
if "mtp" in mcore_weights_name:
return self._convert_mtp_param(mcore_weights_name)
# Otherwise use parent class mapping
return super()._weight_name_mapping_mcore_to_hf(mcore_weights_name)
def _convert_mtp_param(self, name: str) -> list[str]:
"""Convert MTP layer parameters from MCore to HF format."""
# For now, assume single MTP layer support
if "mtp.layers." not in name:
raise NotImplementedError(f"Invalid MTP parameter name: {name}")
# Get the MTP layer index
parts = name.split(".")
mtp_layer_idx = parts[2] # mtp.layers.{idx}
# Direct mappings for MTP-specific components
direct_name_mapping = {
f"mtp.layers.{mtp_layer_idx}.enorm.weight": f"model.mtp_layers.{mtp_layer_idx}.token_layernorm.weight",
f"mtp.layers.{mtp_layer_idx}.hnorm.weight": f"model.mtp_layers.{mtp_layer_idx}.hidden_layernorm.weight",
f"mtp.layers.{mtp_layer_idx}.eh_proj.weight": f"model.mtp_layers.{mtp_layer_idx}.input_proj.weight",
f"mtp.layers.{mtp_layer_idx}.final_layernorm.weight": f"model.mtp_layers.{mtp_layer_idx}.final_layernorm.weight",
}
if name in direct_name_mapping:
return [direct_name_mapping[name]]
# Handle transformer components within MTP
# Check if this is a transformer_layer component
if "transformer_layer" in name:
# Create a proxy name to use with parent class methods
# Convert mtp.layers.{idx}.transformer_layer.* to decoder.layers.{idx}.*
proxy_name = name.replace(
f"mtp.layers.{mtp_layer_idx}.transformer_layer",
f"decoder.layers.{mtp_layer_idx}",
)
if "self_attention" in proxy_name or "input_layernorm.weight" in proxy_name:
convert_names = super()._weight_name_mapping_attention(proxy_name)
elif "mlp" in proxy_name:
convert_names = super()._weight_name_mapping_mlp(proxy_name)
else:
raise NotImplementedError(f"Unsupported transformer component in MTP: {name}")
# Replace the layer index in converted names to point to mtp_layers
convert_names = [
cn.replace(f"model.layers.{mtp_layer_idx}", f"model.mtp_layers.{mtp_layer_idx}")
for cn in convert_names
]
return convert_names
else:
raise NotImplementedError(f"Unsupported MTP parameter name: {name}")
return convert_names
def _weight_to_mcore_format(self, mcore_weights_name: str, hf_weights: list[torch.Tensor]) -> torch.Tensor:
"""Swap halves of eh_proj weights before handing off to Megatron-Core."""
weight = super()._weight_to_mcore_format(mcore_weights_name, hf_weights)
if mcore_weights_name.endswith("eh_proj.weight"):
first_half, second_half = weight.chunk(2, dim=1)
weight = torch.cat([second_half, first_half], dim=1)
return weight
def _weight_to_hf_format(
self, mcore_weights_name: str, mcore_weights: torch.Tensor
) -> tuple[list[str], list[torch.Tensor]]:
"""Swap halves back when exporting eh_proj weights to HuggingFace format."""
if mcore_weights_name.endswith("eh_proj.weight"):
first_half, second_half = mcore_weights.chunk(2, dim=1)
mcore_weights = torch.cat([second_half, first_half], dim=1)
return super()._weight_to_hf_format(mcore_weights_name, mcore_weights)

View File

@@ -0,0 +1,104 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
from mbridge.core import register_model
from mbridge.models import Qwen2MoEBridge
@register_model("qwen3_next")
class Qwen3NextBridge(Qwen2MoEBridge):
_ATTENTION_MAPPING = (
Qwen2MoEBridge._ATTENTION_MAPPING
| {
f"self_attention.{weight_name}": ["model.layers.{layer_number}." + weight_name]
for weight_name in [
"input_layernorm.weight",
# linear attn
"linear_attn.A_log",
"linear_attn.conv1d.weight",
"linear_attn.dt_bias",
"linear_attn.in_proj_ba.weight",
"linear_attn.in_proj_qkvz.weight",
"linear_attn.norm.weight",
"linear_attn.out_proj.weight",
# gated attn
"self_attn.k_norm.weight",
"self_attn.k_proj.weight",
"self_attn.o_proj.weight",
"self_attn.q_norm.weight",
"self_attn.q_proj.weight",
"self_attn.v_proj.weight",
]
}
| {
"self_attention.linear_qkv.layer_norm_weight": ["model.layers.{layer_number}.input_layernorm.weight"],
"self_attention.linear_qkv.weight": [
"model.layers.{layer_number}.self_attn.q_proj.weight",
"model.layers.{layer_number}.self_attn.k_proj.weight",
"model.layers.{layer_number}.self_attn.v_proj.weight",
],
}
)
def _weight_to_mcore_format(
self, mcore_weights_name: str, hf_weights: list[torch.Tensor]
) -> tuple[list[str], list[torch.Tensor]]:
if "self_attention.linear_qkv." in mcore_weights_name and "layer_norm" not in mcore_weights_name:
# merge qkv
assert len(hf_weights) == 3
num_key_value_heads = self.hf_config.num_key_value_heads
hidden_dim = self.hf_config.hidden_size
num_attention_heads = self.hf_config.num_attention_heads
num_querys_per_group = num_attention_heads // self.hf_config.num_key_value_heads
head_dim = getattr(self.hf_config, "head_dim", hidden_dim // num_attention_heads)
group_dim = head_dim * num_attention_heads // num_key_value_heads
q, k, v = hf_weights
# q k v might be tp split
real_num_key_value_heads = q.shape[0] // (2 * group_dim)
q = (
q.view(
[
real_num_key_value_heads,
num_querys_per_group,
2,
head_dim,
-1,
]
)
.transpose(1, 2)
.flatten(1, 3)
)
k = k.view([real_num_key_value_heads, head_dim, -1])
v = v.view([real_num_key_value_heads, head_dim, -1])
out_shape = [-1, hidden_dim] if ".bias" not in mcore_weights_name else [-1]
qgkv = torch.cat([q, k, v], dim=1).view(*out_shape).contiguous()
return qgkv
return super()._weight_to_mcore_format(mcore_weights_name, hf_weights)
def _build_config(self):
return self._build_base_config(
use_cpu_initialization=False,
# MoE specific
moe_ffn_hidden_size=self.hf_config.moe_intermediate_size,
moe_router_bias_update_rate=0.001,
moe_router_topk=self.hf_config.num_experts_per_tok,
num_moe_experts=self.hf_config.num_experts,
moe_aux_loss_coeff=self.hf_config.router_aux_loss_coef,
# moe_router_load_balancing_type="aux_loss",
moe_router_load_balancing_type="none", # default None for RL
moe_grouped_gemm=True,
moe_router_score_function="softmax",
# Other optimizations
persist_layer_norm=True,
bias_activation_fusion=True,
bias_dropout_fusion=True,
# Qwen specific
moe_router_pre_softmax=False,
qk_layernorm=True,
# Qwen3 Next specific
attention_output_gate=True,
moe_shared_expert_gate=True,
)

View File

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

View File

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

View File

@@ -0,0 +1,17 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec
def get_glm_spec(args, config, vp_stage):
transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec(
num_experts=args.num_experts,
moe_grouped_gemm=args.moe_grouped_gemm,
qk_layernorm=args.qk_layernorm,
multi_latent_attention=args.multi_latent_attention,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
post_self_attn_layernorm=args.post_self_attn_layernorm,
post_mlp_layernorm=args.post_mlp_layernorm,
)
return transformer_layer_spec

View File

@@ -0,0 +1,118 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
import torch
import torch.distributed as dist
from megatron.core import mpu, tensor_parallel
from megatron.core.inference.contexts import BaseInferenceContext
from megatron.core.packed_seq_params import PackedSeqParams
from megatron.core.transformer.module import MegatronModule
from transformers import AutoConfig
class HuggingfaceAttention(MegatronModule, ABC):
"""Attention layer abstract class.
This layer only contains common modules required for the "self attn" and
"cross attn" specializations.
"""
def __init__(
self,
args,
config,
layer_number: int,
cp_comm_type: str = "p2p",
pg_collection=None,
):
super().__init__(config=config)
self.args = args
self.config = config
# Note that megatron layer_number starts at 1
self.layer_number = layer_number
self.hf_layer_idx = layer_number - 1
self.hf_config = AutoConfig.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
# hardcode to fa2 at the moment.
self.hf_config._attn_implementation = "flash_attention_2"
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
key_value_states: torch.Tensor | None = None,
inference_context: BaseInferenceContext | None = None,
rotary_pos_emb: torch.Tensor | tuple[torch.Tensor, torch.Tensor] | None = None,
rotary_pos_cos: torch.Tensor | None = None,
rotary_pos_sin: torch.Tensor | None = None,
rotary_pos_cos_sin: torch.Tensor | None = None,
attention_bias: torch.Tensor | None = None,
packed_seq_params: PackedSeqParams | None = None,
sequence_len_offset: int | None = None,
*,
inference_params: BaseInferenceContext | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
assert packed_seq_params is not None
cu_seqlens = packed_seq_params.cu_seqlens_q
if self.args.sequence_parallel:
hidden_states = tensor_parallel.gather_from_sequence_parallel_region(
hidden_states, group=mpu.get_tensor_model_parallel_group()
)
if mpu.get_context_parallel_world_size() > 1:
cp_size = mpu.get_context_parallel_world_size()
hidden_states_list = dist.nn.all_gather(
hidden_states,
group=mpu.get_context_parallel_group(),
)
# TODO: preprocess this for each batch to prevent tolist in the training step
whole_hidden_states_list = []
local_cu_seqlens = cu_seqlens // cp_size
for i in range(len(cu_seqlens) - 1):
seqlen = cu_seqlens[i + 1] - cu_seqlens[i]
chunk_size = seqlen // 2 // cp_size
whole_hidden_states_list.extend(
[
hidden_states_list[cp_rank][local_cu_seqlens[i] : local_cu_seqlens[i] + chunk_size]
for cp_rank in range(cp_size)
]
+ [
hidden_states_list[cp_rank][local_cu_seqlens[i] + chunk_size : local_cu_seqlens[i + 1]]
for cp_rank in range(cp_size)
][::-1],
)
hidden_states = torch.cat(whole_hidden_states_list, dim=0)
hidden_states = hidden_states.permute(1, 0, 2) # [bsz, seq_len, hidden_dim]
output = self.hf_forward(hidden_states, packed_seq_params)
bias = None
output = output.permute(1, 0, 2) # [seq_len, bsz, hidden_dim]
if mpu.get_context_parallel_world_size() > 1:
cp_rank = mpu.get_context_parallel_rank()
output_list = []
for i in range(len(cu_seqlens) - 1):
seqlen = cu_seqlens[i + 1] - cu_seqlens[i]
chunk_size = seqlen // 2 // cp_size
seq = output[cu_seqlens[i] : cu_seqlens[i + 1]]
chunks = torch.chunk(seq, 2 * cp_size, dim=0)
output_list.append(chunks[cp_rank])
output_list.append(chunks[2 * cp_size - 1 - cp_rank])
output = torch.cat(output_list, dim=0)
if self.args.sequence_parallel:
output = tensor_parallel.scatter_to_sequence_parallel_region(
output, group=mpu.get_tensor_model_parallel_group()
)
return output, bias
@abstractmethod
def hf_forward(self, hidden_states, packed_seq_params):
"""Huggingface forward function"""

View File

@@ -0,0 +1,229 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import copy
import torch
import torch.nn as nn
import torch.nn.functional as F
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_decoder_block_spec
from megatron.core.transformer.spec_utils import ModuleSpec
from megatron.core.transformer.transformer_block import get_num_layers_to_build
from megatron.core.transformer.transformer_layer import get_transformer_layer_offset
from transformers import AutoConfig
from transformers.activations import ACT2FN
try:
from fla.modules import FusedRMSNormGated, ShortConvolution
from fla.ops.gated_delta_rule import chunk_gated_delta_rule
from transformers.models.qwen3_next.modeling_qwen3_next import Qwen3NextAttention, Qwen3NextRMSNorm
except ImportError:
pass
from .hf_attention import HuggingfaceAttention
# adapt from https://github.com/huggingface/transformers/blob/38a08b6e8ae35857109cedad75377997fecbf9d0/src/transformers/models/qwen3_next/modeling_qwen3_next.py#L564
class Qwen3NextGatedDeltaNet(nn.Module):
"""
Qwen3NextGatedDeltaNet with varlen support
"""
def __init__(self, config, layer_idx: int):
super().__init__()
self.hidden_size = config.hidden_size
self.num_v_heads = config.linear_num_value_heads
self.num_k_heads = config.linear_num_key_heads
self.head_k_dim = config.linear_key_head_dim
self.head_v_dim = config.linear_value_head_dim
self.key_dim = self.head_k_dim * self.num_k_heads
self.value_dim = self.head_v_dim * self.num_v_heads
self.conv_kernel_size = config.linear_conv_kernel_dim
self.layer_idx = layer_idx
self.activation = config.hidden_act
self.act = ACT2FN[config.hidden_act]
self.layer_norm_epsilon = config.rms_norm_eps
# QKV
self.conv_dim = self.key_dim * 2 + self.value_dim
self.conv1d = ShortConvolution(
hidden_size=self.conv_dim,
bias=False,
kernel_size=self.conv_kernel_size,
)
# projection of the input hidden states
projection_size_qkvz = self.key_dim * 2 + self.value_dim * 2
projection_size_ba = self.num_v_heads * 2
self.in_proj_qkvz = nn.Linear(self.hidden_size, projection_size_qkvz, bias=False)
self.in_proj_ba = nn.Linear(self.hidden_size, projection_size_ba, bias=False)
# time step projection (discretization)
# instantiate once and copy inv_dt in init_weights of PretrainedModel
self.dt_bias = nn.Parameter(torch.ones(self.num_v_heads))
A = torch.empty(self.num_v_heads).uniform_(0, 16)
self.A_log = nn.Parameter(torch.log(A))
self.norm = FusedRMSNormGated(
self.head_v_dim,
eps=self.layer_norm_epsilon,
activation=self.activation,
device=torch.cuda.current_device(),
dtype=config.dtype if config.dtype is not None else torch.get_current_dtype(),
)
self.out_proj = nn.Linear(self.value_dim, self.hidden_size, bias=False)
def fix_query_key_value_ordering(self, mixed_qkvz, mixed_ba):
"""
Derives `query`, `key` and `value` tensors from `mixed_qkvz` and `mixed_ba`.
"""
new_tensor_shape_qkvz = mixed_qkvz.size()[:-1] + (
self.num_k_heads,
2 * self.head_k_dim + 2 * self.head_v_dim * self.num_v_heads // self.num_k_heads,
)
new_tensor_shape_ba = mixed_ba.size()[:-1] + (self.num_k_heads, 2 * self.num_v_heads // self.num_k_heads)
mixed_qkvz = mixed_qkvz.view(*new_tensor_shape_qkvz)
mixed_ba = mixed_ba.view(*new_tensor_shape_ba)
split_arg_list_qkvz = [
self.head_k_dim,
self.head_k_dim,
(self.num_v_heads // self.num_k_heads * self.head_v_dim),
(self.num_v_heads // self.num_k_heads * self.head_v_dim),
]
split_arg_list_ba = [self.num_v_heads // self.num_k_heads, self.num_v_heads // self.num_k_heads]
query, key, value, z = torch.split(mixed_qkvz, split_arg_list_qkvz, dim=3)
b, a = torch.split(mixed_ba, split_arg_list_ba, dim=3)
# [b, sq, ng, np/ng * hn] -> [b, sq, np, hn]
value = value.reshape(value.size(0), value.size(1), -1, self.head_v_dim)
z = z.reshape(z.size(0), z.size(1), -1, self.head_v_dim)
b = b.reshape(b.size(0), b.size(1), self.num_v_heads)
a = a.reshape(a.size(0), a.size(1), self.num_v_heads)
return query, key, value, z, b, a
def forward(
self,
hidden_states: torch.Tensor,
cu_seqlens: torch.Tensor = None,
):
projected_states_qkvz = self.in_proj_qkvz(hidden_states)
projected_states_ba = self.in_proj_ba(hidden_states)
query, key, value, z, b, a = self.fix_query_key_value_ordering(projected_states_qkvz, projected_states_ba)
query, key, value = (x.reshape(x.shape[0], x.shape[1], -1) for x in (query, key, value))
mixed_qkv = torch.cat((query, key, value), dim=-1)
mixed_qkv, _ = self.conv1d(
x=mixed_qkv,
cu_seqlens=cu_seqlens,
)
query, key, value = torch.split(
mixed_qkv,
[
self.key_dim,
self.key_dim,
self.value_dim,
],
dim=-1,
)
query = query.reshape(query.shape[0], query.shape[1], -1, self.head_k_dim)
key = key.reshape(key.shape[0], key.shape[1], -1, self.head_k_dim)
value = value.reshape(value.shape[0], value.shape[1], -1, self.head_v_dim)
beta = b.sigmoid()
# If the model is loaded in fp16, without the .float() here, A might be -inf
g = -self.A_log.float().exp() * F.softplus(a.float() + self.dt_bias)
if self.num_v_heads // self.num_k_heads > 1:
query = query.repeat_interleave(self.num_v_heads // self.num_k_heads, dim=2)
key = key.repeat_interleave(self.num_v_heads // self.num_k_heads, dim=2)
core_attn_out, last_recurrent_state = chunk_gated_delta_rule(
query,
key,
value,
g=g,
beta=beta,
initial_state=None,
output_final_state=False,
use_qk_l2norm_in_kernel=True,
)
z_shape_og = z.shape
# reshape input data into 2D tensor
core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1])
z = z.reshape(-1, z.shape[-1])
core_attn_out = self.norm(core_attn_out, z)
core_attn_out = core_attn_out.reshape(z_shape_og)
core_attn_out = core_attn_out.reshape(core_attn_out.shape[0], core_attn_out.shape[1], -1)
output = self.out_proj(core_attn_out)
return output
class Attention(HuggingfaceAttention):
def __init__(
self,
args,
config,
layer_number: int,
cp_comm_type: str = "p2p",
pg_collection=None,
):
super().__init__(
args,
config,
layer_number,
cp_comm_type,
pg_collection,
)
if Qwen3NextAttention is None:
raise ImportError("Please install transformers>=4.35.0 to use Qwen3NextAttention.")
self.linear_attn = Qwen3NextGatedDeltaNet(self.hf_config, self.hf_layer_idx)
self.input_layernorm = Qwen3NextRMSNorm(self.hf_config.hidden_size, eps=self.hf_config.rms_norm_eps)
def hf_forward(self, hidden_states, packed_seq_params):
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.linear_attn(
hidden_states=hidden_states,
cu_seqlens=packed_seq_params.cu_seqlens_q,
)
return hidden_states
def get_qwen3_next_spec(args, config, vp_stage):
# always use the moe path
if not args.num_experts:
config.moe_layer_freq = [0] * config.num_layers
# Define the decoder block spec
kwargs = {
"use_transformer_engine": True,
}
if vp_stage is not None:
kwargs["vp_stage"] = vp_stage
transformer_layer_spec = get_gpt_decoder_block_spec(config, **kwargs)
assert config.pipeline_model_parallel_layout is None, "not support this at the moment"
# Slice the layer specs to only include the layers that are built in this pipeline stage.
# Note: MCore layer_number starts at 1
num_layers_to_build = get_num_layers_to_build(config, vp_stage=vp_stage)
offset = get_transformer_layer_offset(config, vp_stage=vp_stage)
hf_config = AutoConfig.from_pretrained(args.hf_checkpoint, trust_remote_code=True)
for layer_id in range(num_layers_to_build):
if hf_config.layer_types[layer_id + offset] == "linear_attention":
layer_specs = copy.deepcopy(transformer_layer_spec.layer_specs[layer_id])
layer_specs.submodules.self_attention = ModuleSpec(
module=Attention,
params={"args": args},
)
transformer_layer_spec.layer_specs[layer_id] = layer_specs
return transformer_layer_spec

View File

@@ -0,0 +1,22 @@
accelerate
blobfile
datasets
httpx[http2]
mcp[cli]
megatron-bridge @ git+https://github.com/fzyzcjy/Megatron-Bridge.git@dev_rl
memray
nvidia-modelopt[torch]>=0.37.0
omegaconf
pillow
pylatexenc
pyyaml
ray[default]
ring_flash_attn
sglang-router>=0.2.3
tensorboard
transformers
wandb
liger_kernel
[fsdp]
torch>=2.0

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

View File

@@ -0,0 +1,2 @@
slime
slime_plugins