Trains a pigauto model: a gated ensemble of a phylogenetic baseline
and an attention-based graph neural network correction, implemented
as an internal torch module (ResidualPhyloDAE; "Residual" here
refers to the ResNet-style skip connections in the GNN layers, not
to a statistical residual). For continuous, count, and ordinal traits
the baseline is Brownian motion (phylogenetic correlation matrix); for binary
and categorical traits it is phylogenetic label propagation. Supports
all five trait types via a unified latent space.
Usage
fit_pigauto(
data,
tree,
splits = NULL,
graph = NULL,
baseline = NULL,
gnn = FALSE,
baseline_full = NULL,
hidden_dim = 64L,
k_eigen = "auto",
n_gnn_layers = 2L,
gate_cap = 0.8,
use_attention = TRUE,
use_transformer_blocks = TRUE,
n_heads = 4L,
ffn_mult = 4L,
use_trait_attention = FALSE,
n_trait_heads = 2L,
trait_embed_dim = 32L,
dropout = 0.1,
lr = 0.003,
weight_decay = 1e-04,
epochs = 3000L,
corruption_rate = 0.55,
corruption_start = 0.2,
corruption_ramp = 500L,
refine_steps = 8L,
lambda_shrink = 0.03,
lambda_gate = 0.01,
warmup_epochs = 200L,
edge_dropout = 0.1,
eval_every = 100L,
patience = 10L,
clip_norm = 1,
conformal_method = c("split", "bootstrap", "mondrian"),
conformal_bootstrap_B = 500L,
conformal_split_val = FALSE,
gate_method = c("cv_folds", "median_splits", "single_split"),
gate_splits_B = 31L,
gate_cv_folds = 5L,
safety_floor = TRUE,
phylo_signal_gate = TRUE,
phylo_signal_threshold = 0.2,
phylo_signal_method = c("lambda", "blomberg_k"),
min_val_cells = 20L,
lambda_mode = c("estimate", "fixed_1", "cv", "bayes"),
joint_solver = c("inhouse", "rphylopars"),
predict_method = c("auto", "exact", "per_column"),
joint_refine_iter = 0L,
verbose = TRUE,
seed = NULL
)Arguments
- data
object of class
"pigauto_data".- tree
object of class
"phylo".- splits
list (output of
make_missing_splits) orNULL.- graph
list (output of
build_phylo_graph) orNULL.- baseline
list (output of
fit_baseline) orNULL.- gnn
logical. When
TRUE, trains the attention-based GNN correction as described above. WhenFALSE(default), no GNN is constructed or trained andfit_pigauto()makes zerotorch::calls – the fit is the phylogenetic baseline alone, optionally re-weighted against a grand-mean floor.safety_floorandphylo_signal_gatekeep their usual semantics: the blend collapses to \(r_{BM} \cdot \mu_{BM} + r_{MEAN} \cdot \mu_{MEAN}\) (the GNN corner degenerates to the baseline, so its calibrated weight is folded into \(r_{BM}\)). The pure traditional-stats arm isgnn = FALSE, safety_floor = FALSE, phylo_signal_gate = FALSE.conformal_method = "mondrian"is not supported whengnn = FALSE(its locality statistic conditions on a calibrated GNN prediction surface that does not exist here) and raises an error. Usercovariatesenter pigauto only through the GNN, so undergnn = FALSEthey are ignored (with a warning). Without a validation split (splits = NULL) the fit is pure baseline (\(r_{BM} = 1\)) and carries no conformal scores. DefaultFALSEsince pigauto 0.11.0.9001: on simulated and real data at the current baseline defaults the GNN did not lower imputation error and took 20 to 90 times longer to fit (docs/dev-log/arc/2026-10-02-campaign-gnn-rerun.md). Setgnn = TRUEto train it, for example to usecovariates, which only the GNN uses, or to restore the behaviour of earlier versions.- baseline_full
list (output of
fit_baselinewithsplits = NULL) orNULL. Only used whengnn = FALSE: the production-mode baseline, fit on ALL observed cells (no val/test hold-out), used bypredict.pigauto_fit()for ordinary (non-evaluation) predictions.baselineitself always stays the held-out fit and is what every scorer (val_rmse,test_rmse,conformal_scores,evaluate()) reads. WhenNULL(default) andgnn = FALSE, it is computed internally viafit_baseline(data, tree, splits = NULL, ...); ignored whengnn = TRUEunless explicitly supplied (e.g. byimpute()), in which case it is stored but not used for training.integer. Hidden layer width (default
64).- k_eigen
integer. Number of spectral node features (default
8).- n_gnn_layers
integer. Number of graph message-passing layers (default
2). Each layer has its own learnable alpha gate, layer normalisation, and ResNet-style skip connection.- gate_cap
numeric. Upper bound for the per-column blend gate (default
0.8). Safety comes from regularisation, not the cap.- use_attention
logical. Use attention in the GNN layers (default
TRUE).- use_transformer_blocks
logical. Replace the legacy attention stack with pre-norm transformer-encoder blocks (multi-head attention
FFN + two residual skips). Default
TRUE. SetFALSEto reconstruct pre-v0.9.0 fits (single-head attention with a learnable alpha gate per layer).
- n_heads
integer. Number of attention heads when
use_transformer_blocks = TRUE(default4). Each head learns its own phylogenetic bandwidth (B2 rate-aware attention).- ffn_mult
integer. Feed-forward width multiplier inside each transformer block (default
4, givinghidden_dim * 4).- use_trait_attention
logical. Opt-in within-row cross-trait self-attention (B3, v0.9.3). When
TRUE(defaultFALSE), the model builds per-trait tokens from each row's latent values (linear projection + learnable positional embedding), applies one multi-head self-attention block over the trait sequence, mean-pools to atrait_embed_dimfeature, and concatenates it alongside(x, coords, covs)at the encoder input. Intended for trait sets with within-row functional coupling not represented by the selected baseline. Assess this opt-in architecture with held-out evaluation on the intended data. DefaultFALSEpreserves v0.9.2 behaviour exactly.- n_trait_heads
integer. Number of attention heads in the within-row self-attention block when
use_trait_attention = TRUE. Default2. Ignored whenuse_trait_attention = FALSE.- trait_embed_dim
integer. Embedding dim per trait token in the within-row self-attention block. Default
32. Ignored whenuse_trait_attention = FALSE.- dropout
numeric. Dropout rate (default
0.10).- lr
numeric. AdamW learning rate (default
0.003).- weight_decay
numeric. AdamW weight decay (default
1e-4).- epochs
integer. Maximum training epochs (default
3000).- corruption_rate
numeric. Final corruption fraction if
corruption_ramp > 0; otherwise the fixed corruption rate per epoch (default0.55).- corruption_start
numeric. Initial corruption fraction for the curriculum schedule (default
0.20). Ignored ifcorruption_ramp = 0.- corruption_ramp
integer. Epochs over which corruption linearly ramps from
corruption_starttocorruption_rate(default500). Set to0for fixed corruption.- refine_steps
integer. Iterative refinement steps at inference (default
8). Gate calibration and conformal scoring use the same number of steps so the calibrated surface matches predict.- lambda_shrink
numeric. Weight on the shrinkage penalty
||delta - baseline||^2that keeps the GNN correction close to the phylogenetic baseline (default0.03).- lambda_gate
numeric. Weight on the gate regularisation penalty that pushes learnable gates toward zero. Prevents gates from staying open when the GNN provides no useful correction (default
0.01).- warmup_epochs
integer. Linear learning-rate warmup over the first N epochs (default
200). After warmup, a cosine schedule decays the LR to1e-5.- edge_dropout
numeric. Fraction of adjacency edges randomly zeroed each training epoch for graph regularisation (default
0.1). Set to0to disable.- eval_every
integer. Evaluate on val every N epochs (default
100).- patience
integer. Early-stopping patience in eval cycles (default
10).- clip_norm
numeric. Gradient clip norm (default
1.0).- conformal_method
character. How the conformal residual score is estimated from held-out validation cells.
"split"(default, backward-compatible) takes a sample quantile;"bootstrap"averages quantiles fromconformal_bootstrap_Bbootstrap resamples; and"mondrian"stratifies validation cells by phylogenetic sampling locality before estimating a score per stratum. These choices support nominal held-out diagnostics, not package-certified coverage."mondrian"is single-observation only because its locality is species-level. If either stratum is too small, that trait uses the global"split"score instead.- conformal_bootstrap_B
integer. Bootstrap resamples used when
conformal_method = "bootstrap"; default500. Ignored otherwise.- conformal_split_val
logical. Default
FALSE: gate selection and residual-score estimation share validation cells. WhenTRUE, each eligible latent column separates those tasks into calibration and residual-estimation subsets. This separation does not certify package-wide coverage. Columns without enough validation cells retain the shared-data path. With either setting, usecross_validate()and report held-out diagnostics for the intended data.- gate_method
character. How the per-trait calibrated gate is chosen.
"single_split"runs the grid search on a single random half-A / half-B split of the val rows;"median_splits"repeats the whole procedure forgate_splits_Brandom splits and takes the medianbest_g."cv_folds"(default since PR #102, 2026-05-17) partitions val cells intogate_cv_folds(default 5) deterministic non-overlapping folds and runs the grid + half-B-verify procedure once per fold (training set = K-1 folds, held-out = remaining fold), taking the componentwise median of K winning weight vectors."cv_folds"uses larger training sets per split (K-1/K vs 1/2 inmedian_splits) and has a standard cross-validation interpretation, motivated by the open val→test drift observed on 4/32 binary cells in the discrete-bench memo."median_splits"repeats the split procedure before taking a componentwise median. Assess calibration settings on held-out data.- gate_splits_B
integer. Random splits used when
gate_method = "median_splits"; default31(odd so the median is well-defined).- gate_cv_folds
integer. Number of CV folds when
gate_method = "cv_folds"; default5, must be>= 2. Capped atn_valper trait so each fold has at least 1 cell. When effective K< 2(e.g.n_val = 1), the code falls back to a single split.- safety_floor
logical. When
TRUE(default), post-training calibration searches a 3-way simplex of BM, GNN, and grand-mean candidates. Because the grand-mean corner is always in the grid, the selected candidate cannot be worse than that corner on the validation cells under the calibration metric. WhenFALSE, the v0.9.1 1-D calibration is used exactly (r_MEAN = 0).- phylo_signal_gate
logical. When
TRUE(default since v0.9.1.9003), compute per-trait Pagel's \(\lambda\) on training-observed cells before fitting; for traits withlambda < phylo_signal_threshold, force(r_cal_bm = 0, r_cal_gnn = 0, r_cal_mean = 1)directly and skip BM + GNN training on those traits. Requires thephytoolspackage. Falls back to safety-floor-only behaviour (phylo_signal_gate = FALSEeffective) whenphytoolsis absent.- phylo_signal_threshold
numeric, default
0.2. Traits with Pagel's \(\lambda\) below this value are routed to the grand-mean corner of the safety-floor simplex.- phylo_signal_method
character, currently only
"lambda"is fully implemented. Reserved"blomberg_k"path returns Blomberg's K viaphytools::phylosig()but uses the same threshold — which is NOT dimensionally comparable; users selecting K must supply a K-appropriate threshold.- min_val_cells
integer. Warn at fit time if any trait has fewer than
min_val_cellsvalidation cells available for gate calibration and conformal-score estimation. Default20: with fewer than 19 cells the split-conformal level \(n/(n+1)\) is below 0.95; with 38 or fewer cells the conformal score is the largest validation residual; and with very few cells gate calibration becomes essentially a coin flip between0andgate_cap. Recommended operational target isn_val >= 20-30per trait; achieve this by increasingmissing_fracor collecting more species. See Calibration at small n below.- lambda_mode
character. Pagel-lambda mode for the BM baseline.
"estimate"(default) fits a per-trait Pagel's lambda on each continuous-family (BM-eligible) latent column – continuous, count, proportion, and zi_count magnitude, via the joint MVN / threshold-joint baseline's ownlambda_colsmachinery when that joint path fires, or a per-column re-fit otherwise. Discrete traits (binary, categorical, zi gate) AND ordinal have no discrete-trait analogue of Pagel's lambda of their own – there is no lambda_k to estimate for them. Underpredict_method = "per_column"(see below) they stay fixed at lambda = 1 in every path, matching pre-S3 behaviour. Underpredict_method = "exact"(and for traits the default"auto"routes to"exact") they instead share the joint fit'slambda_block– the exact conditional's covariance model uses ONE shared phylogenetic correlation matrixR(lambda_block)for every column, so the Sigma estimate feeding it must itself come from an internally consistent init, not a mix ofR(1)for discrete columns andR(lambda_k)for continuous ones (Shinichi's decision, docs/dev-log/exact-default/S3-default-report.md)."fixed_1"preserves the pre-lambda Brownian correlation matrix everywhere;"cv"and"bayes"are alternative per-column estimators that force continuous-family columns onto the per-column BM path (no joint analogue). Whenpredict_method = "exact"orjoint_refine_iter > 0, the joint (multi-trait) prediction path additionally uses the single sharedlambda_blockfor cross-trait computations that need one common phylogenetic correlation matrix, never overriding a continuous-family column's own estimated lambda_k. Passed tofit_baselineand stored in the fitted model config. Whencovariatesare supplied, the covariate-aware BM path (bm_impute_col_with_cov()) accepts a numeric lambda or"estimate", solambda_mode %in% c("estimate", "fixed_1")reaches it and each covariate-aware BM-eligible column gets its own estimated lambda; it does NOT accept"cv"/"bayes", which fall back to lambda = 1 for BM-eligible columns with a warning fromfit_baseline.- joint_solver
character. Which solver estimates the joint MVN / threshold-joint / OVR categorical baselines.
"inhouse"(default) is byte-identical to prior releases;"rphylopars"delegates toRphylopars::phylopars()'s converged REML fit, with automatic fallback to"inhouse"on failure. Passed tofit_baselineand stored in the fitted model config. Seedocs/dev-log/2026-08-16-continuous-gap-diagnosis.md.- predict_method
character. Prediction route for the in-house joint solver, passed to
fit_baseline."auto"(default, S5b) fits the baseline with both the"exact"and"per_column"routes and picks, per trait, whichever has the lower loss on that trait's validation cells (mean squared error on the z-scored latent scale for continuous/count/proportion/ordinal/zi_count magnitude; mean log-loss for binary/zi_count gate; mean multinomial log-loss for categorical), falling back to"exact"on a tie or when fewer than 5 validation cells belong to that trait (including when there is no validation split at all)."exact"uses the full cross-trait conditional mean and variance ofvec(L) ~ MVN(0, Sigma %x% R(lambda_block))(Hadfield & Nakagawa, 2010 sparse precision form), each column GLS-mean-centred atlambda_blockbefore the solve; discrete liability columns sharelambda_blockunder this route (seelambda_modeabove). Falls back to"per_column"above roughly 20000 unknown cells (roughly 4000 species at 5 traits), on a singular/unusable Sigma, or when fewer than 2 joint columns or no Henderson sparse precision are available; prints a one-timemessage()per R session whenpredict_methodwas left at its default, or awarning()every time when"exact"was requested explicitly."per_column"retains the original per-column conditional prediction route (no cross-trait borrowing in the prediction step). None of the three options change covariance estimation or the"rphylopars"solver. The route actually used is recorded in$model_config$predict_method_used("exact","per_column", or"auto") and, per trait, in$model_config$predict_method_by_trait. Under"auto"with real validation cells,$model_config $route_val_n/$score_val_nrecord, per trait, how many of its validation cells chose the route versus were reserved for gate calibration and conformal scoring (seefit_baseline'spredict_methoddocs); both areNULLunder an explicit"exact"/"per_column"request.- joint_refine_iter
integer, default
0L. Enables cross-trait refinement of the joint baseline's cell imputations using the estimated Sigma (the in-house solver'smax_iterEM cell-refinement;R/joint_mvn_solver.R).0Lpreserves current behaviour byte-for-byte. The refinement is guarded: the Sigma step must shrink each iteration, or the loop rolls back to the last good iterate and sets$diverged. Assess this opt-in control with held-out evaluation on the intended data. Has no effect on any trait predicted by the"exact"route (seefit_baseline'spredict_methoddocs); under the defaultpredict_method = "auto"it therefore only applies to traits"auto"routes to"per_column".- verbose
logical. Print training progress (default
TRUE).- seed
optional integer. When supplied, makes stochastic training and calibration reproducible; the default
NULLuses the current RNG stream.
Details
Blend formulation: The prediction is \(\hat{x} = (1-r)\mu + r\delta\), where \(\mu\) is the BM baseline, \(\delta\) is the model's direct prediction, and \(r = \sigma(\rho) \times \mathrm{cap}\) is a per-column learnable gate bounded in \((0, \mathrm{gate\_cap})\). When \(r = 0\), the prediction collapses to the baseline. The gate is regularised toward zero via the shrinkage penalty on \(\delta - \mu\), so the model defaults to the baseline unless the GNN's correction demonstrably helps on the validation set.
Training objective (per epoch):
A random subset of observed cells is corrupted with a learnable mask token.
The model predicts \(\delta\) from graph context.
Loss = type-specific reconstruction on corrupted cells +
lambda_shrink* MSE(\(\delta - \mu\)) +lambda_gate* MSE(\(r\)).
The gate penalty on \(r\) is necessary because when \(\delta = \mu\) (the BM-optimal solution for observed cells), the reconstruction and shrinkage losses both equal zero regardless of \(r\), leaving no gradient to close the gate. The explicit penalty ensures gates default toward zero when the GNN correction provides no benefit.
Type-specific losses:
- continuous/count/ordinal
MSE
- binary
BCE with logits
- categorical
cross-entropy over K latent columns
Calibration at small n
Conformal bounds use the empirical (1 - alpha) quantile of
\(|y - \hat y|\) on held-out validation cells. They are nominal
held-out diagnostics and do not certify package-wide coverage.
Small validation sets make both the residual quantile and gate selection
unstable. The min_val_cells warning identifies traits where this is
especially relevant. conformal_split_val = TRUE can separate gate
selection from residual estimation for eligible columns, but it does not
remove the small-sample limitation. Use cross_validate() on the
intended data and report the resulting held-out diagnostics rather than
treating an interval target as certified.
Examples
# \donttest{
data(avonet300, tree300)
tree <- ape::keep.tip(tree300, tree300$tip.label[seq_len(30L)])
traits <- avonet300[match(tree$tip.label, avonet300$Species_Key),
c("Mass", "Wing.Length"), drop = FALSE]
rownames(traits) <- tree$tip.label
data <- preprocess_traits(traits, tree)
splits <- make_missing_splits(data$X_scaled, trait_map = data$trait_map,
seed = 1)
fit <- fit_pigauto(data, tree, splits = splits, epochs = 5L,
verbose = FALSE, seed = 1)
#> Warning: phylo_signal_gate requires the 'phytools' package; returning NA for all traits.
#> Warning: Small validation set for 2 trait(s): Mass (n=1), Wing.Length (n=2). Calibrated gate and conformal scores will be noisy for these trait(s). 95% split-conformal coverage is NOT achievable for 2 trait(s) with fewer than 19 validation cells (Mass (n=1), Wing.Length (n=2)): the achievable ceiling is n_val / (n_val + 1), which only reaches 0.95 at n_val >= 19. See `?fit_pigauto` under 'Calibration at small n' for smoothing options.
# }
