TensorFlow.js模型在Teachable Machine可用但React Native应用中预测失效
问题分析与解决方案
问题背景
我训练了一个识别乐高零件的模型,在Google Teachable Machine测试时能100%准确预测,但部署到Expo Go的React Native应用后,预测结果几乎全错。比如拍摄绿色零件,模型却预测为红色(预测顺序是grey、tan、red、green),推测问题出在张量图像处理环节,但不确定具体原因。
核心问题排查
你的代码存在几个关键问题,直接导致移动端预测失效:
- 模型重复加载:
handleCameraStream中每次循环都重新加载模型,不仅浪费资源,还会导致模型状态不稳定。 - 缺少Teachable Machine要求的图像预处理:Teachable Machine训练时,输入图像会被归一化到
[0,1]范围且默认使用RGB通道顺序,但相机输入的张量是[0,255]的uint8格式,通道顺序也可能不匹配。 - 张量资源管理混乱:循环中未正确释放所有临时张量,易引发内存泄漏,干扰预测逻辑。
修正步骤
- 提前加载模型:在组件初始化时仅加载一次模型,避免重复加载导致的状态异常。
- 添加标准化图像预处理:将相机输入张量转换为模型要求的格式(归一化、通道匹配)。
- 规范张量资源释放:使用
tf.tidy和tf.dispose自动清理临时张量,避免内存泄漏。 - 对齐输入格式:确保输入张量的形状、数值范围与Teachable Machine训练时完全一致。
修正后的代码
import React, {useRef, useState, useEffect} from 'react'; import {View,StyleSheet,Dimensions,Pressable,Modal,Text,ActivityIndicator,} from 'react-native'; import * as MediaLibrary from 'expo-media-library'; import {getModel,convertBase64ToTensor,startPrediction} from '../../helpers/tensor-helper'; import {cropPicture} from '../../helpers/image-helper'; import {Camera} from 'expo-camera'; import * as tf from "@tensorflow/tfjs"; import { cameraWithTensors } from '@tensorflow/tfjs-react-native'; import {bundleResourceIO, decodeJpeg} from '@tensorflow/tfjs-react-native'; const initialiseTensorflow = async () => { await tf.ready(); await tf.setBackend('rn-webgl'); // 明确设置后端,保障兼容性 } const TensorCamera = cameraWithTensors(Camera); const modelJson = require('../../model/model.json'); const modelWeights = require('../../model/weights.bin'); const modelMetaData = require('../../model/metadata.json'); const RESULT_MAPPING = ['grey', 'tan', 'red','green']; const CameraScreen = () => { const [hasCameraPermission, setHasCameraPermission] = useState(); const [hasMediaLibraryPermission, setHasMediaLibraryPermission] = useState(); const [isProcessing, setIsProcessing] = useState(false); const [presentedShape, setPresentedShape] = useState(''); const [model, setModel] = useState<tf.LayersModel | null>(null); // 存储模型实例 useEffect(() => { (async () => { const cameraPermission = await Camera.requestCameraPermissionsAsync(); const mediaLibraryPermission = await MediaLibrary.requestPermissionsAsync(); setHasCameraPermission(cameraPermission.status === "granted"); setHasMediaLibraryPermission(mediaLibraryPermission.status === "granted"); // 初始化TensorFlow并一次性加载模型 await initialiseTensorflow(); const loadedModel = await tf.loadLayersModel(bundleResourceIO(modelJson, modelWeights)); setModel(loadedModel); })(); }, []); if (hasCameraPermission === undefined) { return <Text>请求权限中...</Text> } else if (!hasCameraPermission) { return <Text>相机权限未授予,请在设置中修改。</Text> } let frame = 0; const computeRecognitionEveryNFrames = 60; const handleCameraStream = async (images: IterableIterator<tf.Tensor3D>) => { if (!model) return; // 模型未加载时跳过预测 const loop = async () => { if(frame % computeRecognitionEveryNFrames === 0){ const nextImageTensor = images.next().value; if(nextImageTensor){ tf.tidy(() => { // 自动清理临时张量 // 图像预处理:转换为float32并归一化到[0,1],匹配模型输入要求 const preprocessedTensor = nextImageTensor .cast('float32') .div(tf.scalar(255)) .reshape([1, 224, 224, 3]); // 执行预测并解析结果 const prediction = model.predict(preprocessedTensor) as tf.Tensor; prediction.data().then((data) => { const predictedIndex = data.indexOf(Math.max(...data)); const predictedLabel = RESULT_MAPPING[predictedIndex]; console.log(`预测结果: ${predictedLabel}`, data); setPresentedShape(predictedLabel); setIsProcessing(true); }); }); tf.dispose(nextImageTensor); // 释放原始图像张量 } } frame += 1; frame = frame % computeRecognitionEveryNFrames; requestAnimationFrame(loop); } loop(); } return ( <View style={styles.container}> <Modal visible={isProcessing} transparent={true} animationType="slide"> <View style={styles.modal}> <View style={styles.modalContent}> <Text>当前识别的零件: {presentedShape}</Text> {presentedShape === '' && <ActivityIndicator size="large" />} <Pressable style={styles.dismissButton} onPress={() => { setPresentedShape(''); setIsProcessing(false); }}> <Text>关闭</Text> </Pressable> </View> </View> </Modal> <TensorCamera style={styles.camera} type={Camera.Constants.Type.back} onReady={handleCameraStream} resizeHeight={224} resizeWidth={224} resizeDepth={3} autorender={true} cameraTextureHeight={1920} cameraTextureWidth={1080} /> </View> ); }; const styles = StyleSheet.create({ container: { flex: 1, }, camera: { flex: 1, width: Dimensions.get('window').width, height: Dimensions.get('window').height, }, modal: { flex: 1, justifyContent: 'center', alignItems: 'center', backgroundColor: 'rgba(0,0,0,0.5)', }, modalContent: { backgroundColor: 'white', padding: 20, borderRadius: 10, alignItems: 'center', }, dismissButton: { marginTop: 15, padding: 10, backgroundColor: '#eee', borderRadius: 5, }, }); export default CameraScreen;
额外调试建议
- 通道顺序验证:若修正后仍预测错误,可能是相机输出为BGR通道,可尝试反转通道:
const preprocessedTensor = nextImageTensor .cast('float32') .div(tf.scalar(255)) .reverse(3) // 将BGR转为RGB .reshape([1, 224, 224, 3]); - 元数据核对:检查
metadata.json中的输入格式要求,确保预处理步骤与训练时完全一致。 - 性能调整:可降低
computeRecognitionEveryNFrames的值(如30)提升预测频率,同时关注设备性能负载。
内容的提问来源于stack exchange,提问作者deadpixels
相关产品推荐
相关产品推荐

