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

如何为AWS SageMaker部署的HuggingFace token分类模型传递输入及编写推理代码?

关于SageMaker HuggingFace端点的输入格式与推理代码配置

一、predict方法的输入数据结构

由于你使用的是需要图像、文本词、边界框的多模态token分类模型(如LayoutLMv3),SageMaker的predict方法接受JSON格式的输入,结构需与模型输入要求严格匹配,示例如下:

单样本输入

{
  "image": "base64编码的图像字符串",
  "words": ["发票", "编号", "NO.12345", ...],
  "bbox": [[0, 0, 60, 25], [65, 0, 120, 25], ...]
}
  • image:必须转为Base64编码字符串(避免二进制传输问题)
  • bbox:每个元素为[x_min, y_min, x_max, y_max],需与训练时的标注格式一致

批量样本输入

若要批量预测,可传入包含多个样本的列表:

[
  {
    "image": "base64_str_1",
    "words": ["样本1的词序列"],
    "bbox": [[...]]
  },
  {
    "image": "base64_str_2",
    "words": ["样本2的词序列"],
    "bbox": [[...]]
  }
]

二、在train.py中添加推理代码

SageMaker HuggingFace端点依赖model_fn、input_fn、predict_fn三个核心函数实现推理逻辑,以下是适配你的模型的完整代码示例:

1. 导入必要依赖

import torch
import base64
import json
import io
from PIL import Image
from transformers import AutoProcessor, AutoModelForTokenClassification

# 全局变量,用于加载处理器与模型
processor = None
model = None

2. 实现model_fn:加载模型资产

端点启动时自动调用,用于加载训练好的模型和处理器

def model_fn(model_dir):
    global processor, model
    # 加载训练时使用的processor和模型
    processor = AutoProcessor.from_pretrained(model_dir, apply_ocr=False)
    model = AutoModelForTokenClassification.from_pretrained(model_dir)
    return model

3. 实现input_fn:解析并预处理输入数据

负责将JSON输入转换为模型可接受的张量格式

def input_fn(request_body, request_content_type):
    assert request_content_type == "application/json"
    data = json.loads(request_body)

    # 处理批量或单样本输入
    if isinstance(data, list):
        images = []
        words_list = []
        bbox_list = []
        for sample in data:
            # 解码Base64图像
            img_bytes = base64.b64decode(sample["image"])
            img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
            images.append(img)
            words_list.append(sample["words"])
            bbox_list.append(sample["bbox"])
        # 批量预处理
        encoding = processor(images, words_list, boxes=bbox_list, truncation=True, padding="max_length", return_tensors="pt")
    else:
        # 单样本预处理
        img_bytes = base64.b64decode(data["image"])
        img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
        encoding = processor(img, data["words"], boxes=data["bbox"], truncation=True, padding="max_length", return_tensors="pt")
    
    return encoding

4. 实现predict_fn:执行模型推理

def predict_fn(input_data, model):
    model.eval()
    with torch.no_grad():
        outputs = model(**input_data)
    
    # 获取token分类预测结果
    logits = outputs.logits
    predictions = torch.argmax(logits, dim=-1).tolist()

    # 加载训练时保存的标签映射(需提前保存)
    with open(f"{model_dir}/labels.json", "r") as f:
        id2label = json.load(f)
    
    # 将预测ID转换为实际标签,过滤padding部分
    formatted_predictions = []
    for pred in predictions:
        filtered_pred = [id2label[str(p)] for p in pred if p != model.config.pad_token_id]
        formatted_predictions.append(filtered_pred)
    
    return formatted_predictions

5. 可选:实现output_fn:格式化输出结果

若需统一输出格式,可添加此函数:

def output_fn(prediction, response_content_type):
    assert response_content_type == "application/json"
    return json.dumps(prediction), response_content_type

三、关键注意事项

  • 训练时需将id2label字典保存到模型目录,方便推理时映射标签,代码示例:
    import os
    import json
    # 在训练代码的保存环节添加
    with open(os.path.join(args.output_dir, "labels.json"), "w") as f:
        json.dump(model.config.id2label, f)
    
  • 若需要额外依赖(如pillow),可在source_dir目录下创建requirements.txt,列出依赖包名
  • 本地测试输入时,可通过以下代码生成图像的Base64编码:
    import base64
    with open("test_image.jpg", "rb") as f:
        base64_str = base64.b64encode(f.read()).decode("utf-8")
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 03:15:00