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
相关产品推荐
相关产品推荐

