如何用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ı
相关产品推荐
相关产品推荐

