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

如何用TensorFlow.js在React Native中运行Python训练的ResNet50模型

将训练好的ResNet50模型通过TensorFlow.js集成到React Native实现图像相似推荐

需求说明

把Python中基于ResNet50训练完成的时尚图片推荐模型(对应下方Streamlit实现代码),迁移到React Native项目中,实现上传图片提取特征、返回5张相似图片的功能。

原Python端实现代码

import streamlit as st
import os
from PIL import Image
import numpy as np
import pickle
import tensorflow
from tensorflow.keras.layers import GlobalMaxPooling2D
from tensorflow.keras.applications.resnet50 import ResNet50,preprocess_input
from sklearn.neighbors import NearestNeighbors
from numpy.linalg import norm
import cv2

feature_list = np.array(pickle.load(open('featurevector.pkl','rb')))
filenames = pickle.load(open('filenames.pkl','rb'))

model = ResNet50(weights='imagenet',include_top=False,input_shape=(224,224,3))
model.trainable = False

model = tensorflow.keras.Sequential([
    model,
    GlobalMaxPooling2D()
])

st.title('Man & Women Fashion Recommender System')

def save_uploaded_file(uploaded_file):
    try:
        with open(os.path.join('uploads',uploaded_file.name),'wb') as f:
            f.write(uploaded_file.getbuffer())
        return 1
    except:
        return 0

def extract_feature(img_path, model):
    img=cv2.imread(img_path)
    img=cv2.resize(img, (224,224))
    img=np.array(img)
    expand_img=np.expand_dims(img, axis=0)
    pre_img=preprocess_input(expand_img)
    result=model.predict(pre_img).flatten()
    normalized=result/norm(result)
    return normalized

def recommend(features,feature_list):
    neighbors = NearestNeighbors(n_neighbors=6, algorithm='brute', metric='euclidean')
    neighbors.fit(feature_list)

    distances, indices = neighbors.kneighbors([features])

    return indices

# steps
# file upload -> save
uploaded_file = st.file_uploader("Choose an image")
print(uploaded_file)
if uploaded_file is not None:
    if save_uploaded_file(uploaded_file):
        # display the file
        display_image = Image.open(uploaded_file)
        resized_img = display_image.resize((200, 200))
        st.image(resized_img)
        # feature extract
        features = extract_feature(os.path.join("uploads",uploaded_file.name),model)
        #st.text(features)
        # recommendention
        indices = recommend(features,feature_list)
        # show
        col1,col2,col3,col4,col5 = st.columns(5)

        with col1:
            st.image(filenames[indices[0][1]])
        with col2:
            st.image(filenames[indices[0][2]])
        with col3:
            st.image(filenames[indices[0][3]])
        with col4:
            st.image(filenames[indices[0][4]])
        with col5:
            st.image(filenames[indices[0][5]])
    else:
        st.header("Some error occured in file upload")

迁移到React Native的实现步骤

1. 将Keras模型转换为TensorFlow.js格式

首先把Python中的模型导出为TensorFlow.js支持的格式:

  • 安装依赖:
pip install tensorflowjs
  • 执行导出脚本:
import tensorflow as tf
from tensorflow.keras.applications.resnet50 import ResNet50
from tensorflow.keras.layers import GlobalMaxPooling2D

# 加载与原代码一致的模型结构
base_model = ResNet50(weights='imagenet', include_top=False, input_shape=(224,224,3))
base_model.trainable = False
model = tf.keras.Sequential([
    base_model,
    GlobalMaxPooling2D()
])

# 保存为TensorFlow SavedModel格式
tf.saved_model.save(model, './saved_model')
# 转换为TensorFlow.js格式
!tensorflowjs_converter --input_format=tf_saved_model ./saved_model ./tfjs_model

导出后会得到model.json和若干权重分片文件(如group1-shard1ofX.bin),将这些文件放到React Native项目的assets/tfjs_model目录下。

2. React Native项目配置TensorFlow.js

  • 安装所需依赖:
npm install @tensorflow/tfjs-react-native @react-native-community/image-picker react-native-fs
  • 在项目中初始化模型加载:
import * as tf from '@tensorflow/tfjs';
import { bundleResourceIO } from '@tensorflow/tfjs-react-native';

// 导入模型资源
const modelJson = require('./assets/tfjs_model/model.json');
// 注意:替换为实际生成的所有权重文件
const modelWeights = [
  require('./assets/tfjs_model/group1-shard1of2.bin'),
  require('./assets/tfjs_model/group1-shard2of2.bin')
];

// 加载模型的异步函数
async function loadModel() {
  await tf.ready(); // 等待TensorFlow.js环境初始化
  const model = await tf.loadLayersModel(bundleResourceIO(modelJson, modelWeights));
  return model;
}

