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

M2 Max上PyTorch中MPS GPU加载Dataloader迭代报错求助

解决PyTorch MPS环境下Dataloader直接加载数据到GPU的RuntimeError问题

错误原因

PyTorch的MPS设备不支持跨进程的张量操作,而Dataloader默认启用多进程(num_workers>0)加载数据。当设置pin_memory_device="mps"时,工作进程在CPU上生成张量后,试图将其固定到MPS设备,跨进程的设备操作导致了张量存储设备不匹配的错误。

解决方案

方案1:单进程Dataloader + Dataset内直接加载到MPS

禁用Dataloader的多进程,同时在自定义Dataset的__getitem__方法中直接将数据转换为MPS张量,确保迭代出的batch直接在GPU上:

from torch.utils.data import Dataset, DataLoader
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor

class MNISTMPSDataset(Dataset):
    def __init__(self, base_dataset):
        self.base_dataset = base_dataset
    
    def __len__(self):
        return len(self.base_dataset)
    
    def __getitem__(self, idx):
        data, target = self.base_dataset[idx]
        # 直接将数据转移到MPS设备
        return data.to(device="mps"), target.to(device="mps")

# 初始化原始数据集
mnist_test_dataset = MNIST(root="./data", train=False, download=True, transform=ToTensor())
# 包装为MPS数据集
mps_dataset = MNISTMPSDataset(mnist_test_dataset)
# 禁用多进程
mnist_test_loader = DataLoader(mps_dataset, batch_size=32, shuffle=False, num_workers=0)

# 模型部署到MPS
network.to(device="mps")
for X, y in mnist_test_loader:
    prediction = network(X)
    # 后续计算逻辑

注意:单进程加载可能会成为数据瓶颈,适合数据量较小的场景。

方案2:多进程加载 + 异步批量转移

保留多进程提升数据加载速度,通过pin_memory=True固定CPU内存加速转移,并使用non_blocking=True开启异步转移,让数据转移和模型计算重叠,减少整体耗时:

from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor

mnist_test_dataset = MNIST(root="./data", train=False, download=True, transform=ToTensor())
# 启用多进程和CPU内存固定
mnist_test_loader = DataLoader(mnist_test_dataset, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

network.to(device="mps")
for X, y in mnist_test_loader:
    # 异步转移数据到MPS
    X = X.to(device="mps", non_blocking=True)
    y = y.to(device="mps", non_blocking=True)
    prediction = network(X)
    # 后续计算逻辑

这是MPS环境下兼顾加载效率和转移速度的主流优化方式。

方案3:使用PyTorch 2.0+的AsyncDataLoader(实验性)

PyTorch 2.0及以上版本提供了实验性的AsyncDataLoader,支持异步加载并自动将数据转移到目标设备,适合MPS环境:

from torch.utils.data import AsyncDataLoader
from torchvision.datasets import MNIST
from torchvision.transforms import ToTensor

mnist_test_dataset = MNIST(root="./data", train=False, download=True, transform=ToTensor())
# 指定目标设备为MPS
mnist_test_loader = AsyncDataLoader(mnist_test_dataset, batch_size=32, shuffle=False, num_workers=4, device="mps")

network.to(device="mps")
for X, y in mnist_test_loader:
    prediction = network(X)
    # 后续计算逻辑

注意:该API目前处于实验阶段,可能存在兼容性问题,建议结合你的PyTorch版本测试使用。


内容的提问来源于stack exchange,提问作者Alex_Werben

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 10:43:02