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

Flask应用与请求上下文间sklearn transform_output配置差异问题

问题背景

正在修改一个部署机器学习模型的Flask应用,以支持新模型的预测服务。该模型是使用Pickle序列化的sklearn Pipeline对象,包含Column Transformer、缺失值填充、编码及预测等步骤。由于自定义Column Transformer的实现逻辑,中间步骤必须输出Pandas DataFrame而非数组,这也是模型训练和序列化时的配置。

异常情况

  • 应用上下文加载并预测时,Pipeline可正常运行并返回结果;
  • 请求上下文使用同一加载后的模型和相同数据时,中间Pipeline步骤未输出Pandas格式,导致后续步骤因找不到数组中不存在的列名而失败。

疑问

  1. 为何两种上下文下的行为存在差异?
  2. 如何无需重新配置Pipeline为数组输出并重新训练即可解决该问题?
  3. 是否可仅在请求上下文中配置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

已尝试操作

  1. 在两种上下文下运行模型:
    • 请求上下文内,跟踪发现首个ColumnTransformer输出为数组,而非所需的Pandas DataFrame;
    • 应用上下文内,所有ColumnTransformer及中间步骤输出均为Pandas格式,符合模型配置。
  2. 在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 10:44:52