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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 16:17:36