3. 图像预处理与特征提取

实现和Python端extract_feature函数一致的逻辑,处理React Native获取的图片:

import { Image } from 'react-native';
import * as tf from '@tensorflow/tfjs';

// 预处理图片,匹配ResNet50的输入要求
async function preprocessImage(uri) {
  // 获取图片原始尺寸
  const { width, height } = await new Promise((resolve) => {
    Image.getSize(uri, (w, h) => resolve({ width: w, height: h }));
  });

  // 将图片转换为张量并调整大小
  const imgTensor = await tf.browser.fromPixelsAsync(uri);
  const resizedTensor = tf.image.resizeBilinear(imgTensor, [224, 224]);
  
  // 添加batch维度,并执行ResNet50的预处理(像素值归一化到-1~1,BGR转RGB适配cv2的输入)
  const expandedTensor = resizedTensor.expandDims(0);
  const normalizedTensor = expandedTensor.toFloat().div(tf.scalar(127.5)).sub(tf.scalar(1.0));
  // 因为Python用cv2.imread读的是BGR,这里将RGB转成BGR
  const bgrTensor = normalizedTensor.reverse(-1);

  // 释放无用张量避免内存泄漏
  imgTensor.dispose();
  resizedTensor.dispose();
  expandedTensor.dispose();
  
  return bgrTensor;
}

// 提取图片特征向量并归一化
async function extractFeatures(uri, model) {
  const processedImg = await preprocessImage(uri);
  const rawFeatures = model.predict(processedImg).flatten();
  // L2归一化
  const normalizedFeatures = rawFeatures.div(tf.norm(rawFeatures));

  processedImg.dispose();
  rawFeatures.dispose();
  
  return normalizedFeatures.array();
}

4. 实现相似图片推荐逻辑

由于React Native中无法直接使用sklearn的NearestNeighbors,手动实现欧氏距离计算与排序:

  • 先将Python中的featurevector.pkl和filenames.pkl转换为JSON格式:
import pickle
import json

# 转换特征向量
feature_list = np.array(pickle.load(open('featurevector.pkl','rb')))
with open('feature_list.json', 'w') as f:
    json.dump(feature_list.tolist(), f)

# 转换文件名列表
filenames = pickle.load(open('filenames.pkl','rb'))
with open('filenames.json', 'w') as f:
    json.dump(filenames, f)

将生成的两个JSON文件放到React Native项目的assets目录下,然后实现推荐函数:

const featureList = require('./assets/feature_list.json');
const filenames = require('./assets/filenames.json');

// 根据输入特征计算并返回Top5相似图片路径
async function getSimilarImages(inputFeatures) {
  // 计算每个特征向量与输入特征的欧氏距离
  const distanceList = featureList.map((feat) => {
    let distanceSum = 0;
    for (let i = 0; i < feat.length; i++) {
      distanceSum += Math.pow(feat[i] - inputFeatures[i], 2);
    }
    return Math.sqrt(distanceSum);
  });

  // 将距离与索引绑定,按距离升序排序,取前5个(排除自身)
  const indexedDistances = distanceList.map((dist, idx) => ({ dist, idx }));
  indexedDistances.sort((a, b) => a.dist - b.dist);

  const top5Indices = indexedDistances.slice(1, 6).map(item => item.idx);
  return top5Indices.map(idx => filenames[idx]);
}

5. 整合图片选择与推荐流程

使用图片选择器获取用户上传的图片,串联整个推荐流程:

import ImagePicker from '@react-native-community/image-picker';

// 选择图片并触发推荐
async function pickAndRecommendImage() {
  const pickerResult = await ImagePicker.launchImageLibraryAsync({
    mediaType: 'photo',
    allowsEditing: true,
    aspect: [1, 1],
    quality: 0.8,
  });

  if (!pickerResult.cancelled) {
    try {
      const model = await loadModel();
      const features = await extractFeatures(pickerResult.uri, model);
      const similarImages = await getSimilarImages(features);
      // 这里可以将similarImages渲染到页面上展示
      console.log('相似图片列表:', similarImages);
    } catch (err) {
      console.error('推荐流程出错:', err);
    }
  }
}

注意事项

  • 若权重文件较多,需在metro.config.js中配置允许加载.bin文件:
module.exports = {
  resolver: {
    assetExts: ['bin', 'json', 'png', 'jpg'],
  },
};
  • 频繁处理图片时要注意及时释放TensorFlow.js张量,避免内存泄漏
  • 若图片路径涉及本地文件,需确保React Native有文件访问权限

内容的提问来源于stack exchange,提问作者Ömer Faruk Baltacı

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 18:27:11