You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Mac MPS训练时PyTorch DataLoader逐步变慢问题求助

Mac下PyTorch MPS后端训练时DataLoader逐步变慢的问题排查与解决

问题现象

在Mac设备上使用MPS后端训练关键点检测模型时,出现DataLoader异常:训练初期速度较快,但后续步骤耗时逐渐增加;切换为CPU后端训练时,DataLoader速度始终稳定且快速。已排查张量迁移到MPS设备的耗时,确认该操作耗时极短,并非问题根源。

数据集代码

# just load image rescale it, to tensor and process annotation coord
def load_coord_data(img_path, anno_path, h, w):
  img = cv2.imread(img_path, cv2.IMREAD_COLOR)
  scale = img.shape[0] / h
  img = cv2.resize(img, (w, h))
  img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
  img = img2tensor(img)

  with open(anno_path, 'r') as f:
    anno = json.loads(f.read())
    coords = np.zeros((50, 2), dtype=np.float32) # 修正原代码len(50)的错误
    for idx in anno:
      if idx in anno:  # 修正原代码_class未定义的错误
        coords[idx, 0], coords[idx, 1] = anno[idx]['x'], anno[idx]['y']
      else:
        coords[idx, 0], coords[idx, 1] = -1.0, -1.0

  return img, coords

class KeyPointsDataset(Dataset):
  def __init__(self, h, w, input_dir="xxx"):
    self.h, self.w = h, w
    files = os.listdir(input_dir)
    self.img_files = sorted([os.path.join(input_dir, fn) for fn in files if fn.endswith("jpg")])
    self.anno_files = sorted([os.path.join(input_dir, fn) for fn in files if fn.endswith("json")])
  
  def __getitem__(self, idx):
      img, labels = load_coord_data(self.img_files[idx], self.anno_files[idx], self.h, self.w)
      return img, labels
  
  def __len__(self):
    return len(self.img_files)

训练代码

dataset = KeyPointsDataset(h, w)
dataloader = DataLoader(
 dataset,
  batch_size=1,
  shuffle=True,
  num_workers=1,
  drop_last=False,
  pin_memory=True
)
  
loss = th.nn.MSELoss()
device = th.device('mps')
self.model = self.model.to(device)
self.model.train()

for epoch in range(10):
  for step, (img, labels) in enumerate(dataloader):
    img, labels = img.to(device, non_blocking=True), labels.to(device, non_blocking=True)
    # the time that dataloder took
    if step > 0:
      dataloader_time = round(time.monotonic() - toc, 2)
    else:
      dataloader_time = -1

    tic = time.monotonic()

    pred = self.model(img)
    
    _loss = loss(pred, labels)
    _loss.backward()
    self.opt.step()
    self.opt.zero_grad()

    toc = time.monotonic()
    if step % self.config.LOG_STEPS == 0:
      print('Epoch {:03d} | Step {:05d} | Step Loss {:.6} | Train time {} | Dataloader time {}'.format(
        epoch, step, float(_loss.cpu().detach().numpy()), round(toc - tic, 2), dataloader_time))

MPS训练时的日志(DataLoader逐步变慢)

Epoch 000 | Step 00000 | Step Loss 0.828719 | Train time 3.6 | Dataloader time -1.0
Epoch 000 | Step 00001 | Step Loss 0.708257 | Train time 1.09 | Dataloader time 0.13
Epoch 000 | Step 00002 | Step Loss 0.658343 | Train time 1.15 | Dataloader time 1.27
Epoch 000 | Step 00003 | Step Loss 0.493049 | Train time 1.8 | Dataloader time 2.02
Epoch 000 | Step 00004 | Step Loss 2.28905 | Train time 1.36 | Dataloader time 4.77
Epoch 000 | Step 00005 | Step Loss 0.322044 | Train time 2.05 | Dataloader time 3.58
Epoch 000 | Step 00006 | Step Loss 0.535195 | Train time 1.61 | Dataloader time 5.31
Epoch 000 | Step 00007 | Step Loss 0.647095 | Train time 1.93 | Dataloader time 4.69
Epoch 000 | Step 00008 | Step Loss 0.572585 | Train time 2.03 | Dataloader time 4.75
Epoch 000 | Step 00009 | Step Loss 0.533676 | Train time 5.66 | Dataloader time 6.76
Epoch 000 | Step 00010 | Step Loss 0.569616 | Train time 2.16 | Dataloader time 8.46
Epoch 000 | Step 00011 | Step Loss 0.527826 | Train time 1.95 | Dataloader time 6.09
Epoch 000 | Step 00012 | Step Loss 0.429697 | Train time 2.89 | Dataloader time 5.06
Epoch 000 | Step 00013 | Step Loss 0.463338 | Train time 3.53 | Dataloader time 7.06
Epoch 000 | Step 00014 | Step Loss 0.573107 | Train time 3.31 | Dataloader time 7.57
Epoch 000 | Step 00015 | Step Loss 0.664436 | Train time 2.01 | Dataloader time 6.17
Epoch 000 | Step 00016 | Step Loss 0.420959 | Train time 1.76 | Dataloader time 5.49
Epoch 000 | Step 00017 | Step Loss 0.366839 | Train time 5.88 | Dataloader time 5.55

