机器学习流程卡顿时先查哪里
本文围绕“卡顿时先查哪里”整理可复现的检查思路。所有阈值、配置和结果均应在隔离环境中记录输入、版本与资源条件后再解释;下文示例不对应真实组织、用户、流量或成本数据。
1. 用受控样例界定问题
# 登上卡顿节点,检查 GPU 实时状态与进程堆栈 nvidia-smi # 打印发现 GPU 显存完全被占满 (79GB/80GB),但 Power 和 Volatile GPU-Util 显示 0% 0W2. 诊断工具上手:用 strace、nvidia-smi 和 py-spy 捕捉阻塞死锁
面对没有任何日志输出的卡死现场,必须借助系统级的排障工具。
我们采取了“三步排查法”:
首先,使用py-spy打印 Python 进程的调用栈:
# 获取训练主进程 PID PID=$(pgrep -f "train_llm.py") # 在不中断进程的前提下,实时 Dump 当前进程的 Python 调用栈 py-spy dump --pid $PIDpy-spy出来的堆栈非常直观:主线程卡在了futex_wait_queue_me,等待 DataLoader 子进程的数据返回;而 DataLoader 的 Worker 子进程,则死死锁在multiprocessing/queues.py的Queue.get()方法上。
接下来,使用strace查看 Worker 进程在等待什么系统调用:
strace -p $WORKER_PID -f -e trace=futex,read,write追踪报告显示,Worker 进程正阻塞在从/dev/shm共享内存空间读取数据的系统调用上。根因水落石出:由于/dev/shm共享内存空间不足,PyTorch DataLoader 在多进程传输大张量时触发了 Shared Memory 溢出,造成了死锁!
GPU 算力归零后,应先收集诊断信息,排查 Shared Memory 和 Network I/O 阻塞,再将根因转成容量和超时防线。
3. I/O 吞吐瓶颈与 PyTorch DataLoader 多进程死锁根因分析
深度剖析 PyTorch 的DataLoader源码,会发现它的多进程通信依赖于 Linux 的 Shared Memory (/dev/shm)。
当num_workers > 0时,主进程通过 Queue 派发任务给 Worker 进程,Worker 进程解析数据并创建 PyTorch Tensor。为了避免在 IPC 进程间通信时进行昂贵的内存复制,PyTorch 会将 Tensor 写入/dev/shm共享内存段,仅把句柄传递给主进程。
如果容器或 K8s Pod 启动时没有显式配置--shm-size(Docker 默认仅仅给 64MB),当 DataLoader 尝试塞入大尺寸 Batch 或高维 Image/Text 张量时,Shared Memory 会瞬间被填满。
此时,Worker 进程在写入共享内存时会被系统挂起,等待空间释放;而主进程又在等待 Worker 进程写入完成,双方陷入无限等待的隐性死锁状态。
另外一种常见卡顿发生在磁盘小文件随机 I/O上。如果训练集包含数百万个独立的.jpg或.txt小文件,机械硬盘或网络 NAS 的 IOPs 瞬间被拉满。CPU 绝大部分时间都耗费在了等待磁盘 Seek 操作上,导致 GPU 处于极度的“算力饥饿”状态。
4. 基于 Shared Memory 与 Ray 数据流异步 Prefetch 的防卡死方案
要从根本上解决卡死问题,代码层必须具备数据加载心跳监控与容量自动降级机制。
下述 Python 代码实现了一个带心跳健康检查、共享内存安全校验以及自动超时的健壮 DataLoader 包装器:
import time import logging import threading import multiprocessing as mp import torch from torch.utils.data import DataLoader, Dataset logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") class SafeDataset(Dataset): """模拟包含隐患的数据源""" def __init__(self, size=1000): self.size = size def __getitem__(self, index): # 模拟偶然发生的慢 I/O 阻塞 if index == 500: time.sleep(10) # 模拟卡死 10 秒 return torch.randn(3, 224, 224), torch.tensor(1) def __len__(self): return self.size class HeartbeatDataLoader: """具备心跳监控与防卡死超时的安全 DataLoader 包装器""" def __init__(self, dataset: Dataset, batch_size: int, num_workers: int, timeout_seconds: int = 5): self.dataset = dataset self.batch_size = batch_size self.num_workers = num_workers self.timeout_seconds = timeout_seconds # 自动检验 Linux /dev/shm 共享内存空间 self._verify_shared_memory() def _verify_shared_memory(self): """检查 shared memory 大小,给出警告""" import os if os.path.exists("/dev/shm"): st = os.statvfs("/dev/shm") free_shm_gb = (st.f_bavail * st.f_frsize) / (1024 ** 3) logging.info(f"检测到可用 Shared Memory (/dev/shm): {free_shm_gb:.2f} GB") if free_shm_gb < 2.0: logging.warning("Shared Memory 小于 2GB!建议加大 --shm-size 挂载,防止 DataLoader 死锁。") def get_dataloader(self) -> DataLoader: return DataLoader( self.dataset, batch_size=self.batch_size, num_workers=self.num_workers, pin_memory=True, timeout=self.timeout_seconds, # 关键:配置 PyTorch 底层 C++ 队列超时机制 drop_last=True ) def run_training_loop_with_guard(): dataset = SafeDataset(size=1000) # 建立包装器 guard_loader = HeartbeatDataLoader(dataset, batch_size=32, num_workers=2, timeout_seconds=3) dataloader = guard_loader.get_dataloader() logging.info("启动训练循环防线监控...") data_iter = iter(dataloader) step = 0 while True: try: start_t = time.time() # 尝试获取下一个 Batch,底层如果超时会强行抛出 RuntimeError inputs, targets = next(data_iter) cost = time.time() - start_t step += 1 if step % 50 == 0: logging.info(f"Step {step} 顺利完成 | 数据 Fetch 耗时: {cost:.4f}s") except StopIteration: logging.info("数据迭代正常结束。") break except RuntimeError as err: # 捕获 C++ DataLoader 抛出的 Timeout 异常 logging.error(f"严重警告:检测到 DataLoader 线程卡死/超时!错误信息: {err}") logging.warning("触发应急止损策略:正在强行重启 DataLoader 进程池...") # 重新实例化 DataLoader 救场 dataloader = guard_loader.get_dataloader() data_iter = iter(dataloader) if __name__ == "__main__": run_training_loop_with_guard()代码中设置timeout=3非常关键。默认的timeout=0会导致底层 C++ 队列在读取失败时无限期挂起;而设置了具体的timeout秒数后,一旦子进程超过阈值没有返回数据,PyTorch 会直接抛出RuntimeError,让上层的 Python 捕捉并执行重启逻辑。
5. 目标环境训练卡顿排查 SOP 与监控探针沉淀
为了在团队内推广科学排障,我们制定了一套“训练卡顿标准排查 SOP”:
- 查 /dev/shm 挂载:在容器启动参数里确保显式配置了
--shm-size=64g,严禁使用默认的 64MB。 - 查网络 NCCL 通信:分布式训练卡死时,配置环境变量
export NCCL_DEBUG=INFO和export TORCH_DISTRIBUTED_DEBUG=DETAIL,排查是否有特定 Rank 的节点发出了 Socket 超时。 - 探针心跳告警:在训练主循环中引入超时 Watchdog 线程,只要超过 180 秒没有更新下一个 Batch,自动Dump 当前进程堆栈并向飞书/钉钉群推送告警。
定位卡顿就像医生做 CT 扫描,不能凭空靠感觉猜测。利用好探针与系统调用工具,再隐蔽的死锁和 I/O 瓶颈也会无所遁形。