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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:48:30