如何为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响应格式的结果。
注意事项
- 确保TorchServe运行环境已安装
transformers库,否则tokenizer无法加载。 - 若你trace模型时使用了不同的输入参数结构,需对应调整
preprocess方法的返回内容。 - 生成模型归档时,确保Handler文件名与你执行
torch-model-archiver命令时指定的handler.py一致。
内容的提问来源于stack exchange,提问作者maxwellspi
相关产品推荐
相关产品推荐

