初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD Source: Original Platform
This commit is contained in:
39
slime_plugins/PKG-INFO
Normal file
39
slime_plugins/PKG-INFO
Normal 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
140
slime_plugins/SOURCES.txt
Normal 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
|
||||
3
slime_plugins/__init__.py
Normal file
3
slime_plugins/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
BIN
slime_plugins/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime_plugins/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
1
slime_plugins/dependency_links.txt
Normal file
1
slime_plugins/dependency_links.txt
Normal file
@@ -0,0 +1 @@
|
||||
|
||||
9
slime_plugins/mbridge/__init__.py
Normal file
9
slime_plugins/mbridge/__init__.py
Normal 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"]
|
||||
BIN
slime_plugins/mbridge/__pycache__/__init__.cpython-312.pyc
Normal file
BIN
slime_plugins/mbridge/__pycache__/__init__.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime_plugins/mbridge/__pycache__/glm4.cpython-312.pyc
Normal file
BIN
slime_plugins/mbridge/__pycache__/glm4.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime_plugins/mbridge/__pycache__/glm4moe.cpython-312.pyc
Normal file
BIN
slime_plugins/mbridge/__pycache__/glm4moe.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime_plugins/mbridge/__pycache__/mimo.cpython-312.pyc
Normal file
BIN
slime_plugins/mbridge/__pycache__/mimo.cpython-312.pyc
Normal file
Binary file not shown.
BIN
slime_plugins/mbridge/__pycache__/qwen3_next.cpython-312.pyc
Normal file
BIN
slime_plugins/mbridge/__pycache__/qwen3_next.cpython-312.pyc
Normal file
Binary file not shown.
112
slime_plugins/mbridge/glm4.py
Normal file
112
slime_plugins/mbridge/glm4.py
Normal 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}")
|
||||
125
slime_plugins/mbridge/glm4moe.py
Normal file
125
slime_plugins/mbridge/glm4moe.py
Normal 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,
|
||||
)
|
||||
123
slime_plugins/mbridge/mimo.py
Normal file
123
slime_plugins/mbridge/mimo.py
Normal 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)
|
||||
104
slime_plugins/mbridge/qwen3_next.py
Normal file
104
slime_plugins/mbridge/qwen3_next.py
Normal 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,
|
||||
)
|
||||
3
slime_plugins/megatron_bridge/__init__.py
Normal file
3
slime_plugins/megatron_bridge/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
3
slime_plugins/models/__init__.py
Normal file
3
slime_plugins/models/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
17
slime_plugins/models/glm4.py
Normal file
17
slime_plugins/models/glm4.py
Normal 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
|
||||
118
slime_plugins/models/hf_attention.py
Normal file
118
slime_plugins/models/hf_attention.py
Normal file
@@ -0,0 +1,118 @@
|
||||
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
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"""
|
||||
229
slime_plugins/models/qwen3_next.py
Normal file
229
slime_plugins/models/qwen3_next.py
Normal 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
|
||||
22
slime_plugins/requires.txt
Normal file
22
slime_plugins/requires.txt
Normal 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
|
||||
50
slime_plugins/rollout_buffer/README.md
Normal file
50
slime_plugins/rollout_buffer/README.md
Normal 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
|
||||
```
|
||||
51
slime_plugins/rollout_buffer/README_zh.md
Normal file
51
slime_plugins/rollout_buffer/README_zh.md
Normal 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
|
||||
```
|
||||
343
slime_plugins/rollout_buffer/buffer.py
Normal file
343
slime_plugins/rollout_buffer/buffer.py
Normal 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,
|
||||
)
|
||||
9
slime_plugins/rollout_buffer/generator/__init__.py
Normal file
9
slime_plugins/rollout_buffer/generator/__init__.py
Normal 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",
|
||||
]
|
||||
354
slime_plugins/rollout_buffer/generator/base_generator.py
Normal file
354
slime_plugins/rollout_buffer/generator/base_generator.py
Normal 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
|
||||
310
slime_plugins/rollout_buffer/rollout_buffer_example.py
Normal file
310
slime_plugins/rollout_buffer/rollout_buffer_example.py
Normal 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))
|
||||
137
slime_plugins/rollout_buffer/rollout_buffer_example.sh
Normal file
137
slime_plugins/rollout_buffer/rollout_buffer_example.sh
Normal 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
|
||||
2
slime_plugins/top_level.txt
Normal file
2
slime_plugins/top_level.txt
Normal file
@@ -0,0 +1,2 @@
|
||||
slime
|
||||
slime_plugins
|
||||
Reference in New Issue
Block a user