You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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格式,通道顺序也可能不匹配。
  • 张量资源管理混乱:循环中未正确释放所有临时张量,易引发内存泄漏,干扰预测逻辑。

修正步骤

  1. 提前加载模型:在组件初始化时仅加载一次模型,避免重复加载导致的状态异常。
  2. 添加标准化图像预处理:将相机输入张量转换为模型要求的格式(归一化、通道匹配)。
  3. 规范张量资源释放:使用tf.tidy和tf.dispose自动清理临时张量,避免内存泄漏。
  4. 对齐输入格式:确保输入张量的形状、数值范围与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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.09 16:15:55