PyTorch非零张量发送到MPS设备后变为全0的问题咨询
MPS设备迁移后张量全零问题排查与解决
问题现象
自研CNN模型测试过程中,测试图像批次张量迁移到M1 GPU(MPS后端)后变为全零值,迁移前后非零元素统计结果:
tensor(319845) tensor(0, device='mps:0')
两次统计的唯一代码差异为是否对批次张量调用.to(device)执行设备迁移,复现测试代码如下:
def test(): correct = 0 total = 0 for i in tqdm(range(0, len(test_X), BATCH_SIZE)): print(test_X[i:i+BATCH_SIZE].view(-1, 1, 100, 100).count_nonzero()) batch_X = test_X[i:i+BATCH_SIZE].view(-1, 1, 100, 100).to(device) print(batch_X.count_nonzero()) batch_y = test_y[i:i+BATCH_SIZE].to(device) batch_out = net(batch_X) out_maxes = [torch.argmax(i, axis=0) for i in batch_out] target_maxes = [torch.argmax(i) for i in batch_y] for i,j in zip(out_maxes, target_maxes): if i == j: correct += 1 total += 1 break print(correct, total) print("Accuracy: ", round(correct/total, 3))
运行环境:M1系列芯片,使用预览版PyTorch调用MPS后端,未使用稳定版PyTorch。
排查思路
按优先级从易到难排查:
- 核对张量数据类型:早期MPS预览版对数据类型支持极不完善,
torch.float64、torch.int32等类型的张量迁移时存在已知bug,会直接出现清零、数值错乱问题,先打印CPU侧张量的dtype确认类型。 - 检查张量内存连续性:代码中先对
test_X做切片,再直接调用.view()做形状变换,切片返回的是原张量的非连续内存视图,早期MPS后端对非连续张量的设备拷贝逻辑存在缺陷,会读取错误内存地址导致全零。 - 隔离MPS拷贝逻辑:单独构造一个和输入批次同形状、同dtype的随机非零CPU张量,直接调用
.to("mps")后统计非零值,如果复现全零问题,即可确定是当前预览版PyTorch的MPS后端bug,和模型、数据加载逻辑无关。 - 排查numpy内存共享问题:如果
test_X是从numpy数组转换而来、且转换时开启了内存共享,早期MPS后端和numpy的内存交互存在兼容问题,迁移时会读取到未初始化的空内存块。
解决方案
对应排查结果按顺序尝试:
- 强制内存连续:在设备迁移前调用
.contiguous()将张量转为连续内存格式,修改对应代码行:batch_X = test_X[i:i+BATCH_SIZE].view(-1, 1, 100, 100).contiguous().to(device) - 统一使用MPS兼容数据类型:数据加载阶段就将输入图像张量转为
torch.float32,标签张量转为torch.int64,从根源规避类型兼容问题:test_X = test_X.float() test_y = test_y.long() - 绕开直接拷贝逻辑:如果上述操作仍复现问题,可以先将CPU侧张量转为numpy数组再重新构造MPS张量,绕开
.to()的内置拷贝逻辑:cpu_batch = test_X[i:i+BATCH_SIZE].view(-1, 1, 100, 100).contiguous() batch_X = torch.tensor(cpu_batch.numpy(), device=device) - 版本适配:PyTorch 1.12之前的MPS预览版存在大量张量拷贝、算子兼容bug,优先升级到PyTorch 1.13及以上的正式稳定版本,这类清零问题在正式版中已全部修复。暂时无法升级的话,可以临时切换到CPU设备跑测试,等版本更新后再切回MPS。
内容的提问来源于stack exchange,提问作者gokulsrin
相关产品推荐
相关产品推荐

