如何为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
相关产品推荐
相关产品推荐

