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
相关产品推荐
相关产品推荐

