library(FNN)
library(glmnet)
set.seed(43252)
n <- 400
p <- 30
lambda <- 10^seq(0, -4, length.out = 41)
results <- data.frame(
Setting = c("Independent", "Correlated"),
KNN_MSE = NA_real_, Lasso_MSE = NA_real_,
Lambda = NA_real_, Nonzero = NA_integer_
)
for (l in 1:2) {
Sigma <- diag(p)
if (l == 2) Sigma <- 0.8 + 0.2 * diag(p)
X <- matrix(rnorm((n + 1000) * p), ncol = p) %*% chol(Sigma)
y <- 0.5 * X[, 1] + sin(X[, 2]) - 0.3 * X[, 3]^2 +
rnorm(n + 1000, sd = 0.1)
test.X <- X[(n + 1):(n + 1000), ]
test.y <- y[(n + 1):(n + 1000)]
X <- X[1:n, ]
y <- y[1:n]
knn.fit <- knn.reg(X, test.X, y, k = 5, algorithm = "brute")
results$KNN_MSE[l] <- mean((test.y - knn.fit$pred)^2)
# Estimate means and scales separately in each training fold.
fold <- sample(rep(1:10, each = n / 10))
cv.mse <- matrix(NA_real_, nrow = 10, ncol = length(lambda))
for (m in 1:10) {
X.mean <- colMeans(X[fold != m, ])
X.fit <- sweep(X[fold != m, ], 2, X.mean, "-")
X.sd <- sqrt(colMeans(X.fit^2))
X.fit <- sweep(X.fit, 2, X.sd, "/")
X.valid <- sweep(X[fold == m, ], 2, X.mean, "-")
X.valid <- sweep(X.valid, 2, X.sd, "/")
y.mean <- mean(y[fold != m])
lasso.fit <- glmnet(
X.fit, y[fold != m] - y.mean, alpha = 1, lambda = lambda,
intercept = FALSE, standardize = FALSE
)
y.hat <- predict(lasso.fit, newx = X.valid, s = lambda) + y.mean
cv.mse[m, ] <- colMeans((y[fold == m] - y.hat)^2)
}
lambda.min <- lambda[which.min(colMeans(cv.mse))]
# Refit using all training observations and keep these transformations.
X.mean <- colMeans(X)
X.fit <- sweep(X, 2, X.mean, "-")
X.sd <- sqrt(colMeans(X.fit^2))
X.fit <- sweep(X.fit, 2, X.sd, "/")
test.X.fit <- sweep(test.X, 2, X.mean, "-")
test.X.fit <- sweep(test.X.fit, 2, X.sd, "/")
y.mean <- mean(y)
lasso.fit <- glmnet(
X.fit, y - y.mean, alpha = 1, lambda = lambda,
intercept = FALSE, standardize = FALSE
)
test.pred <- as.vector(
predict(lasso.fit, newx = test.X.fit, s = lambda.min)
) + y.mean
results$Lasso_MSE[l] <- mean((test.y - test.pred)^2)
results$Lambda[l] <- lambda.min
results$Nonzero[l] <- sum(coef(lasso.fit, s = lambda.min)[-1, 1] != 0)
}
knitr::kable(results, digits = 4)