TrOCR在Mac M4芯片(MPS)上微调时出现RuntimeError问题求助
Hey,我之前在MPS设备上跑PyTorch模型时也遇到过一模一样的错误,来帮你捋捋问题根源和解决办法~
首先,这个错误的核心原因很明确:MPS对张量的内存连续性要求比CPU严格得多。view()操作要求张量在内存中是连续存储的(可以用tensor.is_contiguous()验证),但有些操作比如切片、转置、甚至squeeze()都可能让张量变成非连续的。CPU对这种情况比较宽容,能勉强处理非连续张量的view(),但MPS会直接抛出错误,而reshape()会自动重新整理张量的内存布局,所以报错提示你改用它。
针对你的TrOCR微调场景,给你几个优先级从高到低的解决思路:
1. 给输入张量强制加上.contiguous()
这是最简单直接的尝试,修改你main函数里加载batch数据的两行代码:
pixel_values = batch["pixel_values"].to(device).contiguous() labels = batch["labels"].to(device).contiguous()
这样能确保传入模型的张量在MPS内存中是连续的,大概率能直接解决这个view不兼容的问题。
2. 在Dataset阶段提前处理张量连续性
看你的OCRDataset的__getitem__方法里,用了squeeze()去掉多余维度,这个操作可能会破坏张量的连续性。可以在squeeze()之后加上.contiguous():
return { "pixel_values": pixel_values.squeeze().contiguous(), "labels": labels.squeeze().contiguous() }
提前在数据加载阶段把张量处理成连续的,避免后续设备转移时出现问题。
3. 更新transformers库到最新版本
TrOCR属于transformers库中的模型,旧版本对MPS的支持可能存在一些兼容性小bug,官方后续可能已经修复了这类view相关的问题。你可以执行下面的命令更新:
pip install --upgrade transformers
4. 临时替换模型内部的view为reshape(应急方案)
如果上面的方法都无效,那大概率是模型内部某些层硬写了view操作。你可以用猴子补丁的方式,把模型中所有的view替换成reshape,放在加载模型之前执行:
import torch.nn as nn # 替换view为reshape def patched_view(self, *args): return self.reshape(*args) nn.Module.view = patched_view
这个方法稍显粗暴,但能快速绕过MPS的兼容性限制。
你可以先从第一种方法开始试,基本能解决大部分这类MPS上的张量连续性问题。另外,M4芯片的MPS已经支持大部分PyTorch操作,但部分预训练模型的层可能还是存在小的适配问题,这些方案应该能帮你搞定。
备注:内容来源于stack exchange,提问作者New

