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

R语言批量梯度下降函数迭代存储与可视化问题求助

Hey there, let's fix up your batch gradient descent function and add all the features you're looking for—tracking iteration history, calculating test error, and visualizing cost convergence. Let's break down the issues in your code first, then build the corrected version step by step.

First, Fix the Core Errors in Your Code

Your main issues came from a few small but critical mistakes:

  • Case mismatch in error.cost: You used uppercase X instead of lowercase x (the parameter passed to the function), which caused R to look for an undefined variable.
  • Global variable pitfalls: Defining cost_history and theta_history outside the function means they'll retain values from previous runs and can cause dimension mismatches if iterations don't hit num_iters.
  • Incorrect theta initialization: matrix(c(1,1), ncol(x), 1) only works if your input x has exactly 2 features. Use matrix(1, ncol(x), 1) to initialize a theta vector of 1s matching the number of features (including the intercept).
  • Inefficient cost storage: Using append() repeatedly is slow; pre-allocating vectors for history is better.

Refactored Batch Gradient Descent Function

Here's a revised version that fixes these issues, stores iteration history, calculates test error, and returns all the data you need for visualization:

# Cost function (fixed case issue)
error_cost <- function(x, y, theta) {
  # x should already include the intercept column when passed here
  sum((x %*% theta - y)^2) / (2 * length(y))
}

# Batch Gradient Descent with history tracking
grad_d <- function(x_train, y_train, x_test = NULL, y_test = NULL, 
                   alpha = 0.006, epsilon = 1e-10, max_iter = 2000) {
  # Add intercept column to training data
  x_train <- cbind(rep(1, nrow(x_train)), x_train)
  n_features <- ncol(x_train)
  n_samples <- nrow(x_train)
  
  # Initialize theta, history storage
  theta <- matrix(1, nrow = n_features, ncol = 1)
  cost_history <- double(max_iter)
  theta_history <- vector("list", max_iter)
  test_error_history <- double(max_iter)
  
  # Initial cost calculation
  cost_history[1] <- error_cost(x_train, y_train, theta)
  # Initial test error (if test data provided)
  if (!is.null(x_test) && !is.null(y_test)) {
    x_test <- cbind(rep(1, nrow(x_test)), x_test)
    test_error_history[1] <- error_cost(x_test, y_test, theta)
  }
  
  delta <- 1
  iter <- 1
  
  while (delta > epsilon && iter < max_iter) {
    iter <- iter + 1
    # Update theta
    theta <- theta - (alpha / n_samples) * t(x_train) %*% (x_train %*% theta - y_train)
    
    # Calculate current training cost
    cost_history[iter] <- error_cost(x_train, y_train, theta)
    # Store current theta
    theta_history[[iter]] <- theta
    
    # Calculate test error if test data is provided
    if (!is.null(x_test) && !is.null(y_test)) {
      test_error_history[iter] <- error_cost(x_test, y_test, theta)
    }
    
    # Check cost change
    delta <- abs(cost_history[iter] - cost_history[iter - 1])
    
    # Stop if cost increases
    if (cost_history[iter] > cost_history[iter - 1]) {
      cat("Warning: Cost is increasing. Try reducing alpha.\n")
      # Trim history to actual iterations before returning
      cost_history <- cost_history[1:iter]
      theta_history <- theta_history[1:iter]
      test_error_history <- test_error_history[1:iter]
      return(list(theta = theta, iterations = iter, 
                  train_cost_history = cost_history,
                  test_error_history = test_error_history))
    }
  }
  
  # Trim history to actual iterations completed
  cost_history <- cost_history[1:iter]
  theta_history <- theta_history[1:iter]
  test_error_history <- test_error_history[1:iter]
  
  cat(sprintf("Completed in %i iterations.\n", iter))
  return(list(theta = theta, iterations = iter, 
              train_cost_history = cost_history,
              test_error_history = test_error_history))
}

# Prediction function (fixed to match intercept handling)
t_predict <- function(theta, x) {
  x <- cbind(rep(1, nrow(x)), x)
  return(x %*% theta)
}

Key Improvements Explained

  • Local history storage: All history vectors/lists are initialized inside the function, so no global variable conflicts.
  • Test error support: The function accepts optional test data and tracks test error alongside training cost.
  • Max iteration limit: Added max_iter to prevent infinite loops if epsilon is too small.
  • Clean history trimming: After convergence or early stop, we trim the history to only the iterations that actually ran.
  • Consistent naming: Switched to snake_case for R conventions (optional, but makes code more readable).

How to Use This Function

Let's test it with sample data (e.g., simple linear regression):

# Create sample training data
set.seed(123)
x_train <- matrix(rnorm(100), ncol = 1)
y_train <- 2 + 3*x_train + rnorm(100, 0, 0.5)

# Create sample test data
x_test <- matrix(rnorm(50), ncol = 1)
y_test <- 2 + 3*x_test + rnorm(50, 0, 0.5)

# Run gradient descent
gd_result <- grad_d(x_train, y_train, x_test, y_test, alpha = 0.05)

# View final theta
gd_result$theta

# View number of iterations
gd_result$iterations

Visualize Cost Convergence

Add this code to plot training cost (and test error if available):

library(ggplot2)

# Create data frame for plotting
plot_data <- data.frame(
  iteration = 1:gd_result$iterations,
  train_cost = gd_result$train_cost_history
)

# Add test error if available
if (!all(gd_result$test_error_history == 0)) {
  plot_data$test_error <- gd_result$test_error_history
}

# Plot
ggplot(plot_data, aes(x = iteration)) +
  geom_line(aes(y = train_cost, color = "Training Cost"), linewidth = 1) +
  {if (exists("test_error", plot_data)) geom_line(aes(y = test_error, color = "Test Error"), linewidth = 1)} +
  labs(title = "Cost Convergence During Batch Gradient Descent",
       x = "Iteration", y = "Cost/Error", color = "Metric") +
  theme_minimal()

This will give you a clear plot showing how training cost (and test error) decreases over iterations, confirming your gradient descent is working correctly.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 14:58:00