modelscope / modelscope/ms-swift
GRPO多轮训练,rollout成功拉起,但是训练端卡死。(训练异常的慢)
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 15.7k
- Forks
- 1.7k
- Avg merge
- 1d 16h
- Merged PRs (30d)
- 136
Description
Checklist / 检查清单
- I have searched existing issues, and this is a new bug report. / 我已经搜索过现有的 issues,确认这是一个新的 bug report。
Bug Description / Bug 描述
在使用GRPO训练qwen3omni的多轮音频理解能力的时候,rollout成功拉起并且完成一轮rollout,但是训练端会在forward或者完成一次forward后的backward卡死。
首先尝试了使用deepspeed zero3,卡在了get_audio_feature的forward。grpo节点显示gpu利用率全部100,rollout端全部是0. (后来得知zero3在处理多模态数据会卡住,下边尝试了其他方案)
grpo端卡在
rollout卡在
然后尝试切换到zero2,在把batchszie和generatenum以及max_len开小之后,可以多打印几次日志,但是还是卡死
grpo端日志是
rollout端是
堆栈信息看是卡在了backward
========== PID 5986 ==========
Process 5986: /mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/bin/python3.10 /mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/swift/cli/rlhf.py --rlhf_type grpo --model /mnt/data/share-oss/user/shaomingchen/ckpt/long-audio/time-aware/ceshi/exp6-agentic-TAG/v1-20260207-183313/checkpoint-900 --template qwen3_omni --agent_template hermes --freeze_vit false --gradient_checkpointing true --vit_gradient_checkpointing false --dataloader_num_workers 8 --dataset /mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/data/agentic_data/TAG-RL/rl_selected_for_sft.jsonl --val_dataset /mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/data/agentic_data/TAG-RL/dev.jsonl --external_plugins /mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/my-RL/exp9_plugin.py --reward_funcs format_reward_func iou_reward_func think_content_reward_func --max_turns 5 --num_generations 8 --per_device_train_batch_size 1 --gradient_accumulation_steps 4 --max_length 32768 --use_vllm true --vllm_mode server --vllm_server_host 10.34.87.97 --vllm_server_port 8866 --vllm_server_pass_dataset true --sleep_level 0 --vllm_server_timeout 120 --vllm_server_group_port 32145 --vllm_use_async_engine false --deepspeed zero2 --tuner_type lora --bf16 true --attn_impl flash_attention_2 --padding_free true --output_dir /mnt/data/share-oss/user/shaomingchen/ckpt/long-audio/time-aware/ceshi/exp9-fsdp-grpo-rollout
Python v3.10.19 (/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/bin/python3.10)
Thread 5986 (idle): "MainThread"
_engine_run_backward (torch/autograd/graph.py:841)
backward (torch/autograd/init.py:354)
backward (torch/_tensor.py:625)
backward (deepspeed/runtime/engine.py:2532)
wrapped_fn (deepspeed/utils/nvtx.py:20)
backward (accelerate/utils/deepspeed.py:270)
backward (accelerate/accelerator.py:2844)
training_step (transformers/trainer.py:4071)
training_step (swift/rlhf_trainers/grpo_trainer.py:1859)
_inner_training_loop (transformers/trainer.py:2674)
train (transformers/trainer.py:2325)
train (swift/trainers/mixin.py:936)
train (swift/pipelines/train/sft.py:264)
run (swift/pipelines/train/sft.py:198)
wrapper (swift/ray/base.py:169)
main (swift/pipelines/base.py:47)
rlhf_main (swift/pipelines/train/rlhf.py:240)
(swift/cli/rlhf.py:7)
另外还试了fsdp2,直接切换会报错
依旧报错
File "/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/swift/pipelines/train/sft.py", line 264, in train
[rank2]: trainer.train(resume_checkpoint)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/swift/trainers/mixin.py", line 936, in train
[rank2]: res = super().train(*args, **kwargs)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/transformers/trainer.py", line 2325, in train
[rank2]: return inner_training_loop(
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/transformers/trainer.py", line 2480, in _inner_training_loop
[rank2]: model, self.optimizer = self.accelerator.prepare(self.model, self.optimizer)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/accelerate/accelerator.py", line 1555, in prepare
[rank2]: result = self._prepare_fsdp2(*args)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/accelerate/accelerator.py", line 1711, in _prepare_fsdp2
[rank2]: model = fsdp2_prepare_model(self, model)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/accelerate/utils/fsdp_utils.py", line 675, in fsdp2_prepare_model
[rank2]: fully_shard(module, **fsdp2_kwargs)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/torch/distributed/_composable/contract.py", line 150, in wrapper
[rank2]: updated = func(inp_module, *args, **kwargs)
[rank2]: File "/mnt/data/share-ssd/user/shaomingchen/miniconda3/envs/swift/lib/python3.10/site-packages/torch/distributed/fsdp/_fully_shard/_fully_shard.py", line 190, in fully_shard
[rank2]: raise ValueError(
[rank2]: ValueError: fully_shard does not support containers that do not implement forward: ModuleList(
[rank2]: (0-47): 48 x CheckpointWrapper(
[rank2]: (_checkpoint_wrapped_module): Qwen3OmniMoeThinkerTextDecoderLayer(
[rank2]: (self_attn): Qwen3OmniMoeThinkerTextAttention(
[rank2]: (q_proj): lora.Linear(
[rank2]: (base_layer): Linear(in_features=2048, out_features=4096, bias=False)
[rank2]: (lora_dropout): ModuleDict(
[rank2]: (default): Dropout(p=0.05, inplace=False)
[rank2]: )
[rank2]: (lora_A): ModuleDict(
[rank2]: (default): Linear(in_features=2048, out_features=8, bias=False)
[rank2]: )
[rank2]: (lora_B): ModuleDict(
[rank2]: (default): Linear(in_features=8, out_features=4096, bias=False)
[rank2]: )
[rank2]: (lora_embedding_A): ParameterDict()
[rank2]: (lora_embedding_B): ParameterDict()
[rank2]: (lora_magnitude_vector): ModuleDict()
[rank2]: )
How to Reproduce / 如何复现
我的rollout脚本是
export CUDA_VISIBLE_DEVICES=0,1,2,3,4,5,6,7
export VLLM_SERVER_GROUP_HOST=0.0.0.0
export VLLM_SERVER_GROUP_PORT=33388
SWIFT_WORKSPACE="/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9"
export PYTHONPATH="${SWIFT_WORKSPACE}:${PYTHONPATH}"
export VLLM_ATTENTION_BACKEND=FLASH_ATTN
export LD_LIBRARY_PATH="/mnt/data/share-ks3/shaomingchen/qh_transf/pure_compiler:$LD_LIBRARY_PATH"
export NCCL_IB_DISABLE=1
export NCCL_P2P_DISABLE=1
export NCCL_DEBUG=INFO
swift rollout \
--model "/mnt/data/share-oss/user/shaomingchen/ckpt/long-audio/time-aware/ceshi/exp6-agentic-TAG/v1-20260207-183313/checkpoint-900" \
--vllm_tensor_parallel_size 8 \
--vllm_max_model_len 32768 \
--vllm_gpu_memory_utilization 0.9 \
--multi_turn_scheduler AudioAgentScheduler \
--vllm_limit_mm_per_prompt '{"audio": 6}' \
--vllm_enforce_eager false \
--vllm_engine_kwargs '{"attention_config": {"backend": "FLASH_ATTN"}}' \
--external_plugins "${SWIFT_WORKSPACE}/examples/train/multimodal/my-RL/exp9_plugin.py" \
--port 9999 \
--vllm_use_async_engine true \
--host 0.0.0.0
grpo脚本是
#!/bin/bash
export DEBUG_LOG_1CF970=/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/debug-chaoshi.log
export TRITON_CACHE_DIR="/tmp/triton_cache_$(whoami)_node_${RANK}"
export TORCHINDUCTOR_CACHE_DIR="/tmp/torch_inductor_cache_$(whoami)_node_${RANK}"
# export NCCL_DEBUG=INFO
# export NCCL_DEBUG_SUBSYS=ALL
LOG_DIR="/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/logs"
mkdir -p "$LOG_DIR"
CURRENT_TIME=$(date "+%Y%m%d_%H%M%S")
LOG_FILE="${LOG_DIR}/${CURRENT_TIME}_grpo_rollout_node_${RANK}.log"
SWIFT_WORKSPACE="/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9"
MY_RL_PATH="${SWIFT_WORKSPACE}/examples/train/multimodal/my-RL"
export PYTHONPATH="${MY_RL_PATH}:${SWIFT_WORKSPACE}:${MEGATRON_LM_PATH}:${PYTHONPATH}"
export MS_OFFLINE=1
export HF_HUB_OFFLINE=1
export PYTORCH_CUDA_ALLOC_CONF='expandable_segments:True'
MODEL_PATH="/mnt/data/share-oss/user/shaomingchen/ckpt/long-audio/time-aware/ceshi/exp6-agentic-TAG/v1-20260207-183313/checkpoint-900"
TRAIN_DATASET="/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/data/agentic_data/TAG-RL/rl_selected_for_sft.jsonl"
VAL_DATASET="/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/ms-swift-2-9/examples/train/multimodal/data/agentic_data/TAG-RL/dev.jsonl"
OUTPUT_DIR="/mnt/data/share-oss/user/shaomingchen/ckpt/long-audio/time-aware/ceshi/exp9-fsdp-grpo-rollout"
export VLLM_SERVER_GROUP_HOST=10.34.50.242
export VLLM_SERVER_GROUP_PORT=33388
ROLLOUT_HOST="10.34.50.242"
ROLLOUT_PORT=9999
export MASTER_ADDR=127.0.0.1
export MASTER_PORT=29501
torchrun \
--nproc_per_node=8 \
--nnodes=1 \
--node_rank=0 \
--master_addr=$MASTER_ADDR \
--master_port=$MASTER_PORT \
$(which swift) rlhf \
\
--rlhf_type grpo \
--model "$MODEL_PATH" \
--template qwen3_omni \
--agent_template hermes \
--freeze_vit false \
\
--gradient_checkpointing true \
--vit_gradient_checkpointing false \
--dataloader_num_workers 8 \
\
--dataset "$TRAIN_DATASET" \
--val_dataset "$VAL_DATASET" \
\
--external_plugins "${MY_RL_PATH}/exp9_plugin.py" \
\
--reward_funcs format_reward_func iou_reward_func think_content_reward_func \
\
--max_turns 5 \
\
--num_generations 4 \
--per_device_train_batch_size 1 \
--gradient_accumulation_steps 1 \
--max_length 4096 \
\
--use_vllm true \
--vllm_mode server \
--vllm_server_host $ROLLOUT_HOST \
--vllm_server_port $ROLLOUT_PORT \
--vllm_server_pass_dataset true \
--sleep_level 0 \
--vllm_server_timeout 120 \
--vllm_server_group_port 33388 \
--vllm_use_async_engine false \
\
--deepspeed zero2 \
--tuner_type lora \
\
--bf16 true \
--attn_impl flash_attention_2 \
--padding_free true \
--output_dir "$OUTPUT_DIR" \
\
2>&1 | tee "$LOG_FILE"
# --deepspeed "${MY_RL_PATH}/exp9_deepspeed.json" \
# --multi_turn_scheduler AudioAgentScheduler \
# --deepspeed zero3 \ --fsdp "${MY_RL_PATH}/fsdp2_size_wrap.json" \
# --sequence_parallel_size 8 \
plugin代码是
import re
import json
import sys
import copy
import os
import uuid
import librosa
import soundfile as sf
from typing import Dict, List, Any, Optional
# ================= 1. 环境与路径配置 =================
# 务必确保该路径在所有机器(Node 0-3)上都能读写
SHARED_TMP_ROOT = "/mnt/data/share-ssd/user/shaomingchen/code/long-audio/time-aware/ft_omni/tmp_audio_crops"
os.makedirs(SHARED_TMP_ROOT, exist_ok=True)
from swift.infer_engine.protocol import (
ChatCompletionResponseChoice,
RolloutInferRequest
)
from swift.rollout.multi_turn import MultiTurnScheduler, multi_turns
from swift.rewards import ORM, orms
# #region agent log
def _dbg_log(message: str, data: dict, hypothesis_id: str = "plugin"):
try:
log_path = os.environ.get("DEBUG_LOG_1CF970")
if not log_path:
log_path = os.path.join(os.getcwd(), "debug-1cf970.log")
else:
log_path = os.path.abspath(log_path)
with open(log_path, "a", encoding="utf-8") as f:
f.write(json.dumps({"sessionId": "1cf970", "hypothesisId": hypothesis_id, "location": "exp9_plugin.py", "message": message, "data": data, "timestamp": __import__("time").time() * 1000}, ensure_ascii=False) + "\n")
except Exception:
pass
def _safe_path(p):
if not p or not isinstance(p, str):
return p
return os.path.basename(p)
def _trunc(s, max_len=120):
if s is None:
return None
s = str(s)
return s if len(s) <= max_len else s[:max_len] + "..."
def _step_trajectory_log(current_turn: int, raw_content: str, clean_assistant_content: str, call_info: Optional[Dict],
obs_content: str, next_messages: List[Dict], next_audios: List, state: Dict[str, Any]):
"""每一步完整轨迹写入日志:本轮 response、工具调用、工具结果、messages 快照、audios、state。"""
# AIGC START
try:
has_tool = call_info is not None
has_final = bool(re.search(r'\[\d{2}:\d{2}\s*[-——]\s*\d{2}:\d{2}\]', raw_content or ""))
messages_snapshot = []
for m in (next_messages or []):
role = m.get("role", "")
content = m.get("content", "")
if isinstance(content, dict):
content = f"<dict len token_ids={len(content.get('token_ids', []))}>"
else:
content = _trunc(str(content), 80)
messages_snapshot.append({"role": role, "content_preview": content})
data = {
"current_turn": current_turn,
"response": {
"content_len": len(raw_content or ""),
"clean_len": len(clean_assistant_content or ""),
"has_tool_call": has_tool,
"has_final_pattern": has_final,
"content_preview": _trunc(raw_content or "", 200),
},
"tool_call": call_info if call_info is not None else None,
"tool_result_preview": _trunc(obs_content, 200),
"messages_after": messages_snapshot,
"audios_after": [_safe_path(p) for p in (next_audios or [])],
"state_after": dict(state),
}
_dbg_log("step_trajectory", data, "rollout_detail")
except Exception:
pass
# AIGC END
# #endregion
# ================= 2. AudioAgentScheduler (调度引擎) =================
# AIGC START
# 必须把 infer_engine、tokenizer 等通过 **kwargs 传给父类,否则 self.infer_engine 为 None,导致 collective_rpc 解析失败、GRPO 连 32145 超时
# rollout 侧可能传入 max_turns=None,需在此处兜底为 5,否则 check_finished 里 current_turn >= self.max_turns 会报 TypeError
# AIGC END
class AudioAgentScheduler(MultiTurnScheduler):
def __init__(self, *args, max_turns: int = 5, max_tool_calls: int = 2, **kwargs):
if max_turns is None:
max_turns = 5
super().__init__(*args, max_turns=max_turns, **kwargs)
self.max_tool_calls = max_tool_calls
self.final_pattern = r'\[\d{2}:\d{2}\s*[-——]\s*\d{2}:\d{2}\]'
def _get_state(self, req: RolloutInferRequest) -> Dict[str, Any]:
"""从 data_dict 中维护 Agent 状态"""
if not hasattr(req, "data_dict") or req.data_dict is None:
req.data_dict = {}
# 记录已调用的工具次数和裁剪历史
return req.data_dict.setdefault("audio_agent_state", {"tool_calls": 0, "used_crops": []})
def _parse_first_tool(self, content: str) -> Optional[Dict]:
"""解析模型生成的第一个 tool_call 标签"""
match = re.search(r'<tool_call>(.*?)</tool_call>', content, re.DOTALL)
if match:
try:
raw_json = match.group(1).replace("```json", "").replace("```", "").strip()
return json.loads(raw_json)
except:
pass
return None
def _do_physical_crop(self, source_path: str, start: float, end: float) -> str:
"""执行音频物理裁剪并保存至共享盘"""
# #region agent log
_dbg_log("_do_physical_crop entry", {"source": _safe_path(source_path), "start_sec": start, "end_sec": end}, "rollout_detail")
# #endregion
try:
duration = max(0.5, end - start)
# sr=None 保持原始采样率,offset/duration 单位为秒
y, sr = librosa.load(source_path, sr=None, offset=start, duration=duration)
out_filename = f"crop_{uuid.uuid4().hex[:12]}.wav"
out_path = os.path.join(SHARED_TMP_ROOT, out_filename)
sf.write(out_path, y, sr)
# #region agent log
_dbg_log("_do_physical_crop ok", {"out": _safe_path(out_path), "samples": len(y), "sr": sr}, "rollout_detail")
# #endregion
return out_path
except Exception as e:
# #region agent log
_dbg_log("_do_physical_crop fail", {"error": str(e), "source": _safe_path(source_path), "start": start, "end": end}, "rollout_detail")
# #endregion
print(f"❌ Audio Crop Error: {e}")
return source_path
def check_finished(self, infer_request, response_choice, current_turn) -> bool:
"""判定多轮 rollout 是否终止"""
# AIGC START
# 防御:self.max_turns 可能为 None(父类允许),避免 int >= NoneType 报错
# AIGC END
if (self.max_turns is not None and current_turn >= self.max_turns) or not response_choice:
# #region agent log
_dbg_log("check_finished=True", {"reason": "max_turns_or_no_choice", "current_turn": current_turn, "max_turns": self.max_turns}, "H1")
# #endregion
return True
content = str(response_choice.message.content or "")
state = self._get_state(infer_request)
# 条件1:模型给出了最终格式的答案 [00:10 - 00:20]
if re.search(self.final_pattern, content):
# #region agent log
_dbg_log("check_finished=True", {"reason": "final_pattern", "current_turn": current_turn, "content_len": len(content), "content_tail": _trunc(content[-200:], 100), "state": state}, "H1")
_dbg_log("turn_final_response", {"current_turn": current_turn, "reason": "final_pattern", "response_preview": _trunc(content, 300), "state": state}, "rollout_detail")
# #endregion
return True
# 条件2:没有工具调用了或达到工具调用上限
call = self._parse_first_tool(content)
if not call or state["tool_calls"] >= self.max_tool_calls:
# #region agent log
_dbg_log("check_finished=True", {"reason": "no_tool_or_limit", "current_turn": current_turn, "tool_calls": state["tool_calls"], "max_tool_calls": self.max_tool_calls, "has_parsed_call": call is not None, "content_len": len(content)}, "H1")
_dbg_log("turn_final_response", {"current_turn": current_turn, "reason": "no_tool_or_limit", "response_preview": _trunc(content, 300), "state": state}, "rollout_detail")
# #endregion
return True
# #region agent log
_dbg_log("check_finished=False", {"current_turn": current_turn, "reason": "continue_tool", "content_len": len(content), "content_tail": _trunc(content[-150:], 80), "state": state}, "H1")
# #endregion
return False
def step(self, infer_request: RolloutInferRequest, response_choice, current_turn) -> Dict:
"""处理 Tool Call -> 执行工具 -> 构造下一轮请求"""
state = self._get_state(infer_request)
raw_content = str(response_choice.message.content or "")
# #region agent log
token_ids = getattr(response_choice, "token_ids", None) or []
has_lp = bool(getattr(response_choice, "logprobs", None) and isinstance(getattr(response_choice, "logprobs"), dict) and (response_choice.logprobs or {}).get("content"))
audios_in = getattr(infer_request, "audios", None) or []
_dbg_log("step entry", {
"current_turn": current_turn,
"len_messages_in": len(infer_request.messages),
"len_audios_in": len(audios_in),
"audios_in_basename": [_safe_path(a) for a in audios_in[:5]],
"state_tool_calls": state["tool_calls"],
"state_used_crops": state["used_crops"],
"raw_content_len": len(raw_content),
"raw_content_tail": _trunc(raw_content[-200:], 100),
"len_token_ids": len(token_ids),
"has_logprobs": has_lp,
"has_tokenizer": self.tokenizer is not None,
}, "H2")
# #endregion
# --- 熔断:截断模型在生成工具调用后的废话 ---
tc_end_pos = raw_content.find("</tool_call>")
clean_assistant_content = raw_content[:tc_end_pos + 12] if tc_end_pos != -1 else raw_content
call_info = self._parse_first_tool(clean_assistant_content)
next_audios = copy.deepcopy(infer_request.audios)
if call_info:
args = call_info.get("arguments", {})
start = float(args.get("start_sec", 0))
end = float(args.get("end_sec", 0))
# 物理切片:推理端执行裁剪
source_path = infer_request.audios[0] # 始终以原始长音频为基准切片
new_audio_path = self._do_physical_crop(source_path, start, end)
# 更新音频列表,确保文本中的 <audio> 标签能匹配到这个新路径
next_audios.append(new_audio_path)
state["tool_calls"] += 1
state["used_crops"].append(f"{start}_{end}")
# 构造工具观测值,放入 <audio> 标签引导多模态模型加载
obs_content = json.dumps({"result": f"已成功提取 {start}s 到 {end}s 片段:<audio>"}, ensure_ascii=False)
# #region agent log
_dbg_log("step after_crop", {
"current_turn": current_turn,
"parsed_start_sec": start,
"parsed_end_sec": end,
"source_basename": _safe_path(source_path),
"new_audio_basename": _safe_path(new_audio_path),
"next_audios_len": len(next_audios),
"next_audios_basename": [_safe_path(p) for p in next_audios],
"state_tool_calls": state["tool_calls"],
"state_used_crops": state["used_crops"],
"obs_content": _trunc(obs_content, 80),
}, "rollout_detail")
# #endregion
else:
obs_content = "Format error: Please use <tool_call> tags."
# #region agent log
_dbg_log("step no_tool_call", {"current_turn": current_turn, "clean_content_len": len(clean_assistant_content), "clean_tail": _trunc(clean_assistant_content[-150:], 80)}, "rollout_detail")
# #endregion
# 构造对话历史(框架在 step() 前已把本轮 completion 写入最后一条 assistant,这里替换为截断内容并追加 tool,避免连续两条 assistant 导致 template 配对 assert)
# AIGC START
next_messages = copy.deepcopy(infer_request.messages)
if next_messages and next_messages[-1].get("role") == "assistant":
next_messages[-1]["content"] = clean_assistant_content
else:
next_messages.append({"role": "assistant", "content": clean_assistant_content})
next_messages.append({"role": "tool", "content": obs_content})
# AIGC END
# #region agent log — 每一步完整轨迹:response / tool_call / tool_result / messages 快照 / audios / state
_step_trajectory_log(
current_turn, raw_content, clean_assistant_content, call_info, obs_content,
next_messages, next_audios, state,
)
# #endregion
# AIGC START
# 与多轮文档一致:返回 response_token_ids / response_loss_mask / rollout_logprobs(与截断内容对齐),
# 以及 rollout_infos['audios'],避免 completion 与 logprobs 数量不一致导致 collective 异常。
# #region agent log
_token_ids_for_log = getattr(response_choice, "token_ids", None) or []
_lp_content = (getattr(response_choice, "logprobs", None) or {}) if getattr(response_choice, "logprobs", None) else {}
if isinstance(_lp_content, dict):
_lp_content = _lp_content.get("content") or []
else:
_lp_content = []
# #endregion
result = {
"infer_request": RolloutInferRequest(
messages=next_messages,
tools=infer_request.tools,
audios=next_audios,
data_dict=infer_request.data_dict
),
"rollout_infos": {"audios": next_audios},
}
token_ids = _token_ids_for_log
logprobs_content = _lp_content if isinstance(_lp_content, list) else None
n = None # 用于对齐截断长度,未对齐时为 None
if token_ids and self.tokenizer and clean_assistant_content:
# 找到与 clean_assistant_content 对齐的 token 数量 N
for k in range(1, len(token_ids) + 1):
try:
decoded = self.tokenizer.decode(token_ids[:k], skip_special_tokens=True)
if decoded == clean_assistant_content or decoded.rstrip() == clean_assistant_content.rstrip():
n = k
break
except Exception:
continue
if n is not None and n > 0:
result["response_token_ids"] = list(token_ids[:n])
result["response_loss_mask"] = [1] * n
if logprobs_content and len(logprobs_content) >= n:
result["rollout_logprobs"] = [item.get("logprob") for item in logprobs_content[:n]]
elif logprobs_content:
result["rollout_logprobs"] = [item.get("logprob") for item in logprobs_content]
# 无法对齐时不返回 token_ids/logprobs,避免与截断后 messages 编码长度不一致
# #region agent log
return_keys = list(result.keys())
n_tok = len(result.get("response_token_ids", []))
n_lp = len(result.get("rollout_logprobs", []))
next_msgs = result["infer_request"].messages
_dbg_log("step return", {
"current_turn": current_turn,
"return_keys": return_keys,
"n_aligned": n,
"len_response_token_ids": n_tok,
"len_rollout_logprobs": n_lp,
"len_audios": len(next_audios),
"next_audios_basename": [_safe_path(p) for p in next_audios],
"next_messages_len": len(next_msgs),
"next_messages_roles": [m.get("role") for m in next_msgs],
"last_assistant_content_len": len(next_msgs[-2]["content"]) if len(next_msgs) >= 2 and next_msgs[-2].get("role") == "assistant" else None,
"last_tool_content": _trunc(next_msgs[-1]["content"], 60) if next_msgs and next_msgs[-1].get("role") == "tool" else None,
"state_tool_calls": state["tool_calls"],
"state_used_crops": state["used_crops"],
}, "H2")
# #endregion
# AIGC END
return result
# 注册调度器
multi_turns["AudioAgentScheduler"] = AudioAgentScheduler
# ================= 3. Reward Functions (三向全能奖励) =================
class AudioFormatORM(ORM):
"""格式奖励:必须包含思考和标准时间戳"""
def __call__(self, completions, **kwargs) -> List[float]:
pattern = r"<think>.*?</think>.*?\[\d{2}:\d{2}\s*[-——]\s*\d{2}:\d{2}\]"
return [1.0 if re.search(pattern, str(c), re.DOTALL) else 0.0 for c in completions]
class AudioIoUORM(ORM):
"""准确度奖励:计算预测时间戳与 GT 的 IoU"""
def __call__(self, completions, ground_truth, **kwargs) -> List[float]:
rewards = []
pattern = r"\[(\d+):(\d+)\s*[-——]\s*(\d+):(\d+)\]"
def parse_to_seconds(text):
matches = re.findall(pattern, str(text))
if not matches: return None
t = matches[-1] # 取最后输出的答案
return int(t[0])*60 + int(t[1]), int(t[2])*60 + int(t[3])
for c, gt in zip(completions, ground_truth):
p, g = parse_to_seconds(c), parse_to_seconds(gt)
if not p or not g:
rewards.append(0.0); continue
inter = max(0, min(p[1], g[1]) - max(p[0], g[0]))
union = (p[1] - p[0]) + (g[1] - g[0]) - inter
iou = inter/union if union > 0 else 0.0
# 阶梯奖励
if iou > 0.8: r = 2.0
elif iou > 0.4: r = 1.0
elif iou > 0.1: r = 0.5
else: r = 0.0
rewards.append(r)
return rewards
class AudioThinkContentORM(ORM):
"""质量奖励:评估思考过程的深度与防复读"""
def __call__(self, completions, **kwargs) -> List[float]:
rewards = []
for content in completions:
c_str = str(content)
think = re.search(r"<think>(.*?)</think>", c_str, re.DOTALL)
if not think:
rewards.append(0.0); continue
text = think.group(1).strip()
length = len(text)
# 基础分
score = 0.1
if 150 <= length <= 800: score = 0.5 # 理想长度
elif length > 800: score = 0.2 # 太啰嗦扣分
# 复读机检查 (词熵)
words = text.split()
if len(words) > 10 and (len(set(words)) / len(words)) < 0.35:
score = 0.0
rewards.append(score)
return rewards
# 注册奖励函数
orms["format_reward_func"] = AudioFormatORM
orms["iou_reward_func"] = AudioIoUORM
orms["think_content_reward_func"] = AudioThinkContentORM
print("✅ [Plugin] Multi-turn Audio Agent Plugin loaded successfully.")
数据样例是这样的
{"tools": "[{\"type\": \"function\", \"function\": {\"name\": \"crop_audio\", \"description\": \"裁剪音频\", \"parameters\": {\"type\": \"object\", \"properties\": {\"start_sec\": {\"type\": \"number\"}, \"end_sec\": {\"type\": \"number\"}}, \"required\": [\"start_sec\", \"end_sec\"]}}}]", "audios": ["/mnt/data/share-ks3/shaomingchen/long-audio-data/agent-SFT/audio_cuts/audio_cuts/s4_01298_B5-15_S0.mp3"], "messages": [{"role": "user", "content": "<audio>\n请仔细聆听音频,并定位出庄园主一家在餐桌上首次将混血儿比作愚蠢且无法生育的骡子,以此来讨论其种族优越性的具体片段。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。"}], "ground_truth": "[01:25 - 01:50]"}
Additional Information / 补充信息
在等待两个小时后打印了一步的训练日志。
[DEBUG _get_std_messages] n=5 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 134, 'preview': '\n请仔细聆听音频,并定位出中老年男性演讲者在用充满活力的语气阐述自己的人生观时,通过列举具体的体能成就(如单手俯卧撑和哑铃重量)来论证年龄数字对他毫无意义的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 234, 'preview': '\n用户需要定位一段关于演讲者人生观的论述,核心论据是具体的体能成就(单手俯卧撑、哑铃重量)。我初步扫描音频,发现在 01:30 附近,演讲者开始系统性地介绍自己的职业生涯,这通常是个人价值观的前奏。我先截取这一段进行确认,看是否是目标段落的铺垫。\n<tool_call>\n{"name": "crop_audio", "arguments": {"start_sec":...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}]
[DEBUG _get_std_messages] n=5 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 134, 'preview': '\n请仔细聆听音频,并定位出中老年男性演讲者在用充满活力的语气阐述自己的人生观时,通过列举具体的体能成就(如单手俯卧撑和哑铃重量)来论证年龄数字对他毫无意义的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 234, 'preview': '\n用户需要定位一段关于演讲者人生观的论述,核心论据是具体的体能成就(单手俯卧撑、哑铃重量)。我初步扫描音频,发现在 01:30 附近,演讲者开始系统性地介绍自己的职业生涯,这通常是个人价值观的前奏。我先截取这一段进行确认,看是否是目标段落的铺垫。\n<tool_call>\n{"name": "crop_audio", "arguments": {"start_sec":...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}]
Train: 0%| | 1/3747 [1:55:34<7215:36:30, 6934.38s/it]
{'loss': 1e-08, 'grad_norm': 0.43653446, 'learning_rate': 0.0001, 'completions/mean_length': 372.125, 'completions/min_length': 323.0, 'completions/max_length': 412.0, 'completions/clipped_ratio': 0.0, 'num_turns': 3.0, 'reward': 1.54999995, 'reward_std': 1.34667385, 'frac_reward_zero_std': 0.0, 'rewards/AudioFormatORM/mean': 0.5, 'rewards/AudioFormatORM/std': 0.53452247, 'rewards/AudioIoUORM/mean': 0.75, 'rewards/AudioIoUORM/std': 0.88640529, 'rewards/AudioThinkContentORM/mean': 0.30000001, 'rewards/AudioThinkContentORM/std': 0.213809, 'kl': 0.0, 'clip_ratio/low_mean': 0.0, 'clip_ratio/low_min': 0.0, 'clip_ratio/high_mean': 0.0, 'clip_ratio/high_max': 0.0, 'clip_ratio/region_mean': 0.0, 'epoch': 0.0, 'global_step/max_steps': '1/3747', 'percentage': '0.03%', 'elapsed_time': '1h 55m 34.89s', 'remaining_time': '300d 16h 8m 8.1s', 'memory(GiB)': 71.83, 'train_speed(s/it)': 6934.887373}
Train: 0%| | 1/3747 [1:55:34<7215:36:30, 6934.38s/it]
Train: 0%| | 1/3747 [1:55:34<7215:36:30, 6934.38s/it][DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 152, 'preview': '\n请仔细聆听音频。在背景音乐变为轻快且滑稽的管弦乐后,女主持人开始讽刺男主角廉价而霸道的示爱方式。请定位出她列举具体事例(如关空调、全城搜人)来论证一个心机深沉的女主不应被这些虚伪算计所迷惑的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 241, 'preview': '\n用户需要定位一个包含特定行为论证的段落。关键线索是“关空调”、“全城搜人”以及背景音乐变为“轻快滑稽的管弦乐”。我先在全局寻找与“空调”和“霸道”相关的关键词。音频的语音文本显示,在2分20秒附近,ASR中出现了“空调”一词。我先截取这一小段进行初步确认。\n<tool_call>\n{"name": "crop_audio", "arguments": {"start...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 152, 'preview': '\n请仔细聆听音频。在背景音乐变为轻快且滑稽的管弦乐后,女主持人开始讽刺男主角廉价而霸道的示爱方式。请定位出她列举具体事例(如关空调、全城搜人)来论证一个心机深沉的女主不应被这些虚伪算计所迷惑的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 243, 'preview': '\n用户需要定位女主持人讽刺男主角示爱方式的片段,关键特征是“轻快且滑稽的管弦乐”和具体事例。我初步扫描音频,在 1分58秒 附近检测到女主持人开始讨论男主角的示爱行为,并且背景音乐风格有显著变化。我先从这里切入,提取 01:55 到 02:15 的片段进行确认。\n<tool_call>\n{"name": "crop_audio", "arguments": {"sta...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 127, 'preview': '\n请仔细聆听音频,并定位出在激烈的剥蛋人机大赛结束后,旁白宣布最终比赛结果,并明确指出机器在“质量”方面(光滑完整、无带壳)最终优于人工队的结论性话语?请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 268, 'preview': '\n用户需要定位“剥蛋人机大赛”结束后的“最终比赛结果”和“质量优于人工队”的结论。我将首先在音频中搜索与“比赛”相关的语音片段。根据全局扫描,音频后半段有密集的机器运行声和比赛解说。我先从一个听起来像是比赛正式开始的片段切入,以建立上下文。我选择在 13:30 附近开始,那里有明确的“最终的比赛开始了”的宣告。\n<tool_call>\n{"name": "crop_a...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 152, 'preview': '\n请仔细聆听音频。在背景音乐变为轻快且滑稽的管弦乐后,女主持人开始讽刺男主角廉价而霸道的示爱方式。请定位出她列举具体事例(如关空调、全城搜人)来论证一个心机深沉的女主不应被这些虚伪算计所迷惑的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 232, 'preview': '\n用户需要定位一个包含特定音乐和特定论点的片段。关键词是“背景音乐变为轻快且滑稽的管弦乐”和“女主持人”讽刺男主角。我在音频的01:58秒附近检测到一个明显的场景转换标志,紧接着是女主持人的声音。我先从这个时间点开始,截取一段进行初步分析。\n<tool_call>\n{"name": "crop_audio", "arguments": {"start_sec": 11...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 152, 'preview': '\n请仔细聆听音频。在背景音乐变为轻快且滑稽的管弦乐后,女主持人开始讽刺男主角廉价而霸道的示爱方式。请定位出她列举具体事例(如关空调、全城搜人)来论证一个心机深沉的女主不应被这些虚伪算计所迷惑的完整段落。请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 258, 'preview': '\n用户需要定位女主持人讽刺男主角示爱方式的片段,关键信息点是“背景音乐变为轻快且滑稽的管弦乐”以及具体事例“关空调”、“全城搜人”。我首先在音频中检索与“男主”、“爱”相关的语音内容。在 01:40 附近,我注意到主持人开始讨论男主的动机。这似乎是整个论证的起点,我先从这里切入进行初步分析。\n<tool_call>\n{"name": "crop_audio", "ar...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 127, 'preview': '\n请仔细聆听音频,并定位出在激烈的剥蛋人机大赛结束后,旁白宣布最终比赛结果,并明确指出机器在“质量”方面(光滑完整、无带壳)最终优于人工队的结论性话语?请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 212, 'preview': '\n用户正在寻找“剥蛋人机大赛”的最终结果,关键词是“质量优于人工队”。音频中存在多场比赛,我先定位到“人机大赛”这个事件,初步检索到大约在 3 分钟附近有“比赛”、“开始”等语音,我先从这里开始探查。\n<tool_call>\n{"name": "crop_audio", "arguments": {"start_sec": 190, "end_sec": 220}}\n...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 127, 'preview': '\n请仔细聆听音频,并定位出在激烈的剥蛋人机大赛结束后,旁白宣布最终比赛结果,并明确指出机器在“质量”方面(光滑完整、无带壳)最终优于人工队的结论性话语?请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 255, 'preview': '\n用户需要定位一个关于“剥蛋人机大赛”的“最终比赛结果”,核心关键词是“质量”、“光滑完整”、“无带壳”。我在音频的后半段进行全局搜索,注意到在13分45秒附近有“比赛”、“机器”、“人工”等关键词的密集出现,且背景音乐紧张,这很可能是比赛的高潮部分。我先从这里切入,确认一下具体内容。\n<tool_call>\n{"name": "crop_audio", "argum...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[DEBUG _get_std_messages] n=6 messages=[{'role': 'user', 'content_preview': {'type': 'str', 'len': 127, 'preview': '\n请仔细聆听音频,并定位出在激烈的剥蛋人机大赛结束后,旁白宣布最终比赛结果,并明确指出机器在“质量”方面(光滑完整、无带壳)最终优于人工队的结论性话语?请严格按照 [MM:SS - MM:SS] 的格式输出结果,不要输出任何多余的解释文字。'}}, {'role': 'assistant', 'content_preview': {'type': 'str', 'len': 277, 'preview': '\n用户正在寻找关于一场“剥蛋人机大赛”结果的结论。我需要定位到旁白宣布最终结果,并强调机器在“质量”上优于人工的片段。首先,我需要找到比赛的高潮或结束点。通过对音频的全局扫描,我发现在大约17分25秒附近,旁白情绪激昂地宣布了比赛的完成。这很可能是比赛结果宣布的前奏。我先裁剪 17:25 到 17:55 这一段进行初步确认。\n<tool_call>\n{"name": ...'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}, {'role': 'tool', 'content_preview': {'type': 'list', 'len': 2}}, {'role': 'assistant', 'content_preview': {'type': 'dict'}}]
[WARNING:swift] Rollout logprobs count (327) does not match completion tokens count (482). Skipping rollout importance sampling for this batch.
目前想要请问:
- 在zero3不适配多模态训练的情况下,请问推荐zero2上继续尝试还是fsdp2呢?考虑到之后可能全量微调qwen3omni,哪一个显存更友好呢?
- 目前卡死的情况我怀疑是不是plugin代码哪里有问题,请问导致这种卡死可能是plugin代码导致的嘛,需要重点检查plugin代码的哪些方面呢?
- 这种卡死会不会也是数据可能导致的,数据应该从哪些方面检查可能会导致这种卡死的情况出现?
- 除此之外还有什么可以考虑的排查方面
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start by reproducing the multi-turn Qwen3-Omni GRPO run with the supplied DeepSpeed and FSDP configurations. Read swift/rlhf_trainers/grpo_trainer.py, swift/pipelines/train/sft.py, and the reported DeepSpeed, PEFT, and model stack frames to locate where forward or backward stops. Done means the training run completes without hanging, or the FSDP2 setup error is resolved.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python, pytorch
- Domain
- distributed-systems, machine-learning
- Issue type
- Bug
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100