CPU训练时的日志(DataLoader速度稳定)

Epoch 000 | Step 00000 | Step Loss 0.768135 | Train time 6.0 | Dataloader time -1
Epoch 000 | Step 00001 | Step Loss 0.912373 | Train time 4.49 | Dataloader time 0.0
Epoch 000 | Step 00002 | Step Loss 0.678868 | Train time 5.33 | Dataloader time 0.0
Epoch 000 | Step 00003 | Step Loss 0.518494 | Train time 5.62 | Dataloader time 0.0
Epoch 000 | Step 00004 | Step Loss 0.647296 | Train time 5.34 | Dataloader time 0.0
Epoch 000 | Step 00005 | Step Loss 0.621026 | Train time 4.64 | Dataloader time 0.0
Epoch 000 | Step 00006 | Step Loss 0.611825 | Train time 5.24 | Dataloader time 0.0
Epoch 000 | Step 00007 | Step Loss 0.557198 | Train time 4.31 | Dataloader time 0.0
Epoch 000 | Step 00008 | Step Loss 0.341876 | Train time 5.15 | Dataloader time 0.0
Epoch 000 | Step 00009 | Step Loss 0.425114 | Train time 5.41 | Dataloader time 0.0
Epoch 000 | Step 00010 | Step Loss 0.526096 | Train time 5.83 | Dataloader time 0.0
Epoch 000 | Step 00011 | Step Loss 0.541208 | Train time 4.14 | Dataloader time 0.0

张量迁移到MPS设备的耗时测试

测试代码片段:

...
t1 = time.time()
img, labels = img.to(device, non_blocking=True), labels.to(device, non_blocking=True)
print("to mps time: ", time.time() - t1)
...

测试结果:

to mps time:  0.0009737014770507812
Epoch 000 | Step 00000 | Step Loss 0.779288 | Train time 3.54 | Dataloader time -1
to mps time:  0.0029366016387939453
Epoch 000 | Step 00001 | Step Loss 0.81939 | Train time 1.12 | Dataloader time 0.39
to mps time:  0.0012209415435791016
Epoch 000 | Step 00002 | Step Loss 0.755161 | Train time 1.22 | Dataloader time 1.94
to mps time:  0.0012700557708740234
Epoch 000 | Step 00003 | Step Loss 0.443458 | Train time 3.04 | Dataloader time 2.13
to mps time:  0.0002002716064453125
Epoch 000 | Step 00004 | Step Loss 7.6309 | Train time 3.05 | Dataloader time 6.62
to mps time:  0.0003991127014160156
Epoch 000 | Step 00005 | Step Loss 0.802644 | Train time 1.99 | Dataloader time 6.63
to mps time:  0.0002162456512451172
Epoch 000 | Step 00006 | Step Loss 0.690297 | Train time 1.95 | Dataloader time 5.55
to mps time:  0.00043082237243652344
Epoch 000 | Step 00007 | Step Loss 0.805149 | Train time 3.16 | Dataloader time 6.13
to mps time:  0.00084686279296875
Epoch 000 | Step 00008 | Step Loss 0.729108 | Train time 1.84 | Dataloader time 6.87
to mps time:  0.00031113624572753906
Epoch 000 | Step 00009 | Step Loss 0.575548 | Train time 1.73 | Dataloader time 4.74

原因分析

  • MPS与DataLoader多线程的兼容性问题:PyTorch的MPS后端在与多进程/多线程DataLoader配合时,可能存在进程间内存同步或资源竞争问题,尤其是num_workers>0时,工作进程与主进程的MPS设备数据交互易出现累积性阻塞。
  • Pin Memory适配缺陷:pin_memory=True在CPU训练时可加速数据传输,但MPS后端对固定内存的支持不完善,固定内存张量向MPS迁移时会产生隐性内存拷贝阻塞,且训练过程中内存碎片累积会加剧阻塞。
  • 工作进程资源泄漏:MPS设备资源在工作进程中未被正确释放,随着训练步数增加,资源占用累积,导致DataLoader工作进程获取数据的速度变慢。

解决方案

  • 调整DataLoader参数:
    • 将num_workers设为0,使用主进程加载数据,规避多进程与MPS的兼容性问题。
    • 关闭pin_memory,设置pin_memory=False,消除MPS对固定内存的适配问题。
  • 优化数据集加载逻辑:
    • 若数据集不大,预加载部分数据到内存,减少磁盘IO影响。
    • 替换img2tensor为PyTorch原生torchvision.transforms.ToTensor(),避免第三方库潜在兼容性问题。
    • 修复数据集代码中的错误:将len(50)改为50,_class in anno改为idx in anno(原代码逻辑错误会导致标签处理异常)。
  • 定期清理MPS内存:在每个epoch结束后调用torch.mps.empty_cache(),避免内存碎片累积。
  • 升级PyTorch版本:MPS后端仍在迭代,较新版本可能修复了DataLoader相关兼容性bug。

内容的提问来源于stack exchange,提问作者YQ.Wang

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.01 06:34:56