多线程环境下Flask后端model变量未更新问题求助
Flask多线程模型训练后API仍返回"Model not setup"问题排查
我正在开发一款Web应用,采用Flask后端处理预测模型,SvelteKit前端从后端获取预测数据。由于后端使用的数据会持续更新,已调度每周重训模型的任务。后端启动时,会在单独线程中完成模型的初始训练与后续重训,主线程负责处理API请求。初始时model变量设为None,预期完成训练后指向模型实例,但实际训练完成后,调用http://127.0.0.1:5000/predict接口时,前端仍收到“Error: Model not setup”响应,model变量仍为None。已尝试使用threading.Lock避免竞态条件,但问题依旧。
以下是简化后的Flask后端代码:
app.py
import threading from flask import Flask from flask_cors import CORS import models from routes import routes app = Flask(__name__) CORS(app) app.register_blueprint(routes) def setup_model(): models.setup() # set the flag for setting up the model setup_model_flag = True if setup_model_flag: training_thread = threading.Thread(target=setup_model) training_thread.start() if __name__ == '__main__': app.run(debug=False)
models.py
# setup import threading from datetime import datetime from time import sleep from schedule import every from schedule import repeat from schedule import run_pending model_lock = threading.Lock() model = None def prepare_data(): return "preparing data" class LSTMModel: def __init__(self): self.is_trained = False self.df = prepare_data() self.train() def train(self): print("Training model...") self.is_trained = True def predict(self, num_months): print("Predicting...") def setup(): with model_lock: print("Initializing...") global model model = LSTMModel() print(f"Model trained on {datetime.now()}") # schedule the job to run every sunday @repeat(every().sunday) def job(): with model_lock: # update model and retrain data print("Initializing new model...") global model # initialize new instance of model model = LSTMModel() print(f"Model trained on {datetime.now()}") while True: # print(idle_seconds()) run_pending() sleep(1)
routes.py
from flask import Blueprint from flask import jsonify from flask import request from models import model, model_lock routes = Blueprint('routes', __name__) @routes.route('/predict', methods=['POST']) def predict(): try: data = request.get_json() # Parse JSON data from the request body print('Received data:', data) input_data = data['value'] selected_option = data['type'] # Log a message to indicate that the endpoint is called print('Prediction endpoint called. Input:', input_data, 'Type:', selected_option) except Exception as e: # An error occurred, return an error response with 400 status code return jsonify({'error': 'Invalid JSON format'}), 400 if selected_option == 'Years': num_months = int(input_data) * 12 else: num_months = int(input_data) with model_lock: if model is None: # Model is not setup, return an error response with 500 status code return jsonify({'error': 'Model not setup'}), 500 if not model.is_trained: # Model is not trained, return an error response with 500 status code return jsonify({'error': 'Model not trained'}), 500 else: try: # Perform prediction using the model prediction = model.predict(num_months) except Exception as e: # An error occurred during prediction, return an error response with 500 status code return jsonify({'error': 'Prediction error: {}'.format(str(e))}), 500 # Return the prediction as a JSON response return jsonify({'prediction': prediction})
可通过Postman调用API,请求体JSON格式如下:
内容的提问来源于stack exchange,提问作者Linus Tan
相关产品推荐
相关产品推荐

