在Nvidia Triton Server多模型部署中,能否直接传入文本调用模型?
直接传入文本调用Triton多模型端点的方案
完全支持直接传入文本调用模型,核心是通过Triton的Python Backend封装文本tokenization和模型推理的完整流程,无需在客户端提前处理文本。以下是具体配置和调用示例:
一、调整模型配置与结构
1. 模型目录结构
采用Triton Python Backend要求的目录结构:
bert-model/ ├── 1/ │ └── model.py └── config.pbtxt
2. 配置config.pbtxt
明确输入为字符串类型,输出匹配你的BERT模型实际维度:
name: "bert-model" platform: "python_backend" max_batch_size: 0 input [ { name: "INPUT_TEXT" data_type: TYPE_STRING dims: [1] # 单条文本输入,批量场景可改为[-1] } ] output [ { name: "OUTPUT_1" data_type: TYPE_FP32 dims: [2] # 替换为你的BERT模型实际输出维度 } ]
3. 编写model.py(封装预处理与推理)
这个文件负责加载tokenizer、模型,接收文本后自动完成tokenization并执行推理:
import torch from transformers import BertTokenizer, BertForSequenceClassification import triton_python_backend_utils as pb_utils class TritonPythonModel: def initialize(self, args): # 加载预训练的tokenizer和模型 self.tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') self.model = BertForSequenceClassification.from_pretrained('bert-base-uncased') self.model.eval() # 绑定到GPU(GPU实例下自动生效) self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self.model.to(self.device) def execute(self, requests): responses = [] for request in requests: # 提取输入文本 input_text = pb_utils.get_input_tensor_by_name(request, "INPUT_TEXT").as_numpy()[0].decode('utf-8') # 执行文本tokenization tokenized_input = self.tokenizer( input_text, padding='max_length', max_length=128, truncation=True, return_tensors='pt' ) input_ids = tokenized_input['input_ids'].to(self.device) attention_mask = tokenized_input['attention_mask'].to(self.device) # 模型推理(禁用梯度计算提升性能) with torch.no_grad(): outputs = self.model(input_ids, attention_mask=attention_mask) logits = outputs.logits.cpu().numpy() # 构造响应张量 output_tensor = pb_utils.Tensor("OUTPUT_1", logits) responses.append(pb_utils.InferenceResponse(output_tensors=[output_tensor])) return responses
二、打包部署
将bert-model目录打包为bert-model.tar.gz,上传到S3存储桶,然后通过SageMaker多模型端点部署该模型包。
三、直接传入文本的调用代码
修改客户端调用逻辑,直接传入原始文本:
import boto3 import json # 初始化SageMaker Runtime客户端 sm_client = boto3.client('sagemaker-runtime') endpoint_name = "your-triton-mme-endpoint-name" # 待推理的原始文本 input_text = "Triton Inference Server provides a cloud and edge inferencing solution optimized for both CPUs and GPUs." # 构造请求payload payload = { "inputs": [ {"name": "INPUT_TEXT", "shape": [1], "datatype": "BYTES", "data": [input_text]} ] } # 调用端点 response = sm_client.invoke_endpoint( EndpointName=endpoint_name, ContentType="application/json", Body=json.dumps(payload), TargetModel="bert-model.tar.gz" ) # 解析并打印结果 inference_result = json.loads(response['Body'].read().decode('utf-8')) print("推理输出:", inference_result['outputs'][0]['data'])
关键注意事项
- 依赖管理:确保Triton容器安装了
transformers、torch等依赖,可通过自定义Docker镜像提前预装,避免运行时依赖缺失。 - 批量处理:若需批量传入多条文本,可将config.pbtxt的
dims改为[-1],并在payload的data字段传入文本列表。 - 性能优化:GPU实例下,确保tokenization后的张量直接传入GPU,避免CPU-GPU数据拷贝损耗性能。
内容的提问来源于stack exchange,提问作者haju
相关产品推荐
相关产品推荐

