Skip to contents

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) or NULL.

graph

list (output of build_phylo_graph) or NULL.

baseline

list (output of fit_baseline) or NULL.

gnn

logical. When TRUE, trains the attention-based GNN correction as described above. When FALSE (default), no GNN is constructed or trained and fit_pigauto() makes zero torch:: calls – the fit is the phylogenetic baseline alone, optionally re-weighted against a grand-mean floor. safety_floor and phylo_signal_gate keep 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 is gnn = FALSE, safety_floor = FALSE, phylo_signal_gate = FALSE. conformal_method = "mondrian" is not supported when gnn = FALSE (its locality statistic conditions on a calibrated GNN prediction surface that does not exist here) and raises an error. User covariates enter pigauto only through the GNN, so under gnn = FALSE they are ignored (with a warning). Without a validation split (splits = NULL) the fit is pure baseline (\(r_{BM} = 1\)) and carries no conformal scores. Default FALSE since 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). Set gnn = TRUE to train it, for example to use covariates, which only the GNN uses, or to restore the behaviour of earlier versions.

baseline_full

list (output of fit_baseline with splits = NULL) or NULL. Only used when gnn = FALSE: the production-mode baseline, fit on ALL observed cells (no val/test hold-out), used by predict.pigauto_fit() for ordinary (non-evaluation) predictions. baseline itself always stays the held-out fit and is what every scorer (val_rmse, test_rmse, conformal_scores, evaluate()) reads. When NULL (default) and gnn = FALSE, it is computed internally via fit_baseline(data, tree, splits = NULL, ...); ignored when gnn = TRUE unless explicitly supplied (e.g. by impute()), in which case it is stored but not used for training.

hidden_dim

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. Set FALSE to 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 (default 4). Each head learns its own phylogenetic bandwidth (B2 rate-aware attention).

ffn_mult

integer. Feed-forward width multiplier inside each transformer block (default 4, giving hidden_dim * 4).

use_trait_attention

logical. Opt-in within-row cross-trait self-attention (B3, v0.9.3). When TRUE (default FALSE), 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 a trait_embed_dim feature, 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. Default FALSE preserves 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. Default 2. Ignored when use_trait_attention = FALSE.

trait_embed_dim

integer. Embedding dim per trait token in the within-row self-attention block. Default 32. Ignored when use_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 (default 0.55).

corruption_start

numeric. Initial corruption fraction for the curriculum schedule (default 0.20). Ignored if corruption_ramp = 0.

corruption_ramp

integer. Epochs over which corruption linearly ramps from corruption_start to corruption_rate (default 500). Set to 0 for 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||^2 that keeps the GNN correction close to the phylogenetic baseline (default 0.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 to 1e-5.

edge_dropout

numeric. Fraction of adjacency edges randomly zeroed each training epoch for graph regularisation (default 0.1). Set to 0 to 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 from conformal_bootstrap_B bootstrap 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"; default 500. Ignored otherwise.

conformal_split_val

logical. Default FALSE: gate selection and residual-score estimation share validation cells. When TRUE, 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, use cross_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 for gate_splits_B random splits and takes the median best_g. "cv_folds" (default since PR #102, 2026-05-17) partitions val cells into gate_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 in median_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"; default 31 (odd so the median is well-defined).

gate_cv_folds

integer. Number of CV folds when gate_method = "cv_folds"; default 5, must be >= 2. Capped at n_val per 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. When FALSE, 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 with lambda < 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 the phytools package. Falls back to safety-floor-only behaviour (phylo_signal_gate = FALSE effective) when phytools is 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 via phytools::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_cells validation cells available for gate calibration and conformal-score estimation. Default 20: 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 between 0 and gate_cap. Recommended operational target is n_val >= 20-30 per trait; achieve this by increasing missing_frac or 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 own lambda_cols machinery 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. Under predict_method = "per_column" (see below) they stay fixed at lambda = 1 in every path, matching pre-S3 behaviour. Under predict_method = "exact" (and for traits the default "auto" routes to "exact") they instead share the joint fit's lambda_block – the exact conditional's covariance model uses ONE shared phylogenetic correlation matrix R(lambda_block) for every column, so the Sigma estimate feeding it must itself come from an internally consistent init, not a mix of R(1) for discrete columns and R(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). When predict_method = "exact" or joint_refine_iter > 0, the joint (multi-trait) prediction path additionally uses the single shared lambda_block for cross-trait computations that need one common phylogenetic correlation matrix, never overriding a continuous-family column's own estimated lambda_k. Passed to fit_baseline and stored in the fitted model config. When covariates are supplied, the covariate-aware BM path (bm_impute_col_with_cov()) accepts a numeric lambda or "estimate", so lambda_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 from fit_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 to Rphylopars::phylopars()'s converged REML fit, with automatic fallback to "inhouse" on failure. Passed to fit_baseline and stored in the fitted model config. See docs/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 of vec(L) ~ MVN(0, Sigma %x% R(lambda_block)) (Hadfield & Nakagawa, 2010 sparse precision form), each column GLS-mean-centred at lambda_block before the solve; discrete liability columns share lambda_block under this route (see lambda_mode above). 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-time message() per R session when predict_method was left at its default, or a warning() 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_n record, per trait, how many of its validation cells chose the route versus were reserved for gate calibration and conformal scoring (see fit_baseline's predict_method docs); both are NULL under 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's max_iter EM cell-refinement; R/joint_mvn_solver.R). 0L preserves 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 (see fit_baseline's predict_method docs); under the default predict_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 NULL uses the current RNG stream.

Value

An object of class "pigauto_fit".

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):

  1. A random subset of observed cells is corrupted with a learnable mask token.

  2. The model predicts \(\delta\) from graph context.

  3. 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.
# }