React Native中TensorFlow.js模型预测图像张量报错求助
问题背景
将自定义训练的Keras模型转换为TensorFlow.js模型后,在React Native项目中结合Expo Camera捕获图像进行损伤预测时,调用模型predict方法出现错误,无法得到预期的预测结果数组。
使用版本
- React Native项目依赖:
- expo: >=45.0.0-0 <46.0.0
- expo-camera: ~12.2.0
- expo-gl: ~11.3.0
- expo-gl-cpp: ~11.3.0
- @tensorflow/tfjs: ^4.0.0
- @tensorflow/tfjs-react-native: ^0.8.0
- react-native-fs: ^2.20.0
- Python环境:tensorflow 2.10.0
模型转换命令
tensorflowjs_converter --input_format=keras --weight_shard_size_bytes=419430400 --quantize_float16=* /path/to/model.h5 /path/to/output
React Native中模型加载代码
// Load layers model using model json and weights file const models = await tf.loadLayersModel(bundleResourceIO(modelJSON, weights));
图像张量日志
LOG imageAsTensors: {"kept":false,"isDisposedInternal":false,"shape":[224,224,3],"dtype":"int32","size":150528,"strides":[672,3],"dataId":{"id":670},"id":980,"rankType":"3"} LOG imageTensorReshaped: {"kept":false,"isDisposedInternal":false,"shape":[1,224,224,3],"dtype":"int32","size":150528,"strides":[150528,672,3],"dataId":{"id":670},"id":981,"rankType":"4","scopeId":408}
预测代码
try { // predict against the model const output = await models.predict(imageTensorReshaped, { batchSize: 1 }); return output.dataSync(); } catch (error) { console.log('Error predicting from tensor image', error); }
错误信息
Error predicting from tensor image [TypeError: null is not an object (evaluating 'opHandler.clone')]
解决方案
针对该错误,可按以下步骤排查修复:
- 统一张量数据类型
模型训练时通常使用float32类型,但当前图像张量为int32,类型不匹配会引发运算异常。需将图像张量转换为float32并执行训练时一致的归一化操作:
// 在reshape后添加类型转换与归一化 const imageTensorProcessed = imageTensorReshaped.cast('float32').div(255.0); // 使用处理后的张量执行预测 const output = await models.predict(imageTensorProcessed, { batchSize: 1 });
- 修复TF.js版本兼容性
@tensorflow/tfjs-react-native@0.8.0与@tensorflow/tfjs@4.0.0版本跨度较大,存在兼容问题。建议降级tfjs到匹配版本:
npm install @tensorflow/tfjs@3.21.0 @tensorflow/tfjs-react-native@0.8.0
- 调整模型转换参数
--quantize_float16=*参数可能导致模型权重类型与tfjs-react-native环境不兼容,尝试移除量化参数重新转换模型:
tensorflowjs_converter --input_format=keras --weight_shard_size_bytes=419430400 /path/to/model.h5 /path/to/output
- 确认模型加载状态
在调用predict前,确保模型已完全加载:
if (!models || !models.weights || models.weights.length === 0) { throw new Error('Model not loaded properly'); }
- 优化张量内存管理
在React Native环境中,用tf.tidy包裹运算逻辑,避免内存泄漏引发异常:
try { const result = tf.tidy(() => { const processed = imageTensorReshaped.cast('float32').div(255.0); const output = models.predict(processed, { batchSize: 1 }); return output.dataSync(); }); return result; } catch (error) { console.log('Error predicting from tensor image', error); }
内容的提问来源于stack exchange,提问作者Priyanka Dhumane
相关产品推荐
相关产品推荐

