评价此页

使用 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 测试

  • 熟悉 分布式 PyTorch 概念

您将学到什么#

  • 如何将 FlightRecorderHook 附加到 TorchComms 通信器

  • 如何启动带有定期飞行记录器转储功能的调试服务器

  • 如何读取聚合文本转储以识别缺失的 rank 和不匹配的集合通信

  • 如何使用按 rank 分类的 pickle 追踪文件与 FR CLI 进行跨 rank 分析

  • 如何解读堆栈跟踪快照以定位挂起的具体代码行

飞行记录器概述#

每个飞行记录器条目捕获:

字段

描述

collective_seq_id

单调递增的序列号(对于给定的集合通信,所有 rank 相同)

profiling_name

例如 nccl:all_reduce, nccl:broadcast

state

scheduledstartedcompleted

input_dims / output_dims

张量形状

traceback

调用点的 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.debugtorchcomms 包依赖 tabulate, jinja2aiohttp

运行演示#

启动#

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

变量

默认

描述

FR_DUMP_DIR

/tmp/fr_hang_debug

根转储目录

FR_DUMP_INTERVAL

5

定期转储之间的秒数

COMM_TIMEOUT

30

通信器超时(秒)

HANGING_RANK

-1 (最后一个 rank)

要挂起的 rank

TEST_BACKEND

gloo

通信后端

TEST_DEVICE

cuda

张量设备

预期输出#

[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 表中的 missing_ranks: {N}

Rank N 在发出下一个集合通信之前挂起或崩溃了

Rank X 的最后一条目为 state=started,其他为 completed

Rank X 发出了集合通信,但正在等待未加入的对等方

在同一个 collective_seq_id 处有不匹配的 collective_name

代码路径分歧 — rank 正在调用不同的集合通信

不匹配的 input_sizes / output_sizes

跨 rank 的张量形状不一致

堆栈转储显示 time.sleep 或用户代码(非集合通信)

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 在何处发生了分歧。

另请参阅#