Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 44 additions & 0 deletions tests/testthat/helper-pffr-ncv.R
Original file line number Diff line number Diff line change
Expand Up @@ -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))))
}
Expand Down
52 changes: 42 additions & 10 deletions tests/testthat/test-pffr-ncv.R
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down Expand Up @@ -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")
Expand Down
Loading