Skip to content

Implement staged warm-up adaptation schedule - #127

Merged
matt-graham merged 26 commits into
UCL:mainfrom
jennajiali:adapter-staging
Jul 22, 2026
Merged

matt-graham merged 26 commits into
UCL:mainfrom
jennajiali:adapter-staging

Conversation

@jennajiali

Copy link
Copy Markdown
Contributor

This PR implements the staged warm-up adaptation interface proposed in issue #79, allowing users to specify different adapter configurations at different points during warm-up. It builds on the partial implementation started on the mmg/adapter-staging branch and ensures adapter states correctly carry over across stage boundaries.

Background

Previously, sample_chain() ran all adapters from iteration 1 of warm-up with no way to vary which adapters were active over time. This is suboptimal in practice: the shape adapter (which estimates the proposal covariance) should not start until the scale adapter has had enough iterations to stabilise, because early covariance estimates based on poorly-mixed samples can make the sampler worse rather than better.

Furthermore, previously all adapter initialize() functions ignored the current proposal parameters on re-initialisation, causing the scale and shape estimates accumulated during one staged warm-up stage to be discarded at the start of the next. This contradicted the requirement in issue #79 that adaptation parameters be initialised at the values at the end of the previous stage.

Changes

R/chains.R

  • check_and_process_adapters() — new internal function that normalises the adapters argument into a canonical list of stage specifications, each with $adapters (a list of adapter objects) and $n_iteration (an integer). Accepts three input forms:
    1. A flat list of adapter objects (existing behaviour, now wrapped into a single stage)
    2. A staged list where each element is a list of adapters with an optional trailing integer iteration count; the last stage may omit the count and receives the remaining warm-up iterations
    3. A function accepting n_warm_up_iteration and returning a staged list (used by progressive_adaptation_schedule())
  • sample_chain() — warm-up loop now iterates over stages produced by check_and_process_adapters(), passing each stage's adapter list and iteration count to chain_loop() separately. Several bugs in the initial partial implementation are fixed:
    • Bootstrap warm_up_results initialised with key $final_state (not $state) to be consistent with chain_loop()'s return value
    • Each stage now reads stage$n_iteration rather than always using the full n_warm_up_iteration
    • Each stage now passes stage$adapters (clean list of adapter objects) rather than the whole named stage list
    • Progress bar label now shows current stage number (e.g. Warm-up (stage 1/3))
  • combine_warm_up_results() — rewritten to use rbind instead of mapply(c, ...). This correctly handles the bootstrap case (rbind(NULL, matrix) returns the matrix unchanged) and uses the correct $final_state key throughout
  • combine_stage_results() — fixed $state → $final_state to match chain_loop()'s return value
  • @param adapters documentation — updated to describe all three input forms

R/adaptation.R

  • progressive_adaptation_schedule() — new exported convenience function that returns a schedule constructor. When called with n_warm_up_iteration, it produces a three-stage adaptation schedule following common practice:
    • Stage 1: scale adapter only (shape held fixed at identity), for n_fixed_shape_iteration iterations (default 50)
    • Stage 2: scale adapter + diagonal shape adapter (shape_adapter("variance")), for n_diagonal_shape_iteration iterations (default 50)
    • Stage 3: scale adapter + dense shape adapter (shape_adapter("covariance")), for remaining iterations
    • If the sum of stage 1 and stage 2 counts exceeds n_warm_up_iteration, both are reduced proportionally and stage 3 is skipped. All adapter objects and iteration counts are customisable via arguments. Includes full roxygen documentation and a working example.
  • stochastic_approximation_scale_adapter and dual_averaging_scale_adapter: when initial_scale is not user-specified, initialize() now reads proposal$parameters()$scale and uses it as the starting log-scale if available, falling back to proposal$default_initial_scale() only when no current scale exists. Note that for dual_averaging_scale_adapter, mu, smoothed_log_scale and accept_prob_error already persisted across stages via closure scope and required no change.
  • variance_shape_adapter and covariance_shape_adapter: initialize() now follows a three-priority order:
    1. Explicit initial_shape constructor argument if supplied.
    2. Current proposal shape, guarded by type compatibility checks. variance_shape_adapter requires a vector of matching length to prevent accidentally carrying over a matrix-valued shape. covariance_shape_adapter requires a matrix of matching dimensions (Cholesky factor) to prevent carrying over a vector-valued shape.
    3. Default fallback (unit variances or identity matrix, respectively).
  • Added initial_shape parameter (numeric vector of per-dimension scales for variance; lower-triangular Cholesky factor matrix for covariance) to the respective function signatures and roxygen documentation, addressing the request to allow users to supply domain-knowledge-based starting values.

