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

如何在指定时间后终止Scikit-learn的fit()函数执行?

需求实现方案:Scikit-learn RFECV fit()超时终止机制

核心实现逻辑

使用multiprocessing模块创建独立子进程执行RFECV的fit操作,主进程监控子进程运行时长,超时则强制终止子进程,直接切换到下一个回归器。该方案能可靠解决单进程下无法中途终止阻塞fit()的问题。

完整代码实现

import time
import multiprocessing
from sklearn.datasets import make_regression
from sklearn.feature_selection import RFECV
from sklearn.model_selection import train_test_split
from sklearn.utils import all_estimators
from sklearn.exceptions import ConvergenceWarning
import warnings
warnings.filterwarnings("ignore", category=ConvergenceWarning)

# ------------------------------------------------------------------------------------------------ #
#                                      MAKE REGRESSION DATASET                                     #
# ------------------------------------------------------------------------------------------------ #
X, y = make_regression(n_samples=3000,
                       n_features=250,
                       n_informative=50,
                       n_targets=1,
                       shuffle=True,
                       noise=0.1,
                       coef=False)

# ------------------------------------------------------------------------------------------------ #
#                                         TRAIN TEST SPLIT                                         #
# ------------------------------------------------------------------------------------------------ #
X_train, X_test, y_train, y_test = train_test_split(X,
                                                    y,
                                                    shuffle=False,
                                                    test_size=0.25,
                                                    random_state=0)

# ------------------------------------------------------------------------------------------------ #
#                                    GET ALL SKLEARN REGRESSORS                                    #
# ------------------------------------------------------------------------------------------------ #
all_sklearn_regressors_unfiltered = all_estimators(type_filter='regressor')
all_sklearn_regressors_filtered = []
for name, RegressorClass in all_sklearn_regressors_unfiltered:
    try:
        reg = RegressorClass()
        print("Adding regressor:", name)
        all_sklearn_regressors_filtered.append((name, reg))  # 保存名称和实例
    except Exception as e:
        print("ERROR:", name, ":", e,", not adding it.")

# ------------------------------------------------------------------------------------------------ #
#                          包装RFECV fit操作,用于子进程执行                          #
# ------------------------------------------------------------------------------------------------ #
def fit_rfecv(regressor, X_train, y_train, result_queue):
    try:
        selector = RFECV(regressor, 
                         step=10, 
                         cv=5, 
                         verbose=0, 
                         n_jobs=1)
        selector.fit(X_train, y_train)
        result_queue.put((True, selector))
    except Exception as e:
        result_queue.put((False, str(e)))

# ------------------------------------------------------------------------------------------------ #
#                                带超时机制的回归器迭代逻辑                                #
# ------------------------------------------------------------------------------------------------ #
TIMEOUT_SECONDS = 120  # 超时阈值

for reg_name, reg_instance in all_sklearn_regressors_filtered:
    start_time = time.time()
    result_queue = multiprocessing.Queue()
    
    # 创建子进程执行fit操作
    p = multiprocessing.Process(target=fit_rfecv, args=(reg_instance, X_train, y_train, result_queue))
    p.start()
    
    # 等待子进程完成,超时则终止
    p.join(TIMEOUT_SECONDS)
    
    if p.is_alive():
        # 超时,终止子进程
        p.terminate()
        p.join()  # 等待进程彻底退出
        elapsed_time = time.time() - start_time
        print(f"[超时终止] 回归器 {reg_name} 运行时长 {elapsed_time:.2f} 秒,已终止")
    else:
        # 正常完成或异常退出
        success, result = result_queue.get()
        elapsed_time = time.time() - start_time
        if success:
            print(f"[正常完成] 回归器 {reg_name} 运行时长 {elapsed_time:.2f} 秒")
            # 这里可以处理result(即训练好的selector)
        else:
            print(f"[执行异常] 回归器 {reg_name} 运行时长 {elapsed_time:.2f} 秒,错误信息: {result}")

关键说明

  • 每个回归器的RFECV fit操作都在独立子进程中执行,主进程通过join(timeout)监控时长
  • 超时后调用terminate()强制终止子进程,再调用join()确保资源回收
  • 使用Queue传递子进程的执行结果(成功/失败、训练好的模型或错误信息)
  • 保存回归器名称便于更清晰的日志输出

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.23 12:36:32