执行摘要
新增流式生成示例,支持 abort 时保留部分状态
每块 SSE 数据直接写入 sample,中断时已积累部分状态,不依赖 abort 请求返回的收集文本。
建议阅读 slime/rollout/sglang_streaming_rollout.py 中基于 base snapshot + chunk delta 的增量状态累加设计,该模式在需要中断恢复的场景下值得推广。
本案无审阅讨论,由 author zhuzilin 独立提交并合并到主分支。
每块 SSE 数据直接写入 sample,中断时已积累部分状态,不依赖 abort 请求返回的收集文本。
建议阅读 slime/rollout/sglang_streaming_rollout.py 中基于 base snapshot + chunk delta 的增量状态累加设计,该模式在需要中断恢复的场景下值得推广。
本案无审阅讨论,由 author zhuzilin 独立提交并合并到主分支。
slime/rollout/sglang_streaming_rollout.py,定义异步函数 generate_streaming:复用 GenerateState 和 _prepare_prompt_ids,构建 stream=True 的 HTTP 请求,通过 aiter_lines 循环解析 SSE 行,增量累积 call_tokens、call_log_probs、text,最后合并到 sample。tests/test_qwen3_4B_streaming_partial_rollout.py 集成测试:设置 --custom-generate-function-path 指向新函数,配置 --over-sampling-batch-size > --rollout-batch-size 并启用 --partial-rollout,确保每个 rollout 步骤触发 abort 以检验流式恢复。pr-test.yml 和 pr-test.yml.j2)的测试矩阵中新增该测试条目,分配 8 GPU 自动执行。| 文件 | 模块 | 状态 | 重要度 |
|---|---|---|---|
slime/rollout/sglang_streaming_rollout.py |
流式生成 | added | 8.08 |
tests/test_qwen3_4B_streaming_partial_rollout.py |
流式测试 | added | 7.11 |
.github/workflows/pr-test.yml |
CI 配置 | modified | 2.55 |
.github/workflows/pr-test.yml.j2 |
CI 配置 | modified | 2.24 |
tests/test_qwen3_4B_streaming_partial_rollout.py
test-coverage
CI 集成测试,验证 streaming + partial rollout 组合在 abort 场景下的正确性和奖励信号。
"""CI smoke test for the streaming sglang rollout path.
Wires slime.rollout.sglang_streaming_rollout.generate_streaming as the
per-sample generate function, with --over-sampling-batch-size >
--rollout-batch-size and --partial-rollout enabled so the rollout
loop must abort in-flight requests every step — exercising the streaming
abort path (partial state should already be on the sample when the SSE is
cut, then the partial groups get recycled into the data buffer).
Uses Qwen3-4B so responses on dapo-math are long enough to actually
trigger mid-stream aborts.
"""
import os
import slime.utils.external_utils.command_utils as U
TIGHT_HOST_MEMORY = U.get_bool_env_var("SLIME_TEST_TIGHT_HOST_MEMORY", "1")
MODEL_NAME = "Qwen3-4B"
MODEL_TYPE = "qwen3-4B"
NUM_GPUS = 8
def prepare():
# 准备模型权重、数据集和 checkpoint 转换
U.exec_command("mkdir -p /root/models /root/datasets")
U.exec_command(f"hf download Qwen/{MODEL_NAME} --local-dir /root/models/{MODEL_NAME}")
U.hf_download_dataset("zhuzilin/dapo-math-17k")
U.convert_checkpoint(model_name=MODEL_NAME, megatron_model_type=MODEL_TYPE, num_gpus_per_node=NUM_GPUS)
def execute():
ckpt_args = f"--hf-checkpoint /root/models/{MODEL_NAME}/ " f"--ref-load /root/{MODEL_NAME}_torch_dist "
rollout_args = (
# 使用流式生成函数作为每个样本的 generate 函数
"--custom-generate-function-path slime.rollout.sglang_streaming_rollout.generate_streaming "
"--prompt-data /root/datasets/dapo-math-17k/dapo-math-17k.jsonl "
"--input-key prompt "
"--label-key label "
"--apply-chat-template "
"--rollout-shuffle "
"--rm-type deepscaler "
"--num-rollout 2 "
"--rollout-batch-size 4 "
# over-sampling 2x 确保半数 in-flight 组必须 abort,partial-rollout 回收它们
"--over-sampling-batch-size 8 "
"--partial-rollout "
"--mask-offpolicy-in-partial-rollout "
"--n-samples-per-prompt 4 "
"--rollout-max-response-len 4096 "
"--rollout-temperature 0.8 "
"--global-batch-size 16 "
"--balance-data "
)
perf_args = (
"--tensor-model-parallel-size 2 "
"--sequence-parallel "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 2 "
"--recompute-granularity full "
"--recompute-method uniform "
"--recompute-num-layers 1 "
"--use-dynamic-batch-size "
f"--max-tokens-per-gpu {2048 if TIGHT_HOST_MEMORY else 8192} "
)
# ... 其他参数(优化器、sglang、CI 等)拼接后执行
# 完整配置见测试源文件
当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。
--incremental-streaming-output) 或修改 JSON 结构,该函数需适配。--custom-generate-function-path 显式启用,风险隔离。对现有用户无直接影响,需手动配置才启用。系统新增可选流式生成路径,无额外资源占用。为团队提供了 abort 场景下更可靠的样本状态管理示例,未来可能成为部分 rollout 的推荐配置。
当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。
参与讨论