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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 04:08:18