在AWS批量转换作业中为DistilBERT模型的文本输出添加ID
实现批量转换输出保留ID列的方案
可以实现,核心是通过自定义推理脚本处理输入数据,将原始id字段与提取的文本特征一同输出。以下是具体步骤和示例代码:
1. 创建自定义推理脚本(inference.py)
在你的代码目录下创建inference.py文件,覆盖默认推理逻辑以保留id字段:
from transformers import pipeline import json def model_fn(model_dir): # 加载预训练模型,与原配置保持一致 return pipeline( "feature-extraction", model="distilbert-base-uncased", device=-1 # CPU运行,用GPU则改为0 ) def predict_fn(input_data, model): # 提取输入文本做特征提取 text = input_data["text"] features = model(text)[0] # 合并原始id与特征结果返回 return { "id": input_data["id"], "features": features } def input_fn(input_data, content_type): # 解析JSON格式输入 if content_type == "application/json": return json.loads(input_data) raise ValueError(f"Unsupported content type: {content_type}") def output_fn(prediction, accept): # 将结果序列化为JSON输出 if accept == "application/json": return json.dumps(prediction), accept raise ValueError(f"Unsupported accept type: {accept}")
2. 修改AWS SageMaker管道配置
更新原有代码,在创建HuggingFaceModel时指定自定义推理脚本的路径:
hub = { 'HF_MODEL_ID':'distilbert-base-uncased', 'HF_TASK':'feature-extraction' } # 创建Hugging Face Model Class,引入自定义推理脚本 huggingface_model = HuggingFaceModel( env=hub, role=role, transformers_version="4.26", pytorch_version="1.13", py_version='py39', source_dir="./", # 推理脚本所在目录(当前目录填./即可) entry_point="inference.py" # 指定自定义脚本文件名 ) # 创建Transformer运行批量作业 batch_job = huggingface_model.transformer( instance_count=1, instance_type='ml.m5.xlarge', output_path=output_s3_path, strategy='SingleRecord' )
关键说明
source_dir:填写包含inference.py的目录路径,脚本在当前工作目录时直接填./。predict_fn是核心逻辑:从输入数据中取出id,与特征提取结果合并返回,确保输出同时包含两个字段。- 原有模型版本、实例配置无需改动,仅需添加自定义脚本的引用。
运行修改后的批量转换作业后,输出的JSON文件会同时包含id和features字段,对应每条输入数据的标识与提取的文本特征。
内容的提问来源于stack exchange,提问作者nana1vo
相关产品推荐
相关产品推荐