tests/testthat/test-chains.R

New tests appended after existing tests:

  • Unit tests for check_and_process_adapters() covering all three input forms, the last-stage remainder logic, and all error cases (non-last stage missing count, counts not summing to n_warm_up_iteration, invalid input types)
  • Integration tests for sample_chain() with staged adapters: correct warm-up row counts across two and three stages, correct behaviour with differing adapter sets across stages when trace_warm_up = FALSE, function-form adapters argument, correct final_state type, and invalid adapters error propagation

tests/testthat/test-adaptation.R

New tests appended after existing tests covering progressive_adaptation_schedule():

  • Returns a function
  • Correct stage count and iteration counts for the normal case (n_warm_up > n_fixed + n_diagonal), the exact-boundary case (no dense stage), and the fallback case (n_warm_up < n_fixed + n_diagonal)
  • Edge case of n_warm_up_iteration = 1
  • Custom n_fixed_shape_iteration and n_diagonal_shape_iteration are respected
  • Correct adapter count in each stage
  • Schedule output always parses cleanly through check_and_process_adapters() across a range of n_warm_up_iteration values
  • End-to-end sample_chain() integration test

New tests covering state carry-over across stages:

  • Four tests for scale adapter carry-over: both stochastic_approximation and dual_averaging variants are checked for (a) reading the current proposal scale on re-initialisation and (b) explicit initial_scale still taking priority over the proposal's current value.
  • Four tests for variance_shape_adapter: carry-over from a vector proposal shape; explicit initial_shape override; NULL proposal falls back to unit variances; matrix-valued proposal shape (incompatible type) falls back to unit variances.
  • Four tests for covariance_shape_adapter: carry-over from a matrix proposal shape; explicit initial_shape override; NULL proposal falls back to identity; vector-valued proposal shape (incompatible type) falls back to identity.
  • One end-to-end integration test via sample_chain with trace_warm_up=TRUE: verifies that the log_scale in warm_up_statistics does not jump back to the default value at the stage boundary between two consecutive scale-adapter-only stages.

Usage examples

Fully custom staged schedule:

results <- sample_chain(
  target_distribution,
  initial_state = rnorm(2),
  n_warm_up_iteration = 1000,
  n_main_iteration = 1000,
  adapters = list(
    list(scale_adapter(), 50),                              # scale only
    list(scale_adapter(), shape_adapter("variance"), 50),   # scale + diagonal shape
    list(scale_adapter(), shape_adapter("covariance"))      # scale + dense shape (remainder)
  )
)

Using the convenience constructor with defaults:

results <- sample_chain(
  target_distribution,
  initial_state = rnorm(2),
  n_warm_up_iteration = 1000,
  n_main_iteration = 1000,
  adapters = progressive_adaptation_schedule()
)

Existing flat-list usage is unchanged:

results <- sample_chain(
  target_distribution,
  initial_state = rnorm(2),
  n_warm_up_iteration = 1000,
  n_main_iteration = 1000,
  adapters = list(scale_adapter(), shape_adapter())  # still works as before
)

Known limitation

When trace_warm_up = TRUE and stages use different adapter sets (e.g. scale-only in stage 1, scale + shape in stage 2), the warm-up statistics matrices from each stage have different column counts and cannot be rbind-ed. This causes an error. A future PR could address this by filling NA for missing columns. For now, trace_warm_up = TRUE with staged adapters is only safe when all stages use the same adapter set. This limitation is documented in the new tests.

matt-graham and others added 6 commits June 15, 2026 14:03
…sts, with minor changes of indentation to improve readability.
…itial_shape parameter

Previously all adapter initialize() functions ignored the current proposal
parameters on re-initialisation, causing the scale and shape estimates
accumulated during one staged warm-up stage to be discarded at the start
of the next. This contradicted the requirement in issue UCL#79 that adaptation
parameters be initialised at the values at the end of the previous stage.

