TorchServe如何高效将gRPC返回的bytes转换为torch.Tensor
通过TorchServe部署模型,采用gRPC协议与服务端通信。自定义处理器中定义的postprocess方法最终返回列表,该列表会被转换为字节流用于网络传输,对应服务端postprocess实现如下:
def postprocess(self, data): # data type - torch.Tensor # data shape - [1, 17, 80, 64] and data dtype - torch.float32 return data.tolist()
当前性能瓶颈位于客户端侧:现有将TorchServe返回的字节流转换为torch.Tensor的实现依赖ast.literal_eval,转换效率极低——gRPC推理请求本身仅耗时0.3秒,但字节解析转张量的操作耗时高达0.84秒,对应客户端代码如下:
# 推理请求耗时0.3秒 response = self.inference_stub.Predictions( inference_pb2.PredictionsRequest(model_name=model_name, input=input_data)) # 字节解析转张量耗时0.84秒 predictions = torch.as_tensor(literal_eval( response.prediction.decode('utf-8')))
尝试使用numpy.frombuffer或torch.frombuffer直接解析字节流时均抛出错误,numpy侧报错信息如下:
import numpy as np np.frombuffer(response.prediction) Traceback (most recent call last): File "<string>", line 1, in <module> ValueError: buffer size must be a multiple of element size np.frombuffer(response.prediction, dtype=np.float32) Traceback (most recent call last): File "<string>", line 1, in <module> ValueError: buffer size must be a multiple of element size
torch侧报错信息如下:
import torch torch.frombuffer(response.prediction, dtype = torch.float32) Traceback (most recent call last): File "<string>", line 1, in <module> ValueError: buffer length (2601542 bytes) after offset (0 bytes) must be a multiple of element size (4)
需要找到更高效的方案,实现接收字节流到torch.Tensor的转换。
- 当前
postprocess返回Python原生列表时,TorchServe默认会将列表序列化为UTF-8编码的JSON字符串再放入gRPC响应,返回的根本不是张量的原始二进制字节,因此直接调用frombuffer解析必然报错——拿到的字节流是包含括号、逗号、数字字符的文本内容,不是连续排列的float32二进制值。 - 数值校验可直接验证该结论:目标张量总元素数为
1*17*80*64 = 87040,如果是原始float32二进制格式,总字节数应为87040 * 4 = 348160字节,但实际拿到的buffer长度为2601542字节,远大于原始二进制体积,完全符合JSON文本序列化的体积特征。 ast.literal_eval性能差的核心原因是需要逐字符解析嵌套列表的JSON结构,逐层构造Python对象后再转张量,冗余开销极高。
方案1:修改服务端返回原始二进制字节(性能最高,推荐)
直接调整服务端postprocess逻辑,不要返回Python列表,改为返回序列化后的原始二进制字节,同时附带shape、dtype元信息,客户端拿到后可近乎零拷贝直接转张量,解析耗时可降到毫秒级。
服务端修改示例:
import numpy as np def postprocess(self, data): # data为shape [1,17,80,64]的float32张量 # 转numpy数组后提取原始字节,同步附带元信息 return [ data.numpy().tobytes(), # 张量原始二进制字节 np.array(data.shape, dtype=np.int32).tobytes(), # 张量shape的二进制 b"float32" # 张量dtype标识 ]
客户端解析示例:
import torch import numpy as np from google.protobuf.json_format import MessageToDict # 提取响应内容 resp_list = MessageToDict(response)["prediction"] # 解析张量shape shape = tuple(np.frombuffer(resp_list[1].encode('latin1'), dtype=np.int32).tolist()) # 直接从字节构造张量,无需逐字符解析 predictions = torch.frombuffer( resp_list[0].encode('latin1'), dtype=torch.float32 ).reshape(shape)
注意:gRPC传输字节时默认做latin1编码,客户端解码时必须用latin1编码反解回原始字节,不能用utf-8,否则会破坏二进制数据。
该方案下解析耗时通常在1ms以内,相比原有literal_eval方案性能提升数百倍。
方案2:不修改服务端的快速优化方案
如果暂时无法调整服务端部署逻辑,可替换ast.literal_eval为orjson、ujson这类高性能JSON解析库,解析速度可提升3-5倍,改造成本极低。
示例代码:
import orjson import torch # 字节解码为字符串后用高性能JSON库解析,再转张量 predictions = torch.as_tensor( orjson.loads(response.prediction.decode('utf-8')), dtype=torch.float32 )
该方案无需改动服务端,但性能上限低于直接传输原始二进制的方案。
内容的提问来源于stack exchange,提问作者Mohit Motwani

