如何在指定时间后终止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
相关产品推荐
相关产品推荐

