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

如何在Keras中获取模型预测值?(附R语言代码示例)

如何在R版Keras的5折交叉验证中获取并保存预测值?

我正在学习R语言版本的Keras,想要查看模型输出的预测数值,但当前代码里没有保存预测值的逻辑。以下是我的数据集处理、模型构建及5折交叉验证代码:

数据集处理代码

df <- MASS::Boston
index <- sample(c(TRUE, FALSE), nrow(df), replace=TRUE, prob=c(0.7,0.3))
train_features <- Boston[index,]
test_features <- Boston[!index,]
train_labels <- Boston$medv[index]
test_labels <- Boston$medv[!index]
train_features <- scale(train_features)
train_features <- train_features[,1:ncol(train_features)]
test_features <- scale(test_features)
test_features <- test_features[,1:ncol(test_features)]
mean <- apply(train_features, 2, mean)
sd <- apply(train_features, 2, sd)
train_data <- scale(train_features, center = mean, scale = sd)
test_data <- scale(test_features, center = mean, scale = sd)
train_targets <- Boston$medv[index]
test_targets <- Boston$medv[!index]

模型构建代码

build_model <- function() {
  
  model <- keras_model_sequential() %>%
    layer_dense(64, activation = "relu") %>%
    layer_dense(64, activation = "relu") %>%
    layer_dense(1)
  
  model %>% compile(optimizer = "rmsprop",
                    loss = "mse",
                    metrics = "mse")
  model
}

5折交叉验证代码

k <- 5
fold_id <- sample(rep(1:k, length.out = nrow(train_data)))
num_epochs <- 100
all_scores <- numeric()

for (i in 1:k) {
  cat("Processing fold #", i, "\n")
  val_indices <- which(fold_id == i)
  
  val_data <- train_data[val_indices, ]
  val_targets <- train_targets[val_indices]
  
  partial_train_data <- train_data[-val_indices, ]
  partial_train_targets <- train_targets[-val_indices]
  
  model <- build_model()
  
  model %>% fit (
    partial_train_data,
    partial_train_targets,
    epochs = num_epochs,
    batch_size = 16,
    verbose = 0
  )
  
  results <- model %>%
    evaluate(val_data, val_targets, verbose = 0)
  all_scores[[i]] <- results[['mse']]
}

keras.RMSE <- sqrt(mean(all_scores))

目前遇到的问题:

  • all_scores仅保存了RMSE分数,没有预测值
  • val_targets和预测值维度可能不匹配
  • model$fit不返回预测值,知道用model$predict但不知道在哪里保存

解决方案

要获取并保存预测值,你需要在每折交叉验证的模型训练完成后,调用predict()生成验证集的预测结果,同时保存对应的真实标签。具体修改如下:

  1. 初始化存储预测结果的结构:创建一个列表来保存每折的真实值和预测值,方便后续分析。
  2. 生成并保存预测值:在evaluate()之后,调用predict()生成验证集的预测值,将真实值和预测值存入列表。
  3. 修正维度匹配问题:predict()返回的是矩阵格式,可通过as.vector()转换为向量,和val_targets保持一致。

修改后的5折交叉验证代码:

k <- 5
fold_id <- sample(rep(1:k, length.out = nrow(train_data)))
num_epochs <- 100
all_scores <- numeric()
# 初始化列表存储每折的真实值和预测值
all_predictions <- list()

for (i in 1:k) {
  cat("Processing fold #", i, "\n")
  val_indices <- which(fold_id == i)
  
  val_data <- train_data[val_indices, ]
  val_targets <- train_targets[val_indices]
  
  partial_train_data <- train_data[-val_indices, ]
  partial_train_targets <- train_targets[-val_indices]
  
  model <- build_model()
  
  model %>% fit (
    partial_train_data,
    partial_train_targets,
    epochs = num_epochs,
    batch_size = 16,
    verbose = 0
  )
  
  # 生成验证集预测值并转换为向量
  val_predictions <- as.vector(model %>% predict(val_data, verbose = 0))
  
  results <- model %>%
    evaluate(val_data, val_targets, verbose = 0)
  all_scores[[i]] <- results[['mse']]
  
  # 保存当前折的真实值和预测值
  all_predictions[[i]] <- data.frame(
    true_value = val_targets,
    predicted_value = val_predictions,
    fold = i
  )
}

keras.RMSE <- sqrt(mean(all_scores))

# 将所有折的结果合并为一个数据框,方便整体分析
all_predictions_df <- do.call(rbind, all_predictions)

额外说明

  • 数据处理部分存在冗余:train_features已经通过scale()标准化,后续又用mean和sd再次标准化,这会导致重复处理。可以简化为:
# 仅对训练集计算均值和标准差
train_features <- Boston[index,]
test_features <- Boston[!index,]
mean <- apply(train_features, 2, mean)
sd <- apply(train_features, 2, sd)
# 用训练集的均值和标准差标准化训练集和测试集
train_data <- scale(train_features, center = mean, scale = sd)
test_data <- scale(test_features, center = mean, scale = sd)
train_targets <- Boston$medv[index]
test_targets <- Boston$medv[!index]
  • 如果需要对测试集生成预测值,在交叉验证完成后,重新训练一个完整的模型(用全部训练数据),然后调用predict(test_data)即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 13:05:37