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

如何解决SageMaker部署HuggingFace多模态模型的推理序列化问题?

问题分析与解决方案

核心问题排查

  • 你自定义的TextImageSerializer强制要求只能传question或image中的一个,但Donut模型需要同时接收文本问题和图像输入才能完成文档问答任务
  • JSON无法直接序列化二进制数据,直接把图片二进制塞到JSON里会报错,必须转成base64编码
  • HF_TASK设置为question-answering不符合Donut的任务类型,它属于自定义多模态任务,Hugging Face默认推理容器不支持直接用该任务类型,需要自定义推理脚本处理数据

解决方案步骤

1. 修正序列化器,支持同时传文本和图片

修改序列化器,同时处理文本和图片,将图片转为base64字符串以兼容JSON格式:

from sagemaker.serializers import BaseSerializer
import json
import base64
from io import BytesIO
from PIL import Image

class TextImageSerializer(BaseSerializer):
    CONTENT_TYPE = 'application/json'

    def serialize(self, data: dict) -> bytes:
        # 确保同时提供question和image字段
        if not (data.get('question') and data.get('image')):
            raise ValueError("必须同时提供'question'和'image'字段")
        
        # 处理图片:二进制转base64字符串
        image_data = data['image']
        if isinstance(image_data, bytes):
            img = Image.open(BytesIO(image_data))
            buffer = BytesIO()
            img.save(buffer, format='PNG')
            img_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
        else:
            raise TypeError("image必须是二进制字节流")
        
        # 构造Donut要求的输入格式
        payload = {
            "prompt": data['question'],
            "image": img_base64
        }
        return json.dumps(payload).encode('utf-8')

2. 编写自定义推理脚本inference.py

Donut的预处理、推理逻辑和标准任务不同,需要自定义脚本适配:

from transformers import DonutProcessor, VisionEncoderDecoderModel
import torch
import base64
from io import BytesIO
from PIL import Image

processor = None
model = None

def model_fn(model_dir):
    # 加载模型和处理器
    processor = DonutProcessor.from_pretrained(model_dir)
    model = VisionEncoderDecoderModel.from_pretrained(model_dir)
    device = "cuda" if torch.cuda.is_available() else "cpu"
    model.to(device)
    return processor, model

def predict_fn(input_data, model):
    processor, model = model
    device = "cuda" if torch.cuda.is_available() else "cpu"
    
    # 解析base64编码的图片
    image_base64 = input_data['image']
    image_bytes = base64.b64decode(image_base64)
    image = Image.open(BytesIO(image_bytes)).convert("RGB")
    
    # 构造模型输入
    prompt = input_data['prompt']
    decoder_input_ids = processor.tokenizer(prompt, add_special_tokens=False, return_tensors="pt").input_ids
    pixel_values = processor(image, return_tensors="pt").pixel_values
    
    # 执行推理
    outputs = model.generate(
        pixel_values.to(device),
        decoder_input_ids=decoder_input_ids.to(device),
        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,
    )
    
    # 解析推理结果
    prediction = processor.batch_decode(outputs.sequences)[0]
    prediction = processor.token2json(prediction)
    return prediction

3. 修改部署代码,指定推理脚本

将自定义推理脚本和模型一起部署,调整模型初始化参数:

from sagemaker.huggingface import HuggingFaceModel
from sagemaker.serverless import ServerlessInferenceConfig

# Hub模型配置
hub = {
    'HF_MODEL_ID': 'naver-clova-ix/donut-base-finetuned-docvqa',
    'HF_TASK': 'custom'  # 指定为自定义任务,使用自定义推理脚本
}

# 初始化模型,指定推理脚本所在目录
huggingface_model = HuggingFaceModel(
   env=hub,                      
   role=role,                    
   transformers_version="4.26.0",  # 升级到兼容Donut的版本,避免API兼容问题
   pytorch_version="1.13.1",      
   py_version='py38',
   source_dir=".",  # 推理脚本所在的本地目录
   entry_point="inference.py"  # 自定义推理脚本文件名
)

serverless_config = ServerlessInferenceConfig(
    memory_size_in_mb=4096,  # Donut模型内存需求较高,建议设置为4096MB
    max_concurrency=10,
)

# 部署predictor,使用修正后的序列化器
predictor = huggingface_model.deploy(
    serializer=TextImageSerializer(),
    serverless_inference_config=serverless_config,
)

4. 测试推理

构造符合Donut要求的payload并执行推理:

# 构造payload,注意Donut需要特定的prompt前缀
payload = {
    'question': "<s_docvqa><s_question>what is the document title</s_question><s_answer>",
    'image': open("document.png", "rb").read()
}

# 执行推理并打印结果
res = predictor.predict(data=payload)
print(res)

额外注意事项

  • 确保transformers和pytorch版本兼容Donut模型,旧版本可能存在API差异
  • 服务器内存建议设置为4096MB,避免模型运行时内存不足
  • 输入图片需转为RGB格式,尺寸不要过大,避免推理超时

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 06:32:08