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

R中Keras训练的Sequential回归神经网络模型应选什么数据结构存储

R环境Keras序列回归模型存储方案

针对你通过keras_model_sequential接口训练得到的多个回归任务神经网络模型,根据使用场景分为内存临时存储、持久化落盘存储两类方案,适配模型的Python对象继承特性,不会出现属性丢失、指针失效问题。

内存存储(训练阶段临时调用、批量评估)

  • 优先使用**R原生命名列表(list)**存储,无额外依赖,支持自定义索引名匹配模型对应的变量组合,存取逻辑简单直接:
# 初始化存储容器
model_storage <- list()

# 每训练完一个模型,按标识命名存入,示例为使用3个特征训练的模型
model_storage[["3var_base"]] <- trained_model_1
# 存入第二个使用5个特征训练的模型
model_storage[["5var_tuned"]] <- trained_model_2

# 调用时直接按索引取
pred <- predict(model_storage[["3var_base"]], test_data)
  • 如果需要批量做训练、评估、指标对比,推荐用嵌套tibble数据框存储,可以把模型和对应的输入变量列表、训练集、验证集、评估指标存在同一张表中,搭配purrr包的映射函数可以直接做批量操作:
library(tibble)
library(purrr)

model_tbl <- tibble(
  model_id = c("m1", "m2"),
  used_features = list(c("x1","x2","x3"), c("x1","x2","x3","x4","x5")),
  train_set = list(train_m1, train_m2),
  model_obj = list(trained_model_1, trained_model_2),
  val_mse = c(0.15, 0.08)
)

# 批量做测试集预测
all_pred <- map2(model_tbl$model_obj, list(test_m1, test_m2), ~predict(.x, .y))

注意:不要用原子向量、普通非嵌套data.frame列存储这类Keras模型,你使用的模型本质是R封装的Python外部指针对象,继承层级如下,非列表类结构存储会破坏对象属性:

[1] "keras.engine.sequential.Sequential"
[2] "keras.engine.training.Model"
[3] "keras.engine.network.Network"
[4] "keras.engine.base_layer.Layer"
[5] "tensorflow.python.module.module.Module"
[6] "tensorflow.python.training.tracking.tracking.AutoTrackable"
[7] "tensorflow.python.training.tracking.base.Trackable"
[8] "python.builtin.object"

持久化存储(训练后长期保存、跨会话调用)

  • 不要用R原生的saveRDS()/readRDS()直接存储单个模型,这类方法只能保存R侧的对象封装,会丢失TensorFlow后端关联的计算图、权重参数,必须用Keras自带的save_model_tf()/load_model_tf()接口存储为标准SavedModel格式。
  • 批量存储多个模型时,单独创建模型存储根目录,每个模型对应一个唯一命名的子文件夹,同时保存一份索引表记录模型元信息,避免混淆:
# 创建存储目录
model_root <- "./regression_keras_models/"
dir.create(model_root, showWarnings = FALSE)

# 批量保存所有模型
walk2(model_tbl$model_id, model_tbl$model_obj, function(mid, mod){
  save_model_tf(mod, filepath = file.path(model_root, mid))
})

# 保存不含模型对象的元信息索引表
saveRDS(
  model_tbl %>% select(-model_obj), 
  file.path(model_root, "model_index.rds")
)

# 后续加载时先读索引,再按需加载模型
index <- readRDS(file.path(model_root, "model_index.rds"))
m1_loaded <- load_model_tf(file.path(model_root, "m1"))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 03:54:02