Flask应用与请求上下文间sklearn transform_output配置差异问题
问题背景
正在修改一个部署机器学习模型的Flask应用,以支持新模型的预测服务。该模型是使用Pickle序列化的sklearn Pipeline对象,包含Column Transformer、缺失值填充、编码及预测等步骤。由于自定义Column Transformer的实现逻辑,中间步骤必须输出Pandas DataFrame而非数组,这也是模型训练和序列化时的配置。
异常情况
- 应用上下文加载并预测时,Pipeline可正常运行并返回结果;
- 请求上下文使用同一加载后的模型和相同数据时,中间Pipeline步骤未输出Pandas格式,导致后续步骤因找不到数组中不存在的列名而失败。
疑问
- 为何两种上下文下的行为存在差异?
- 如何无需重新配置Pipeline为数组输出并重新训练即可解决该问题?
- 是否可仅在请求上下文中配置
sklearn.set_config(transform_output="pandas")?
基础代码
import os import pickle import pandas as pd from flask import (Flask, redirect, make_response) import logging #define app app = Flask(__name__) # load the trained Model model = pickle.load(open("model.pkl"), "rb") # load test record test_record = pd.read_csv("test_record.csv") # Make a prediction using the model and test record. This step works. try: model.predict(test_record) except Exception as e: log.debug(str(e), stack_info=True) @app.route('/predict') def predict(): # Make a prediction using the model and test record. This step *doesn't* work. try: prediction = model.predict(test_record) except Exception as e: log.debug(str(e), stack_info=True) return make_response('Test Record Prediction: ' + str(prediction),200) # Start the Flask app if __name__ == '__main__': if os.environ['ENV'] in {'local','local_w_db','DEV'}: app.run(debug=True) else: app.run()
环境配置
python==3.11.7 Flask==3.0.1 scikit-learn==1.3.2 scikit-learn-intelex==2023.2.1 scipy==1.11.4 pandas==2.1.4 category_encoders==2.6.3 werkzeug==3.0.1 joblib==1.2.0
已尝试操作
- 在两种上下文下运行模型:
- 请求上下文内,跟踪发现首个ColumnTransformer输出为数组,而非所需的Pandas DataFrame;
- 应用上下文内,所有ColumnTransformer及中间步骤输出均为Pandas格式,符合模型配置。
- 在Flask代码开头设置全局sklearn配置
sklearn.set_config(transform_output="pandas"),但无效果。
异常堆栈信息
[2024-01-30 08:58:44,666] DEBUG [app.predict:109] - Specifying the columns using strings is only supported for pandas DataFrames Stack (most recent call last): File "c:\Users\user\.vscode\extensions\ms-python.python-2023.22.1\pythonFiles\lib\python\debugpy\_vendored\pydevd\_pydev_bundle\pydev_monkey.py", line 1118, in __call__ ret = self.original_func(*self.args, **self.kwargs) File "..\miniforge3\envs\APP_ENV\Lib\threading.py", line 1002, in _bootstrap self._bootstrap_inner() File "..\miniforge3\envs\APP_ENV\Lib\threading.py", line 1045, in _bootstrap_inner self.run() File "..\miniforge3\envs\APP_ENV\Lib\threading.py", line 982, in run self._target(*self._args, **self._kwargs) File "..\miniforge3\envs\APP_ENV\Lib\socketserver.py", line 691, in process_request_thread self.finish_request(request, client_address) File "..\miniforge3\envs\APP_ENV\Lib\socketserver.py", line 361, in finish_request self.RequestHandlerClass(request, client_address, self) File "..\miniforge3\envs\APP_ENV\Lib\socketserver.py", line 755, in __init__ self.handle() File "..\miniforge3\envs\APP_ENV\Lib\site-packages\werkzeug\serving.py", line 390, in handle super().handle() File "..\miniforge3\envs\APP_ENV\Lib\http\server.py", line 436, in handle self.handle_one_request() File "..\miniforge3\envs\APP_ENV\Lib\http\server.py", line 424, in handle_one_request method() File "..\miniforge3\envs\APP_ENV\Lib\site-packages\werkzeug\serving.py", line 362, in run_wsgi execute(self.server.app) File "..\miniforge3\envs\APP_ENV\Lib\site-packages\werkzeug\serving.py", line 323, in execute application_iter = app(environ, start_response) File "..\miniforge3\envs\APP_ENV\Lib\site-packages\flask\app.py", line 1488, in __call__ return self.wsgi_app(environ, start_response) File "..\miniforge3\envs\APP_ENV\Lib\site-packages\flask\app.py", line 1463, in wsgi_app response = self.full_dispatch_request() File "..\miniforge3\envs\APP_ENV\Lib\site-packages\flask\app.py", line 870, in full_dispatch_request rv = self.dispatch_request() File "..\miniforge3\envs\APP_ENV\Lib\site-packages\flask\app.py", line 855, in dispatch_request return self.ensure_sync(self.view_functions[rule.endpoint])(**view_args) # type: ignore[no-any-return] File "C:\Users\user\app_directory\app.py", line 109, in predict log.debug(str(e), stack_info=True)
解决方案与解释
1. 上下文行为差异的原因
Flask在debug模式下默认启用重载器,会重新加载应用代码。全局作用域加载的模型和sklearn配置会在主进程执行一次,而处理请求的子进程会重新执行整个脚本,导致原有的sklearn配置状态丢失,模型加载时的上下文配置被重置为默认值(输出数组),从而出现两种上下文行为不一致的情况。
2. 无需重新训练的解决方法
方法一:禁用Flask重载器
启动Flask时关闭debug模式的重载功能,确保应用在单个进程中运行,避免配置重复执行丢失:
if __name__ == '__main__': if os.environ['ENV'] in {'local','local_w_db','DEV'}: app.run(debug=True, use_reloader=False) # 关闭重载器 else: app.run()
方法二:强制设置模型组件的输出配置
加载模型后,遍历Pipeline的所有步骤,直接给每个支持的组件设置transform_output参数,不受全局配置或进程重载影响:
import sklearn # load the trained Model model = pickle.load(open("model.pkl"), "rb") # 递归设置所有组件的transform_output为pandas def set_transform_output(obj): if hasattr(obj, 'transform_output'): obj.transform_output = "pandas" # 处理ColumnTransformer内部的子转换器 if hasattr(obj, 'transformers'): for name, trans, cols in obj.transformers: set_transform_output(trans) # 处理Pipeline的每一步 if hasattr(obj, 'steps'): for name, step in obj.steps: set_transform_output(step) set_transform_output(model)
方法三:在请求上下文内配置并刷新组件
在predict函数开头设置sklearn全局配置,并强制刷新模型中ColumnTransformer的配置:
@app.route('/predict') def predict(): import sklearn sklearn.set_config(transform_output="pandas") # 强制更新ColumnTransformer的输出配置 if hasattr(model, 'steps'): for name, step in model.steps: if isinstance(step, sklearn.compose.ColumnTransformer): step.transform_output = "pandas" try: prediction = model.predict(test_record) except Exception as e: log.debug(str(e), stack_info=True) return make_response('Test Record Prediction: ' + str(prediction),200)
3. 仅在请求上下文配置的可行性
可以,但需注意进程隔离问题:如果Flask采用多进程部署(比如debug重载器、生产环境多进程),每个请求进程都需要单独设置该配置。直接在predict函数内添加sklearn.set_config(transform_output="pandas")即可,但结合方法二直接修改模型组件的本地配置,会比依赖全局状态更可靠。
内容的提问来源于stack exchange,提问作者coruble
相关产品推荐
相关产品推荐

