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

从Checkpoint加载后Donut模型model.generate()推理过慢的优化求助

Donut模型加载Checkpoint后推理速度变慢的解决方案

以下是针对你遇到的Donut模型加载后推理速度骤降问题的具体解决方法:

1. 启用混合精度推理

训练阶段通常会启用自动混合精度(AMP)加速,但加载后的推理默认可能以全精度运行,这会大幅拖慢速度。可以通过两种方式开启:

  • 方式一:使用上下文管理器
    在generate调用外层包裹torch.cuda.amp.autocast():
    with torch.cuda.amp.autocast():
        outputs = model.generate(
            pixel_values,
            decoder_input_ids=decoder_input_ids,
            max_length=model.decoder.config.max_position_embeddings,
            early_stopping=True,
            pad_token_id=processor.tokenizer.pad_token_id,
            eos_token_id=processor.tokenizer.eos_token_id,
            use_cache=True,
            num_beams=1,
            bad_words_ids=[[processor.tokenizer.unk_token_id]],
            return_dict_in_generate=True,
        )
    
  • 方式二:直接转换模型为半精度
    将模型和输入数据都转为半精度:
    model = model.half()
    pixel_values = pixel_values.half().to(device)
    

2. 验证模型状态与缓存配置

虽然代码中调用了model.eval(),但加载后需确保decoder的缓存功能正常启用(Transformer生成依赖缓存加速):

model.eval()
# 强制设置decoder的use_cache为True,避免加载时被覆盖
model.decoder.config.use_cache = True

# 确认所有参数都加载到GPU
print(next(model.parameters()).device)  # 输出应为cuda:0或对应GPU设备

3. 处理分布式训练生成的Checkpoint

如果训练时使用了DistributedDataParallel(DDP),保存的Checkpoint会带有module.前缀,加载时需要去除:

model = VisionEncoderDecoderModel.from_pretrained(CKPT_PATH, config=config)
# 若模型包含module属性,取出实际模型
if hasattr(model, 'module'):
    model = model.module
model.to(device)

4. 优化推理批量与数据加载

单张图片推理存在GPU调度开销,建议批量处理多张图片提升效率:

from torch.utils.data import DataLoader

# 构建批量数据加载器
val_loader = DataLoader(val_ds, batch_size=8, shuffle=False)
model.eval()

with torch.no_grad(), torch.cuda.amp.autocast():
    for batch in tqdm(val_loader):
        pixel_values = batch["pixel_values"].to(device)
        # 批量生成decoder输入
        task_prompts = ["<s_fci>"] * pixel_values.shape[0]
        decoder_input_ids = processor.tokenizer(
            task_prompts, 
            add_special_tokens=False, 
            return_tensors="pt", 
            padding=True
        ).input_ids.to(device)
        
        # 批量推理
        outputs = model.generate(
            pixel_values,
            decoder_input_ids=decoder_input_ids,
            max_length=model.decoder.config.max_position_embeddings,
            early_stopping=True,
            pad_token_id=processor.tokenizer.pad_token_id,
            eos_token_id=processor.tokenizer.eos_token_id,
            use_cache=True,
            num_beams=1,
            bad_words_ids=[[processor.tokenizer.unk_token_id]],
            return_dict_in_generate=True,
        )
        
        # 批量解析结果
        seqs = processor.batch_decode(outputs.sequences)
        for seq, gt in zip(seqs, batch["labels"]):
            seq = seq.replace(processor.tokenizer.eos_token, "").replace(processor.tokenizer.pad_token, "")
            seq = re.sub(r"<.*?>", "", seq, count=1).strip()
            seq = processor.token2json(seq)
            seq["class"] = seq.get("class", "other")
            accs.append(float(seq["class"] == gt["class"]))

5. 使用Optimum库做生产级优化

利用HuggingFace Optimum将模型转换为ONNX格式,进一步提升推理速度:

from optimum.onnxruntime import ORTModelForVisionEncoderDecoder

# 导出为ONNX(仅需执行一次)
model = VisionEncoderDecoderModel.from_pretrained(CKPT_PATH, config=config)
model.save_pretrained("donut_onnx_temp")
ort_model = ORTModelForVisionEncoderDecoder.from_pretrained("donut_onnx_temp", export=True)
ort_model.save_pretrained("donut_onnx_optimized")

# 加载优化后的模型推理
ort_model = ORTModelForVisionEncoderDecoder.from_pretrained("donut_onnx_optimized")
ort_model.to(device)

内容的提问来源于stack exchange,提问作者Minh Ngo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 18:18:08