使用Donut模型提取图像文本时遇RuntimeError:输入与偏置类型不匹配
解决Donut Model推理时Input type与bias type不匹配的RuntimeError
这个错误源于模型各组件的张量类型不统一:你在CPU环境下仅将encoder转为torch.bfloat16,但模型的decoder等其他部分仍为默认的float32,同时推理时输入图像生成的张量也是float32,和模型中bfloat16类型的参数(比如bias)冲突。
修复后的完整代码
from donut import DonutModel from PIL import Image import torch model = DonutModel.from_pretrained("naver-clova-ix/donut-base-finetuned-cord-v2") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 统一模型的设备和类型 if torch.cuda.is_available(): model.half() # 使用FP16加速CUDA推理 else: model.to(torch.bfloat16) # CPU下将整个模型转为BF16而非仅encoder model.to(device) model.eval() # 统一设置eval模式,避免训练模式的随机操作 image = Image.open("testfolder/test1.jpg").convert("RGB") # 推理时禁用梯度计算,节省内存并提升效率 with torch.no_grad(): output = model.inference(image=image, prompt="<s_cord-v2>") print(output)
关键修改点
- 统一模型类型与设备:CPU环境下将整个模型转为
torch.bfloat16,而非仅修改encoder;CUDA环境下保持model.half()并将模型移至GPU,确保所有组件类型一致。 - 全局设置eval模式:不管是否有CUDA,都在最后设置
model.eval(),避免训练模式的BatchNorm、Dropout等操作干扰推理结果。 - 添加梯度禁用:用
torch.no_grad()包裹推理代码,减少内存占用,提升推理速度。
内容的提问来源于stack exchange,提问作者Rithwik Babu
相关产品推荐
相关产品推荐

