React Native+YOLOv5转TFJS模型推理报错:pad参数类型异常
问题排查与修复方案
核心错误原因
- 张量类型不兼容:TFJS v3.19.0中部分内置操作(如
pad)对张量内部类型校验严格,原代码的张量处理流程导致张量被异常包装,触发类型误判。 - 未定义变量:直接使用
inputHeight/inputWidth但未赋值,导致输入尺寸错误。 - 重复加载模型:每次调用预测都重新加载模型,引发张量上下文冲突。
- 冗余张量操作:
resizeBilinear后已生成目标尺寸张量,重复reshape破坏了张量结构。 - 内存泄漏风险:未对中间张量进行销毁,导致React Native内存占用过高。
修复后的完整代码
import * as tf from '@tensorflow/tfjs'; import {bundleResourceIO, decodeJpeg} from '@tensorflow/tfjs-react-native'; // 全局模型实例,避免重复加载 let model: tf.GraphModel | null = null; const modelJSON = require('../assets/web_model/model.json'); const modelWeights = [ require('../assets/web_model/group1-shard1of7.bin'), require('../assets/web_model/group1-shard2of7.bin'), require('../assets/web_model/group1-shard3of7.bin'), require('../assets/web_model/group1-shard4of7.bin'), require('../assets/web_model/group1-shard5of7.bin'), require('../assets/web_model/group1-shard6of7.bin'), require('../assets/web_model/group1-shard7of7.bin'), ]; // 初始化模型(仅执行一次) const initModel = async () => { if (!model) { await tf.ready(); model = await tf.loadGraphModel(bundleResourceIO(modelJSON, modelWeights)); } }; const getPredictions = async (dataURL: string) => { // 确保模型已初始化 await initModel(); if (!model) throw new Error('模型加载失败'); return tf.tidy(() => { // 解析Base64图片 const imgB64 = dataURL.split(';base64,')[1]; const raw = tf.util.encodeString(imgB64, 'base64') as Uint8Array; const imagesTensor = decodeJpeg(raw); // 获取模型输入尺寸 const [inputHeight, inputWidth] = model.inputs[0].shape.slice(1, 3) as [number, number]; // 处理输入张量 let processedTensor = tf.image.resizeBilinear(imagesTensor, [inputHeight, inputWidth]) as tf.Tensor<tf.Rank.R3>; processedTensor = tf.cast(processedTensor, 'float32'); processedTensor = tf.div(processedTensor, 255.0); processedTensor = tf.expandDims(processedTensor, 0); // 添加batch维度 // 执行推理 return model.execute(processedTensor) as tf.Tensor[]; }); }; export default getPredictions;
关键修改说明
- 全局模型单例:将模型加载移到全局初始化函数,避免每次预测重复加载,解决上下文冲突。
- 变量赋值修正:从模型输入形状中直接获取
inputHeight/inputWidth,解决未定义变量问题。 - 移除冗余reshape:
resizeBilinear已输出[inputHeight, inputWidth, 3]的张量,无需重复reshape。 - tf.tidy内存管理:用
tf.tidy包裹所有张量操作,自动清理中间张量,避免内存泄漏,同时确保张量类型符合TFJS要求。 - 简化张量解析:直接将
tf.util.encodeString的结果转为Uint8Array,避免多余的buffer转换步骤。
内容的提问来源于stack exchange,提问作者gildniy
相关产品推荐
相关产品推荐

