TorchServe自定义处理器:传递张量列表实现批量推理
TorchServe自定义处理器批量推理实现方案
一、编写支持批量处理的自定义处理器
要实现批量推理,自定义处理器需继承BaseHandler,重写预处理、推理、后处理逻辑,核心是把多个请求的输入打包成统一batch,而非逐个循环处理。
示例代码:
from ts.torch_handler.base_handler import BaseHandler import torch from transformers import TransfoXLModel, TransfoXLTokenizer class TransfoXLHandler(BaseHandler): def initialize(self, context): # 初始化模型与分词器 self.tokenizer = TransfoXLTokenizer.from_pretrained('transfo-xl-wt103') self.model = TransfoXLModel.from_pretrained('transfo-xl-wt103') self.model.eval() self.device = torch.device('cpu') # 按需切换为cuda self.model.to(self.device) def preprocess(self, requests): # 批量预处理:提取所有请求文本,统一转成带padding的张量 texts = [req['body']['text'] for req in requests] inputs = self.tokenizer(texts, padding=True, truncation=True, return_tensors='pt') return inputs.to(self.device) def inference(self, inputs): # 批量推理:直接传入整个batch,无需循环 with torch.no_grad(): outputs = self.model(**inputs) return outputs def postprocess(self, outputs): # 批量后处理:将模型输出拆分回单个请求的结果 results = [] for hidden_state in outputs.last_hidden_state: # 按业务需求处理单条结果,此处示例返回隐藏层张量 results.append({'last_hidden_state': hidden_state.cpu().numpy().tolist()}) return results
二、配置batch_size与max_batch_delay
TorchServe的批量参数可通过配置文件或启动命令设置:
1. 配置文件方式
在config.properties中添加以下配置:
enable_batching=true batch_size=8 max_batch_delay=100
enable_batching:必须设为true以启用批量推理batch_size:单个批量的最大请求数,可根据CPU核心数、内存容量调整(8核CPU可先试8或16)max_batch_delay:等待凑齐批量的最长时间(毫秒),若超时未凑够batch_size,则直接处理当前已收集的请求
2. 启动命令参数方式
启动TorchServe时直接指定参数:
torchserve --start --ncs --model-store model_store --models transfo-xl=transfo-xl.mar --enable-batching --batch-size 8 --max-batch-delay 100
三、CPU使用率问题解析
你提到逐个处理时CPU使用率高,多进程反而效果差,核心原因是transformers库的模型在CPU上默认启用了操作内并行(intra-op parallelism),比如通过OpenMP优化,模型内部运算已自动利用多核心。
Python多进程会导致模型重复加载,且多进程并行与模型内部的并行机制冲突,反而降低效率。TorchServe的批量推理是在单进程内将多个请求打包成batch,让模型一次性处理,既能充分利用模型自身的并行能力,又避免了多进程的额外开销,能有效拉满CPU使用率。
四、批量推理测试
可通过批量发送请求验证效果:
# 后台并行发送多个请求,触发TorchServe批量处理 curl -X POST http://localhost:8080/predictions/transfo-xl -H "Content-Type: application/json" -d '{"text": "第1页OCR文本"}' & curl -X POST http://localhost:8080/predictions/transfo-xl -H "Content-Type: application/json" -d '{"text": "第2页OCR文本"}' & curl -X POST http://localhost:8080/predictions/transfo-xl -H "Content-Type: application/json" -d '{"text": "第3页OCR文本"}' &
内容的提问来源于stack exchange,提问作者Paras Bansal
相关产品推荐
相关产品推荐

