使用 TorchComms 和调试服务器通过飞行记录器调试挂起#
TorchComms 是一个 Python 库,它在 PyTorch 的分布式后端(如 NCCL、Gloo 等)之上提供了一个高级通信器抽象。它使用钩子 (hook) 系统封装了集合通信操作,使你无需修改应用程序代码即可附加监测功能,例如日志记录、性能分析或记录。更多信息,请参阅 TorchComms 文档。
TorchComms 飞行记录器 (Flight Recorder) 就是这样一个钩子(FlightRecorderHook)。它维护一个固定大小的环形缓冲区,静默记录通过 TorchComms 通信器发出的每一个集合通信操作,捕获操作类型、序列号、张量形状、执行状态以及调用点的 Python 堆栈跟踪。有关飞行记录器的参考,请参阅 TorchComms 中的 飞行记录器钩子。
调试服务器 (Debug Server)(torch.distributed.debug)在每个 rank 上运行一个 HTTP 服务器,可以定期将飞行记录器快照和 Python 堆栈跟踪转储到磁盘。有关调试服务器、其端点和定期转储的完整 API 文档,请参阅 torch.distributed 调试 HTTP 服务器文档。
本教程将引导你完成一个具体示例:一个双 rank 任务,其中一个 rank 在集合通信之前挂起,并展示如何分步使用飞行记录器和调试服务器来诊断该问题。
注意
本教程涵盖了针对飞行记录器转储的 TorchComms 调试服务器方法。有关基于环境变量的旧版飞行记录器配置,请参阅 用于调试卡死任务的飞行记录器。
先决条件#
PyTorch 2.5 或更高版本,包含
torch.distributed已安装
torchcomms(pip install torchcomms)拥有 2 个或更多 GPU 的 CUDA 主机,或使用
TEST_BACKEND=gloo TEST_DEVICE=cpu进行纯 CPU 测试
您将学到什么#
如何将
FlightRecorderHook附加到 TorchComms 通信器如何启动带有定期飞行记录器转储功能的调试服务器
如何读取聚合文本转储以识别缺失的 rank 和不匹配的集合通信
如何使用按 rank 分类的 pickle 追踪文件与 FR CLI 进行跨 rank 分析
如何解读堆栈跟踪快照以定位挂起的具体代码行
飞行记录器概述#
每个飞行记录器条目捕获:
字段 |
描述 |
|---|---|
|
单调递增的序列号(对于给定的集合通信,所有 rank 相同) |
|
例如 |
|
|
|
张量形状 |
|
调用点的 Python 堆栈跟踪 |
当调试服务器启用定期转储时,每个转储周期会产生两种类型的输出:
聚合文本文件 (
torchcomms_fr_trace_<ts>.txt) — rank 0 上的前端从所有 rank 获取飞行记录器数据并写入人类可读的表格。按 rank 分类的 pickle 文件 (
per_rank/rank_<N>) — 每个 rank 的工作服务器写入其自己的 pickle 追踪文件。这些文件可以输入到 FR CLI (python -m torch.distributed.flight_recorder.fr_trace) 中,用于自动检测跨 rank 不匹配。
场景#
下面的演示脚本创建了一个两阶段的工作负载:
阶段 1 (所有 rank):3 次
all_reduce+ 1 次broadcast操作正常完成。阶段 2:
挂起的 rank 进入
time.sleep。其他 rank 发起另一次
all_reduce,由于等待挂起的 rank 而超时。
当超时触发时,转储目录包含:
FR_DUMP_DIR/
├── torchcomms_fr_trace_<ts>.txt ← aggregated text
└── per_rank/ ← per-rank pickle files
├── rank_0
└── rank_1
演示脚本#
将以下内容保存为 verify_flight_recorder.py
import os
import time
from datetime import timedelta
import torch
from torch.distributed.debug import start_debug_server
from torchcomms import new_comm, ReduceOp
from torchcomms.hooks import FlightRecorderHook
def main():
backend = os.environ.get("TEST_BACKEND", "gloo")
device = torch.device(os.environ.get("TEST_DEVICE", "cuda"))
dump_dir = os.environ.get("FR_DUMP_DIR", "/tmp/fr_hang_debug")
dump_interval = float(os.environ.get("FR_DUMP_INTERVAL", "5"))
timeout_seconds = int(os.environ.get("COMM_TIMEOUT", "30"))
hanging_rank = int(os.environ.get("HANGING_RANK", "-1"))
os.makedirs(dump_dir, exist_ok=True)
per_rank_dir = os.path.join(dump_dir, "per_rank")
os.makedirs(per_rank_dir, exist_ok=True)
dump_prefix = os.path.join(per_rank_dir, "rank_")
os.environ["TORCHCOMM_FR_DUMP_TEMP_FILE"] = dump_prefix
comm = new_comm(
backend=backend,
device=device,
name="main_comm",
timeout=timedelta(seconds=timeout_seconds),
abort_process_on_timeout_or_error=False,
)
rank = comm.get_rank()
world_size = comm.get_size()
if hanging_rank < 0:
hanging_rank = world_size - 1
num_devices = torch.cuda.device_count()
device_id = rank % num_devices
target_device = torch.device(f"cuda:{device_id}")
print(
f"[Rank {rank}/{world_size}] device={device_id}, "
f"hanging_rank={hanging_rank}, timeout={timeout_seconds}s"
)
# ── Debug Server with Periodic Dumps ──
start_debug_server(
port=25999,
dump_dir=dump_dir,
dump_interval=dump_interval,
enabled_dumps={"torchcomms_fr_trace", "stacks"},
)
if rank == 0:
print(f"[Rank {rank}] Debug server: https://:25999")
print(f"[Rank {rank}] Periodic dumps every {dump_interval}s → {dump_dir}")
print(f"[Rank {rank}] Per-rank pickles → {per_rank_dir}")
# ── Flight Recorder Hook ──
recorder = FlightRecorderHook(max_entries=100)
recorder.register_with_comm(comm)
tensor = torch.full(
(1024,),
float(rank + 1),
dtype=torch.float32,
device=target_device,
)
# ── Phase 1: Successful collectives (all ranks) ──
print(f"[Rank {rank}] Phase 1: Running 3 all_reduce + 1 broadcast")
for _i in range(3):
comm.all_reduce(tensor, ReduceOp.SUM, async_op=False)
comm.broadcast(tensor, root=0, async_op=False)
torch.cuda.current_stream().synchronize()
print(f"[Rank {rank}] Phase 1 complete")
# ── Phase 2: One rank hangs ──
if rank == hanging_rank:
print(f"[Rank {rank}] >>> HANGING – entering infinite sleep <<<")
while True:
time.sleep(1)
print(
f"[Rank {rank}] Phase 2: all_reduce "
f"(rank {hanging_rank} will NOT participate)"
)
print(f"[Rank {rank}] Expecting timeout in ~{timeout_seconds}s ...")
try:
comm.all_reduce(tensor, ReduceOp.SUM, async_op=False)
except Exception as e:
print(f"[Rank {rank}] Caught timeout: {type(e).__name__}: {e}")
recorder.dump_file(rank)
print(f"[Rank {rank}] Pickle trace written to {dump_prefix}{rank}")
recorder.unregister()
comm.finalize()
if __name__ == "__main__":
main()
注意
此脚本需要 torchcomms (pip install torchcomms) 和 torch.distributed.debug。torchcomms 包依赖 tabulate, jinja2 和 aiohttp。
运行演示#
启动#
FR_DUMP_DIR=/tmp/fr_hang_debug \
FR_DUMP_INTERVAL=3 \
COMM_TIMEOUT=15 \
TEST_BACKEND=gloo \
TEST_DEVICE=cpu \
torchrun --nproc_per_node=2 verify_flight_recorder.py
变量 |
默认 |
描述 |
|---|---|---|
|
|
根转储目录 |
|
|
定期转储之间的秒数 |
|
|
通信器超时(秒) |
|
|
要挂起的 rank |
|
|
通信后端 |
|
|
张量设备 |
预期输出#
[Rank 0/2] device=0, hanging_rank=1, timeout=15s
[Rank 1/2] device=1, hanging_rank=1, timeout=15s
[Rank 0] Debug server: https://:25999
[Rank 0] Periodic dumps every 3.0s → /tmp/fr_hang_debug
[Rank 0] Per-rank pickles → /tmp/fr_hang_debug/per_rank
[Rank 0] Phase 1: Running 3 all_reduce + 1 broadcast
[Rank 0] Phase 1 complete
[Rank 0] Phase 2: all_reduce (rank 1 will NOT participate)
[Rank 0] Expecting timeout in ~15s ...
[Rank 1] Phase 1 complete
[Rank 1] >>> HANGING – entering infinite sleep <<<
... periodic mismatch warnings every 3 seconds ...
Not all ranks joining collective, sequence number: 4
collective: nccl:all_reduce
missing ranks: {1}
collective state: scheduled
... ~15 seconds pass ...
[Rank 0] Caught timeout: RuntimeError: Timed out waiting 15000ms for recv operation
[Rank 0] Pickle trace written to /tmp/fr_hang_debug/per_rank/rank_0
读取聚合文本转储#
调试服务器会定期写入聚合所有 rank 数据的文本快照。
$ ls /tmp/fr_hang_debug/torchcomms_fr_trace_*.txt
torchcomms_fr_trace_20260401_192058.txt
torchcomms_fr_trace_20260401_192101.txt
torchcomms_fr_trace_20260401_192104.txt
...
打开挂起期间写入的一个快照:
cat /tmp/fr_hang_debug/torchcomms_fr_trace_20260401_192104.txt
Collectives 表显示了每一个记录的操作。
--- Collectives ---
id group_id pass_check collective_seq_id collective_name collective_state missing_ranks
0 main_comm True 0 nccl:all_reduce scheduled
1 main_comm True 1 nccl:all_reduce scheduled
2 main_comm True 2 nccl:all_reduce scheduled
3 main_comm True 3 nccl:broadcast scheduled
4 main_comm True 4 nccl:all_reduce scheduled {1} ← MISMATCH
NCCL Calls 表显示了哪些 rank 参与了通信。
--- NCCL Calls ---
id collective_id group_id global_rank collective_type
0 0 main_comm 0 nccl:all_reduce
1 0 main_comm 1 nccl:all_reduce
...
6 3 main_comm 0 nccl:broadcast
7 3 main_comm 1 nccl:broadcast
8 main_comm 0 nccl:all_reduce ← Only rank 0!
Dump File 部分确认已写入按 rank 分类的 pickle 文件。
=== TorchComms FR Dump File ===
Rank 0: OK - Flight Recorder debug info written to /tmp/fr_hang_debug/per_rank/rank_0
Rank 1: OK - Flight Recorder debug info written to /tmp/fr_hang_debug/per_rank/rank_1
stacks_*.txt 文件显示了 Python 堆栈跟踪,定位到每个 rank 卡住的具体代码行。
$ cat /tmp/fr_hang_debug/stacks_20260401_192104.txt
=== Rank 0 ===
File "verify_flight_recorder.py", line 148 in main ← all_reduce (waiting)
=== Rank 1 ===
File "verify_flight_recorder.py", line 140 in main ← time.sleep (the hang!)
Rank 1 从未发出 collective_seq_id=4。堆栈转储确认它卡在 time.sleep 中,而不是集合通信中。
对按 rank 分类的 pickle 转储运行 FR CLI#
定期转储还会触发每个 rank 的工作服务器将 pickle 追踪文件写入 per_rank/ 子目录。
$ ls /tmp/fr_hang_debug/per_rank/
rank_0 rank_1
跨 rank 不匹配分析#
python -m torch.distributed.flight_recorder.fr_trace \
/tmp/fr_hang_debug/per_rank -p rank_
输出
Not all ranks joining collective, sequence number: 4
internal record id: 4
group info: main_comm:gloo
collective: nccl:all_reduce
missing ranks: {1}
input sizes: [[1024]]
output sizes: [[1024]]
world size: 2
expected ranks: {0, 1}
collective state: scheduled
CLI 检测到 rank 1 从未发出 collective_seq_id=4。
并排原始条目视图#
python -m torch.distributed.flight_recorder.fr_trace \
/tmp/fr_hang_debug/per_rank -p rank_ -j
输出
Rank 0 Rank 1
------------------------------------------------- -------------------------------------------------
all_reduce(input_sizes=[[1024]], state=scheduled) all_reduce(input_sizes=[[1024]], state=scheduled)
all_reduce(input_sizes=[[1024]], state=scheduled) all_reduce(input_sizes=[[1024]], state=scheduled)
all_reduce(input_sizes=[[1024]], state=scheduled) all_reduce(input_sizes=[[1024]], state=scheduled)
broadcast(input_sizes=[[1024]], state=scheduled) broadcast(input_sizes=[[1024]], state=scheduled)
all_reduce(input_sizes=[[1024]], state=scheduled)
Rank 0 有 5 个条目(3 次 all_reduce + 1 次 broadcast + 卡住的 all_reduce)。Rank 1 只有 4 个条目——第 5 次 all_reduce 丢失了,因为 rank 1 在发出它之前就挂起了。
带堆栈跟踪#
python -m torch.distributed.flight_recorder.fr_trace \
/tmp/fr_hang_debug/per_rank -p rank_ -j --print_stack_trace
这会将 Python 堆栈跟踪添加到每个条目,准确显示用户代码中调用每个集合通信的位置。
需要查找的内容#
症状 |
可能的原因 |
|---|---|
Collectives 表中的 |
Rank N 在发出下一个集合通信之前挂起或崩溃了 |
Rank X 的最后一条目为 |
Rank X 发出了集合通信,但正在等待未加入的对等方 |
在同一个 |
代码路径分歧 — rank 正在调用不同的集合通信 |
不匹配的 |
跨 rank 的张量形状不一致 |
堆栈转储显示 |
rank 卡在计算中,而不是集合通信中 |
FR CLI 快速参考#
# Cross-rank mismatch analysis:
python -m torch.distributed.flight_recorder.fr_trace <dir> -p <prefix>
# Side-by-side raw entries per rank:
python -m torch.distributed.flight_recorder.fr_trace <dir> -p <prefix> -j
# With stack traces:
python -m torch.distributed.flight_recorder.fr_trace <dir> -p <prefix> -j --print_stack_trace
# Best-effort when some rank dumps are missing:
python -m torch.distributed.flight_recorder.fr_trace <dir> -p <prefix> --allow-incomplete-ranks
结论#
在本教程中,你学习了如何使用 TorchComms 飞行记录器和调试服务器来诊断分布式 PyTorch 任务中的单 rank 挂起问题。通过检查聚合文本转储、按 rank 分类的 pickle 追踪和堆栈跟踪快照,你识别出了哪个集合通信卡住了、哪个 rank 未参与,以及导致挂起的具体代码行。你可以将相同的工作流程应用于调试实际的分布式训练挂起问题——用你的任务实际卡住的位置替换模拟的 time.sleep,飞行记录器将向你展示 rank 在何处发生了分歧。
另请参阅#
用于调试卡死任务的飞行记录器 — 基于环境变量的飞行记录器配置