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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 21:20:09