diff --git a/tests/testthat/helper-pffr-ncv.R b/tests/testthat/helper-pffr-ncv.R index f7194553..b9d6fd25 100644 --- a/tests/testthat/helper-pffr-ncv.R +++ b/tests/testthat/helper-pffr-ncv.R @@ -24,6 +24,50 @@ ncv_test_fit <- function(dat, ...) { ) } +# Response on the scale of `family`, driven by the dependent Gaussian data. +ncv_test_glm_data <- function(family) { + dat <- ncv_test_data() + eta <- 0.5 * scale(dat$Y) + y <- if (family$family == "poisson") { + rpois(length(eta), 3 * exp(eta)) + } else { + rbinom(length(eta), 1, plogis(eta - 0.5)) + } + dat$Y <- I(matrix(y, nrow(eta))) + dat +} + +# Total penalty matrix sum_j sp_j S_j of a fitted model, in coefficient order. +ncv_test_penalty <- function(fit) { + p <- length(fit$coefficients) + S <- matrix(0, p, p) + j <- 0L + for (sm in fit$smooth) { + ii <- seq.int(sm$first.para, sm$last.para) + for (penalty in sm$S) { + j <- j + 1L + S[ii, ii] <- S[ii, ii] + fit$sp[j] * penalty + } + } + stopifnot(j == length(fit$sp)) + S +} + +# Penalized IRLS at fixed penalty S (minimizes deviance + beta' S beta). +ncv_test_pirls <- function(X, y, S, family, beta, tol = 1e-12) { + for (iter in 1:100) { + eta <- as.numeric(X %*% beta) + mu <- family$linkinv(eta) + deta <- family$mu.eta(eta) + w <- deta^2 / family$variance(mu) + z <- eta + (y - mu) / deta + new <- solve(crossprod(X, w * X) + S, crossprod(X, w * z)) + if (max(abs(new - beta)) < tol) break + beta <- new + } + as.numeric(new) +} + ncv_test_groups <- function(nei) { split(nei$a, rep(seq_along(nei$ma), diff(c(0L, nei$ma)))) } diff --git a/tests/testthat/test-pffr-ncv.R b/tests/testthat/test-pffr-ncv.R index 80cd6493..74a4801f 100644 --- a/tests/testthat/test-pffr-ncv.R +++ b/tests/testthat/test-pffr-ncv.R @@ -212,16 +212,7 @@ test_that("Gaussian NCV equals brute-force fixed-penalty curve deletion loss", { dat <- ncv_test_data() fit <- ncv_test_fit(dat) X <- predict(fit, type = "lpmatrix", reformat = FALSE) - S <- matrix(0, ncol(X), ncol(X)) - j <- 0L - for (sm in fit$smooth) { - ii <- seq.int(sm$first.para, sm$last.para) - for (penalty in sm$S) { - j <- j + 1L - S[ii, ii] <- S[ii, ii] + fit$sp[j] * penalty - } - } - expect_equal(j, length(fit$sp)) + S <- ncv_test_penalty(fit) expect_equal( as.numeric(solve(crossprod(X) + S, crossprod(X, fit$y))), as.numeric(fit$coefficients), @@ -251,6 +242,47 @@ test_that("Gaussian NCV equals brute-force fixed-penalty curve deletion loss", { expect_equal(predictions, as.numeric(implied), tolerance = 1e-7) }) +test_that("Poisson and binomial NCV approximate fixed-penalty curve deletion", { + skip_on_cran() + skip_if_not_installed("mgcv", "1.9.0") + # For non-Gaussian families mgcv replaces each deletion refit by a Newton + # step from the full fit (?mgcv::NCV), so agreement is approximate. The + # tolerances sit about 2x above the observed error on these data (loss + # within 0.2%, deletion predictions within 0.013 on the link scale), while + # the in-sample predictor is 0.15 to 0.5 away: a check that ignored the + # curve blocks would fail. + for (family in list(poisson(), binomial())) { + set.seed(102) + dat <- ncv_test_glm_data(family) + fit <- ncv_test_fit(dat, family = family) + X <- predict(fit, type = "lpmatrix", reformat = FALSE) + S <- ncv_test_penalty(fit) + expect_equal( + ncv_test_pirls(X, fit$y, S, family, fit$coefficients), + as.numeric(fit$coefficients), + tolerance = 1e-7 + ) + predictions <- numeric(nrow(X)) + for (rows in ncv_test_groups(fit$pffr$ncv$nei)) { + beta <- ncv_test_pirls( + X[-rows, ], + fit$y[-rows], + S, + family, + fit$coefficients + ) + predictions[rows] <- as.numeric(X[rows, ] %*% beta) + } + loss <- sum(family$dev.resids(fit$y, family$linkinv(predictions), 1)) + expect_equal(as.numeric(fit$gcv.ubre), loss, tolerance = 1e-2) + implied <- attr(fit$gcv.ubre, "eta.cv") + if (is.null(implied)) next + error <- max(abs(implied - predictions)) + expect_lt(error, 0.03) + expect_gt(max(abs(implied - fit$linear.predictors)), 5 * error) + } +}) + test_that("NCV supports CL2 coefficients and model covariance fallback", { skip_on_cran() skip_if_not_installed("mgcv", "1.9.0")