Changes in R/adaptation.R:

- stochastic_approximation_scale_adapter: when initial_scale is not
  user-specified, initialize() now reads proposal$parameters()$scale and
  uses it as the starting log-scale if available, falling back to
  proposal$default_initial_scale() only when no current scale exists.

- dual_averaging_scale_adapter: same fix. Note that mu, smoothed_log_scale
  and accept_prob_error already persisted across stages via closure scope
  and required no change.

- variance_shape_adapter: initialize() now follows a three-priority order:
  (1) explicit initial_shape constructor argument if supplied,
  (2) current proposal shape if it is a vector of matching length,
  (3) unit variances as before. The length check prevents accidentally
  carrying over a matrix-valued shape from a covariance adapter stage.
  Adds initial_shape parameter (numeric vector of per-dimension scales)
  to the function signature and roxygen documentation, addressing the
  request to allow users to supply domain-knowledge-based starting values.

- covariance_shape_adapter: same three-priority fix. Carries over the
  current proposal Cholesky factor if it is a matrix of matching dimensions,
  falls back to the identity otherwise. The is.matrix() guard prevents
  carrying over a vector-valued shape from a variance adapter stage.
  Adds initial_shape parameter (lower-triangular Cholesky factor matrix)
  to the function signature and roxygen documentation.

Changes in tests/testthat/test-adaptation.R:

- Four tests for scale adapter carry-over: both stochastic_approximation
  and dual_averaging variants are checked for (a) reading the current
  proposal scale on re-initialisation and (b) explicit initial_scale
  still taking priority over the proposal's current value.

- Four tests for variance_shape_adapter: carry-over from a vector proposal
  shape; explicit initial_shape override; NULL proposal falls back to unit
  variances; matrix-valued proposal shape (incompatible type) falls back
  to unit variances.

- Four tests for covariance_shape_adapter: carry-over from a matrix proposal
  shape; explicit initial_shape override; NULL proposal falls back to
  identity; vector-valued proposal shape (incompatible type) falls back
  to identity.

- One end-to-end integration test via sample_chain with trace_warm_up=TRUE:
  verifies that the log_scale in warm_up_statistics does not jump back to
  the default value at the stage boundary between two consecutive
  scale-adapter-only stages.
@codecov

codecov Bot commented Jun 25, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 100.00%. Comparing base (65b86ab) to head (a700117).

Additional details and impacted files
@@            Coverage Diff             @@
##              main      #127    +/-   ##
==========================================
  Coverage   100.00%   100.00%            
==========================================
  Files           11        11            
  Lines          620       722   +102     
==========================================
+ Hits           620       722   +102     

☔ View full report in Codecov by Harness.
📢 Have feedback on the report? Share it here.

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@matt-graham matt-graham left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is looke great @jennajiali - I've added some suggested changes but mainly minor things. Will also need to resolve the linter warnings

Comment thread R/chains.R
Comment thread R/chains.R Outdated
Comment thread R/adaptation.R Outdated
Comment thread R/adaptation.R Outdated
Comment thread R/adaptation.R Outdated
jennajiali and others added 5 commits July 9, 2026 18:27
Co-authored-by: Matt Graham <matthew.m.graham@gmail.com>
warm_up_results_1$traces and warm_up_results_1$statistics are not guaranteed to be NULL its just that may be

Co-authored-by: Matt Graham <matthew.m.graham@gmail.com>
This will make code a bit more self-documenting and should solve linter line length warning

