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

如何为TorchServe部署HuggingFace模型创建自定义Handler

适配Flan-T5 TorchScript模型的TorchServe Handler代码

以下是针对你导出的Flan-T5 TorchScript模型编写的TorchServe Handler代码,完全匹配模型输入输出结构:

import torch
from ts.torch_handler.base_handler import BaseHandler
from transformers import T5Tokenizer

class FlanT5Handler(BaseHandler):
    def initialize(self, context):
        # 初始化模型与tokenizer
        self.manifest = context.manifest
        properties = context.system_properties
        model_dir = properties.get("model_dir")
        
        # 加载TorchScript格式模型
        self.model = torch.jit.load(f"{model_dir}/traced_flan-t5-base.pt")
        self.model.eval()
        
        # 加载对应tokenizer
        self.tokenizer = T5Tokenizer.from_pretrained("google/flan-t5-base")
        
        # 设置计算设备(GPU/CPU)
        self.device = torch.device("cuda:" + str(properties.get("gpu_id")) if torch.cuda.is_available() else "cpu")
        self.model.to(self.device)
        
        self.initialized = True

    def preprocess(self, data):
        # 解析请求中的输入文本
        input_texts = []
        for req in data:
            input_content = req.get("data") or req.get("body")
            if isinstance(input_content, bytes):
                input_content = input_content.decode("utf-8")
            input_texts.append(input_content)
        
        # 对输入文本进行token化处理
        tokenized_input = self.tokenizer(
            input_texts,
            padding=True,
            truncation=True,
            return_tensors="pt"
        )
        
        # 构造模型trace时要求的decoder初始输入(使用bos token)
        decoder_input_ids = torch.tensor([[self.tokenizer.bos_token_id]] * len(input_texts)).long()
        
        # 返回与模型trace输入匹配的张量元组
        return (
            tokenized_input["input_ids"].to(self.device),
            tokenized_input["attention_mask"].to(self.device),
            decoder_input_ids.to(self.device)
        )

    def inference(self, inputs):
        # 执行推理计算
        with torch.no_grad():
            outputs = self.model(*inputs)
        
        # 从模型输出logits中获取预测token id
        predicted_ids = torch.argmax(outputs.logits, dim=-1)
        return predicted_ids

    def postprocess(self, inference_output):
        # 将预测token id解码为自然语言文本
        results = self.tokenizer.batch_decode(inference_output, skip_special_tokens=True)
        return [{"generated_text": result} for result in results]

关键说明

  • initialize方法:在TorchServe启动模型时执行,负责加载模型、tokenizer并设置计算设备,确保模型处于eval模式。
  • preprocess方法:解析客户端请求的文本输入,转换为模型所需的张量格式,严格匹配你trace模型时的输入结构(input_ids、attention_mask、decoder_input_ids)。
  • inference方法:调用模型执行推理,使用torch.no_grad()禁用梯度计算,提升推理效率。
  • postprocess方法:将模型输出的token id解码为可读文本,返回符合API响应格式的结果。

注意事项

  1. 确保TorchServe运行环境已安装transformers库,否则tokenizer无法加载。
  2. 若你trace模型时使用了不同的输入参数结构,需对应调整preprocess方法的返回内容。
  3. 生成模型归档时,确保Handler文件名与你执行torch-model-archiver命令时指定的handler.py一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 14:05:30