RTMPose Body2D模型ONNX转TFLite时出现张量分配溢出错误排查
解决RTMPose Body2D转TFLite时的张量分配溢出问题
错误根源
报错BytesRequired number of elements overflowed伴随MAX_POOL_2D节点初始化失败,本质是TFLite处理模型时,某张量的形状计算超出了整数范围,主要由两个转换环节的问题导致:
- RTMPose ONNX模型采用NCHW维度顺序(
[batch, channel, height, width]),但TensorFlow/TFLite默认使用NHWC,ONNX转SavedModel时未做显式处理,引发形状计算异常。 - 原生
onnx-tf转换工具对带动态形状的池化算子兼容性不足,生成的SavedModel存在隐性形状错误,转TFLite后无法正确分配内存。
分步解决方案
1. 修正ONNX到SavedModel的转换逻辑
显式指定维度顺序并放宽转换严格性,避免维度自动转换引发的问题:
import onnx from onnx_tf.backend import prepare def convert_onnx_to_saved_model(): model = onnx.load("model.onnx") # strict=False兼容更多ONNX算子,device指定CPU避免GPU相关转换问题 tf_rep = prepare(model, strict=False, device='CPU') tf_rep.export_graph("model")
2. 优化TFLite转换配置
启用算子兼容选项并固定输入形状,消除动态形状带来的计算溢出风险:
import tensorflow as tf def convert_saved_model_to_tflite(): converter = tf.lite.TFLiteConverter.from_saved_model("model") # 允许使用TensorFlow原生算子,覆盖TFLite不支持的算子 converter.target_spec.supported_ops = [ tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS ] converter.experimental_enable_resource_variables = True # 启用默认优化,同时压缩模型并修正形状问题 converter.optimizations = [tf.lite.Optimize.DEFAULT] tflite_model = converter.convert() with open("model.tflite", "wb") as f: f.write(tflite_model)
3. 调整TFLite推理的输入维度匹配
确保输入数据的维度顺序与TFLite模型要求一致:
import numpy as np import tensorflow as tf def load_and_test_tflite(model_path): interpreter = tf.lite.Interpreter(model_path=model_path) print("Interpreter successfully created.") interpreter.allocate_tensors() print("Tensors successfully allocated.") input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() print("Input details:", input_details) print("Output details:", output_details) input_shape = input_details[0]['shape'] # 生成匹配维度的测试输入,注意保持与模型要求的NCHW/NHWC一致 input_data = np.random.random_sample(input_shape).astype(np.float32) interpreter.set_tensor(input_details[0]['index'], input_data) interpreter.invoke() print("Inference successfully run.") output_data = interpreter.get_tensor(output_details[0]['index']) print("Output data shape:", output_data.shape)
额外验证步骤
- 用
onnxruntime先验证原始ONNX模型的推理结果,确认模型本身无问题:import onnxruntime as ort sess = ort.InferenceSession("model.onnx") input_name = sess.get_inputs()[0].name output_names = [o.name for o in sess.get_outputs()] test_input = np.random.randn(1,3,256,192).astype(np.float32) outputs = sess.run(output_names, {input_name: test_input}) print("ONNX模型输出形状:", [o.shape for o in outputs]) - 确认依赖版本兼容性:推荐使用TensorFlow 2.10+、onnx-tf 1.10+、onnxruntime 1.13+。
内容的提问来源于stack exchange,提问作者samedhrmn
相关产品推荐
相关产品推荐