Co-authored-by: Matt Graham <matthew.m.graham@gmail.com>
Factoring out condition here into variable again makes code a bit more readable and may help with line length warning (though haven't counted to check!)

Co-authored-by: Matt Graham <matthew.m.graham@gmail.com>
@matt-graham matt-graham added the enhancement New feature or request label Jul 9, 2026
jennajiali and others added 5 commits July 10, 2026 21:26
@jennajiali
jennajiali requested a review from matt-graham July 13, 2026 20:46
@matt-graham matt-graham linked an issue Jul 16, 2026 that may be closed by this pull request

@matt-graham matt-graham left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks @jennajiali for addressing previous comments. I've added a few more small suggestions to resolve the linter and test coverage fails in the checks. Let me know if anything doesn't make sense.

Comment thread R/chains.R Outdated
Comment thread R/chains.R Outdated
Comment thread R/chains.R
jennajiali and others added 4 commits July 19, 2026 17:10
Split check_and_process_adapters into two functions to reduce
cyclomatic complexity (satisfying cyclocomp_linter). A new private
helper `process_stage_specifications` owns the per-stage-spec walk:
peeling off the optional trailing iteration count, validating stage
adapter elements, allocating the remainder to a trailing no-count
stage, and checking the total sums to n_warm_up_iteration.
`check_and_process_adapters` is now a flat four-branch dispatch that
normalises each of the three accepted input forms into a single
canonical stage-spec-list shape and delegates the rest of the work.

A brief note on the shape choice: `process_stage_specifications` is
written to consume *user-input* stage specs (each a list of adapters
with an optional trailing integer), not the already-normalised
`list(adapters = ..., n_iteration = ...)` output shape. To keep the
helper focused on one input shape, Form 1 (a flat list of adapters)
is wrapped as `list(c(adapters, list(n_warm_up_iteration)))` — a
length-1 list-of-stages whose single stage spec is all the adapters
followed by a trailing count. `c(adapters, list(n_warm_up_iteration))`
is used rather than `c(adapters, n_warm_up_iteration)` so that the
scalar is always spliced in as exactly one element regardless of its
length or type.

Form 3 (schedule constructor function) now hands the function's
return value straight to `process_stage_specifications` rather than
recursing back into `check_and_process_adapters`. This aligns with
the documented contract that a schedule function returns Form 2
(which `progressive_adaptation_schedule` respects) and keeps the
three top-level branches symmetric.
Address review comment and also test the behaviour of
empty-adapter stages suggested by Sam.

check_and_process_adapters unit tests:

* "numeric scalar input raises error" — passes a numeric scalar so
  the top-level "adapters invalid" branch is exercised in a way that
  complements the existing string-input test.
* "staged list with per-stage counts exceeding n_warm_up_iteration
  and no trailing count raises informative error" — passes a schedule
  whose first two stages already over-consume n_warm_up_iteration and
  whose trailing stage omits an iteration count, verifying the
  "Per-stage iteration counts exceeds n_warm_up_iteration" error is
  raised at parse time rather than surfacing later as a cryptic
  sampling-time failure.
* "staged list with an empty middle stage (no adapters, just an
  iteration count) is valid" — confirms that a stage containing only
  a trailing iteration count and no adapters (a "pause" during warm-up
  where the chain moves but no proposal parameters update) parses
  correctly and produces a stage with zero adapters and the requested
  n_iteration.
* "empty top-level list is treated as a single no-adapter stage
  covering all warm-up iterations" — confirms that an entirely empty
  `adapters` argument produces a single zero-adapter stage spanning
  the whole warm-up, i.e. "no adaptation at all". Structurally
  different from the empty-middle-stage case above (Form 1 rather
  than Form 2), so worth locking down separately.

sample_chain integration test:

* "sample_chain with an empty middle stage (no adapters, just an
  iteration count) runs without error" — integration-level check that
  the whole pipeline (parse -> sample -> combine) handles a stage with
  zero adapters. Exercises the `for (adapter in adapters)` no-op path
  in initialize/update/finalize_adapters and the empty-adapter-set
  path in initialize_statistics.

@matt-graham matt-graham left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks @jennajiali for making the additional changes. Factoring out the process_stage_specifications function I think has helped improve the readability a lot - well done for also catching and fixing the unintentional bug in my suggested changes in the handling of the form 1 case 😅. The additional tests for edge cases are also great. I merged in the changes from #130 and made one small fix to a use of = rather than <- that was causing styler check in pre-commit to fail. This all looks good to go to me now so will merge in.

)

test_that(
"check_and_process_adapters: empty top-level list is treated as a single no-adapter stage covering all warm-up iterations",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Great idea to have this as an additional test case!

# ── sample_chain with staged adapters integration tests ───────────────────────

test_that(
"sample_chain with an empty middle stage (no adapters, just an iteration count) runs without error",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is also a great test to have - it's a testament to the generality of your implementation that this works as expected!

@matt-graham
matt-graham merged commit f399737 into UCL:main Jul 22, 2026
9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Allow user to define a schedule for adaptation during burn-in

2 participants