在R中借助foreach多线程实现GAM模型训练的暂停与恢复
解决方案:GAM训练断点续训与集群批量处理适配
一、单个GAM模型的断点续训实现
mgcv的gam()本身无内置断点续训功能,可通过以下两种方式实现:
1. 手动拆分迭代过程实现Checkpoint
mgcv的GAM拟合依赖外点迭代,可手动调用gam.fit()分步执行,每完成若干迭代就保存中间状态:
# 初始化未拟合的GAM框架 init_gam <- mgcv::gam(r ~ te(x1) + te(x2) + te(x1,x2), data = mydata, fit = FALSE) # 加载已保存的中间状态(若存在) checkpoint_path <- "gam_checkpoint.RData" if (file.exists(checkpoint_path)) { load(checkpoint_path) } else { current_fit <- init_gam iter_count <- 0 } # 设置迭代参数与保存间隔 max_iter <- 100 checkpoint_interval <- 5 while (iter_count < max_iter && !current_fit$converged) { # 执行单次迭代 current_fit <- mgcv::gam.fit(current_fit) iter_count <- iter_count + 1 # 定期保存状态 if (iter_count %% checkpoint_interval == 0) { save(current_fit, iter_count, file = checkpoint_path) } # 监控集群剩余运行时间(以SLURM环境为例) remaining_time <- as.numeric(Sys.getenv("SLURM_TIME_REMAINING")) # 剩余1小时时停止训练,保存最终中间状态 if (!is.na(remaining_time) && remaining_time < 3600) { save(current_fit, iter_count, file = paste0(checkpoint_path, "_final")) break } } # 拟合完成后整理为标准GAM对象 final_gam <- mgcv::gam.object(current_fit)
2. 进程级信号触发保存
通过捕获集群终止信号,在作业超时前触发状态保存:
# 定义信号处理函数:收到SIGUSR1信号时保存当前模型 save_on_signal <- function(sig) { save(mygam, file = "gam_intermediate.RData") quit(save = "no") } signal(SIGUSR1, save_on_signal) # 正常执行GAM训练 mygam <- mgcv::gam(r ~ te(x1) + te(x2) + te(x1,x2), data = mydata)
搭配集群作业脚本,在超时前1小时发送SIGUSR1信号即可触发保存(例如SLURM可通过scancel --signal=USR1 <jobid>实现)。
二、结合foreach的批量模型断点续训
针对上百个模型的批量处理,需记录已完成任务,避免重复训练:
# 生成所有待抽样的参数组合 param_grid <- expand.grid(a = a_values, b = b_values) # 加载已完成任务记录 progress_path <- "completed_models.RData" if (file.exists(progress_path)) { load(progress_path) } else { completed_indices <- c() model_results <- list() } # 遍历未完成的参数组合 foreach(i = seq_len(nrow(param_grid)), .packages = "mgcv") %dopar% { if (i %in% completed_indices) next # 替换为你的数据抽样逻辑 mydata <- sample_data(param_grid$a[i], param_grid$b[i]) # 单个模型的checkpoint路径 model_checkpoint <- paste0("gam_checkpoint_", i, ".RData") if (file.exists(model_checkpoint)) { # 从断点恢复训练 load(model_checkpoint) while (!current_fit$converged) { current_fit <- mgcv::gam.fit(current_fit) save(current_fit, iter_count, file = model_checkpoint) remaining_time <- as.numeric(Sys.getenv("SLURM_TIME_REMAINING")) if (!is.na(remaining_time) && remaining_time < 3600) quit(save = "no") } final_gam <- mgcv::gam.object(current_fit) file.remove(model_checkpoint) } else { # 从头开始训练 init_gam <- mgcv::gam(r ~ te(x1) + te(x2) + te(x1,x2), data = mydata, fit = FALSE) current_fit <- init_gam iter_count <- 0 while (!current_fit$converged) { current_fit <- mgcv::gam.fit(current_fit) iter_count <- iter_count + 1 save(current_fit, iter_count, file = model_checkpoint) remaining_time <- as.numeric(Sys.getenv("SLURM_TIME_REMAINING")) if (!is.na(remaining_time) && remaining_time < 3600) quit(save = "no") } final_gam <- mgcv::gam.object(current_fit) file.remove(model_checkpoint) } # 更新任务进度 model_results[[i]] <- final_gam completed_indices <- c(completed_indices, i) save(model_results, completed_indices, file = progress_path) }
三、集群作业脚本适配
以SLURM为例,编写自动重提交的作业脚本,实现超时后自动续训:
#!/bin/bash #SBATCH --time=24:00:00 # 集群单任务最大时长 #SBATCH --cpus-per-task=4 #SBATCH --mem=16G # 加载R环境 module load R/4.3.1 # 运行训练脚本 Rscript gam_training.R # 若脚本因时间限制退出,自动重新提交作业 if [ $? -ne 0 ]; then sbatch $0 fi
内容的提问来源于stack exchange,提问作者Maki
相关产品推荐
相关产品推荐

