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

如何在React.js/Flask应用中集成Flower框架实现联邦学习图像分类

基于React+Flask+Celery(Flower)+TensorFlow的图像分类Web应用实现方案

一、核心架构梳理

明确各组件职责,避免耦合:

  • React前端:图像上传UI、分类结果展示、异步任务状态查询
  • Flask后端:提供API接口(上传触发、任务查询)、全局加载TensorFlow模型
  • Celery+Redis/RabbitMQ:异步执行图像分类任务,避免阻塞Flask主线程
  • Flower:监控Celery任务队列、查看执行状态与报错日志

二、关键模块实现细节

1. Flask后端与TensorFlow集成

  • 模型全局加载:避免每次请求重复加载模型,提升响应效率

    from flask import Flask, request, jsonify
    from celery import Celery
    import tensorflow as tf
    import numpy as np
    from PIL import Image
    import io
    
    app = Flask(__name__)
    # 配置Celery消息队列与结果存储
    app.config['CELERY_BROKER_URL'] = 'redis://localhost:6379/0'
    app.config['CELERY_RESULT_BACKEND'] = 'redis://localhost:6379/0'
    celery = Celery(app.name, broker=app.config['CELERY_BROKER_URL'])
    celery.conf.update(app.config)
    
    # 启动时加载TensorFlow模型(仅加载一次)
    MODEL_PATH = './trained_image_classifier.h5'
    model = tf.keras.models.load_model(MODEL_PATH)
    # 匹配你的模型分类标签
    CLASS_LABELS = ['猫', '狗', '鸟', '其他']
    
    # 图像预处理函数(严格对齐模型训练时的逻辑)
    def preprocess_image(image_bytes):
        img = Image.open(io.BytesIO(image_bytes)).resize((224, 224))  # 匹配模型输入尺寸
        img_array = np.array(img) / 255.0  # 归一化
        return np.expand_dims(img_array, axis=0)
    
  • 异步分类任务定义:

    @celery.task(bind=True)
    def classify_image_task(self, image_bytes):
        try:
            processed_img = preprocess_image(image_bytes)
            predictions = model.predict(processed_img)
            top_idx = np.argmax(predictions[0])
            confidence = round(float(predictions[0][top_idx]) * 100, 2)
            return {'class': CLASS_LABELS[top_idx], 'confidence': confidence}
        except Exception as e:
            self.update_state(state='FAILURE', meta={'error': str(e)})
            raise
    
  • Flask API接口:

    # 上传图像并触发异步任务
    @app.route('/api/classify', methods=['POST'])
    def trigger_classify():
        if 'image' not in request.files:
            return jsonify({'error': '未上传图像文件'}), 400
        image_file = request.files['image']
        task = classify_image_task.delay(image_file.read())
        return jsonify({'task_id': task.id}), 202
    
    # 查询任务状态与结果
    @app.route('/api/task/<task_id>', methods=['GET'])
    def get_task_result(task_id):
        task = classify_image_task.AsyncResult(task_id)
        if task.state == 'PENDING':
            return jsonify({'state': 'PENDING', 'msg': '任务等待执行'})
        elif task.state == 'SUCCESS':
            return jsonify({'state': 'SUCCESS', 'result': task.result})
        elif task.state == 'FAILURE':
            return jsonify({'state': 'FAILURE', 'error': task.info.get('error', '未知错误')})
        else:
            return jsonify({'state': task.state, 'msg': '任务执行中'})
    

2. React前端实现

  • 核心功能:图像上传、任务状态轮询、结果展示
    import { useState } from 'react';
    
    function ImageClassifier() {
        const [selectedImg, setSelectedImg] = useState(null);
        const [taskId, setTaskId] = useState(null);
        const [status, setStatus] = useState('');
        const [result, setResult] = useState(null);
    
        const handleUpload = async (e) => {
            const file = e.target.files[0];
            if (!file) return;
            setSelectedImg(file);
            const formData = new FormData();
            formData.append('image', file);
    
            try {
                const res = await fetch('/api/classify', { method: 'POST', body: formData });
                const data = await res.json();
                setTaskId(data.task_id);
                setStatus('正在分类,请稍候...');
                // 轮询任务状态
                const interval = setInterval(async () => {
                    const taskRes = await fetch(`/api/task/${data.task_id}`);
                    const taskData = await taskRes.json();
                    if (taskData.state === 'SUCCESS') {
                        setResult(taskData.result);
                        setStatus('分类完成');
                        clearInterval(interval);
                    } else if (taskData.state === 'FAILURE') {
                        setStatus(`分类失败:${taskData.error}`);
                        clearInterval(interval);
                    }
                }, 1000);
            } catch (err) {
                setStatus(`请求失败:${err.message}`);
            }
        };
    
        return (
            <div className="classifier-container">
                <input type="file" accept="image/*" onChange={handleUpload} />
                {selectedImg && <img src={URL.createObjectURL(selectedImg)} alt="上传预览" style={{ width: '200px', marginTop: '10px' }} />}
                <p className="status-text">{status}</p>
                {result && (
                    <div className="result-card">
                        <h3>分类结果</h3>
                        <p>类别:{result.class}</p>
                        <p>置信度:{result.confidence}%</p>
                    </div>
                )}
            </div>
        );
    }
    
    export default ImageClassifier;
    

3. Flower监控配置

  • 启动Flower服务(需确保Celery worker已启动):
    celery -A app.celery flower --port=5555
    
  • 访问http://localhost:5555即可查看任务队列状态、执行日志、失败任务详情,快速定位问题

三、常见瓶颈解决方案

  1. 模型加载慢/内存占用高:
    • 将TensorFlow模型转换为TensorFlow Lite格式,压缩体积并加快加载速度
    • 若部署到服务器,可使用GPU加速(需安装对应TensorFlow GPU版本)
  2. 异步任务阻塞:
    • 启动Celery worker时指定并发数:celery -A app.celery worker --loglevel=info --concurrency=4
    • 避免在任务中执行IO密集型操作,可拆分独立任务处理
  3. 跨域问题:
    • 在Flask中添加跨域支持:
      from flask_cors import CORS
      CORS(app)
      
  4. 图像预处理错误:
    • 严格对齐模型训练时的预处理逻辑(尺寸、归一化、颜色通道顺序)
    • 前端添加图像格式校验,仅允许JPG/PNG等支持的格式

四、调试排障技巧

  • 通过Flower查看任务报错日志,定位代码执行异常点
  • 在Flask后端添加日志记录:app.logger.info(f"处理任务:{task_id}")
  • 单独测试TensorFlow模型的推理功能,排除模型本身问题
  • 用Postman测试API接口,验证后端逻辑正确性

内容的提问来源于stack exchange,提问作者chaima bouabd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 19:06:18