From 617a1fe8cfceb8351bfccfc0a0f1a534e154e78d Mon Sep 17 00:00:00 2001 From: Erik Date: Mon, 17 Aug 2026 15:45:48 +1000 Subject: [PATCH 1/2] Add refit option to cv.balnet --- r-package/balnet/R/cv.balnet.R | 10 +++++++++- r-package/balnet/man/cv.balnet.Rd | 5 +++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/r-package/balnet/R/cv.balnet.R b/r-package/balnet/R/cv.balnet.R index d32c821..dd78336 100644 --- a/r-package/balnet/R/cv.balnet.R +++ b/r-package/balnet/R/cv.balnet.R @@ -5,6 +5,9 @@ #' @param type.measure The loss to minimize for cross-validation. #' Default is balance loss (e.g., Zhiqiang (2020)). #' For "imbalance", the criterion is mean covariate imbalance (e.g., Wang & Zubizarreta (2020)). +#' @param refit Whether to refit the model on each training fold, default is TRUE. +#' If FALSE, weights are computed once on full data and the loss is evaluated +#' per subsample (e.g., imbalance is measured on each data subsample). #' @param nfolds The number of folds used for cross-validation, default is 10. #' @param foldid An optional `n`-vector specifying which fold 1 to `nfold` a sample belongs to. #' If NULL, this defaults to `sample(rep(seq(nfolds), length.out = nrow(X)))`. @@ -45,6 +48,7 @@ cv.balnet <- function( X, W, type.measure = c("balance.loss", "imbalance"), + refit = TRUE, nfolds = 10, foldid = NULL, ... @@ -80,7 +84,11 @@ cv.balnet <- function( X.train <- X[train, , drop = FALSE] W.train <- W[train] dot.args[["sample.weights"]] <- sample.weights[train] - fit.train <- do.call(balnet, c(list(X = X.train, W = W.train, standardize = ".inplace"), dot.args)) + if (refit) { + fit.train <- do.call(balnet, c(list(X = X.train, W = W.train, standardize = ".inplace"), dot.args)) + } else { + fit.train <- fit.full + } X.test <- X[test, , drop = FALSE] W.test <- W[test] diff --git a/r-package/balnet/man/cv.balnet.Rd b/r-package/balnet/man/cv.balnet.Rd index 6ce378e..60e6482 100644 --- a/r-package/balnet/man/cv.balnet.Rd +++ b/r-package/balnet/man/cv.balnet.Rd @@ -8,6 +8,7 @@ cv.balnet( X, W, type.measure = c("balance.loss", "imbalance"), + refit = TRUE, nfolds = 10, foldid = NULL, ... @@ -22,6 +23,10 @@ cv.balnet( Default is balance loss (e.g., Zhiqiang (2020)). For "imbalance", the criterion is mean covariate imbalance (e.g., Wang & Zubizarreta (2020)).} +\item{refit}{Whether to refit the model on each training fold, default is TRUE. +If FALSE, weights are computed once on full data and the loss is evaluated +per subsample (e.g., imbalance is measured on each data subsample).} + \item{nfolds}{The number of folds used for cross-validation, default is 10.} \item{foldid}{An optional \code{n}-vector specifying which fold 1 to \code{nfold} a sample belongs to. From 646e2b5e031bacb665549d9a575d7d693679cf7d Mon Sep 17 00:00:00 2001 From: Erik Date: Mon, 17 Aug 2026 15:46:13 +1000 Subject: [PATCH 2/2] -> tuning --- r-package/balnet/R/cv.balnet.R | 2 +- r-package/balnet/man/cv.balnet.Rd | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/r-package/balnet/R/cv.balnet.R b/r-package/balnet/R/cv.balnet.R index dd78336..bba72c5 100644 --- a/r-package/balnet/R/cv.balnet.R +++ b/r-package/balnet/R/cv.balnet.R @@ -1,4 +1,4 @@ -#' Cross-validation for balnet. +#' Tuning for balnet. #' #' @param X A numeric matrix or data frame with pre-treatment covariates. #' @param W Treatment vector (0: control, 1: treated). diff --git a/r-package/balnet/man/cv.balnet.Rd b/r-package/balnet/man/cv.balnet.Rd index 60e6482..c2bfc83 100644 --- a/r-package/balnet/man/cv.balnet.Rd +++ b/r-package/balnet/man/cv.balnet.Rd @@ -2,7 +2,7 @@ % Please edit documentation in R/cv.balnet.R \name{cv.balnet} \alias{cv.balnet} -\title{Cross-validation for balnet.} +\title{Tuning for balnet.} \usage{ cv.balnet( X, @@ -38,7 +38,7 @@ If NULL, this defaults to \code{sample(rep(seq(nfolds), length.out = nrow(X)))}. A fit cv.balnet object. } \description{ -Cross-validation for balnet. +Tuning for balnet. } \examples{ \donttest{