Phantom collection and evidence conditioning

Phantom states are eligible intermediate transitions from a constrained slice chain. They condition the Monte Carlo shrinkage model, but they are not classic race-tree samples and do not contribute posterior coordinates or posterior effective sample size.

For the statistical derivation and experimental limitations of phantom conditioning, read the JAXNS v3 paper, Phantom-Conditioned Nested Sampling.

Collection owns memory

Enable collection on jaxns.core.NestedSampler:

nested_sampler = NestedSampler(
    model=model,
    collect_phantom_samples=True,
)
state = nested_sampler.run(key=jax.random.PRNGKey(0))
results = state.to_result().trim()

The runner stores every intermediate transition from each generated chain. The final transition remains the classic replacement and is never stored as a phantom. A chain with num_slices=s*D therefore retains s*D - 1 states. Root prior draws do not produce phantoms. Keeping these likelihoods increases state and checkpoint memory in proportion to the retained count.

The runner owns collection even when an explicit UniDimSliceSampler is supplied. It requests either all intermediate states or zero, according to collect_phantom_samples. For direct low-level calls, the sampler accepts num_phantom_samples and validates that it lies between zero and num_slices - 1. A shorter evidence-time prefix remains an independent choice.

Continuation requires the saved phantom width to match this policy. Older runs that retained only a shorter prefix remain readable for analysis, but should be continued with the code and collection settings that produced them. Missing intermediate states cannot be reconstructed retrospectively.

Conditioning owns computation

The completed state or results can reuse any leading part of the retained prefix:

all_saved = results.sample_evidence(
    num_samples=4096,
    phantom_conditioning=True,
    num_phantoms=None,
    key=jax.random.PRNGKey(1),
)
first_four = results.sample_evidence(
    num_samples=4096,
    phantom_conditioning=True,
    num_phantoms=4,
    key=jax.random.PRNGKey(1),
)
classic = results.sample_evidence(
    num_samples=4096,
    phantom_conditioning=False,
    key=jax.random.PRNGKey(1),
)

None uses all saved states. An explicit positive count uses log_L_phantom[:, :num_phantoms] before the MC kernel is compiled, so an unused suffix does not add device work. Classic conditioning is the default, independent of phantom storage. Phantom conditioning requires phantom_conditioning=True.

state.expected_log_Z_mean and state.expected_log_Z_uncert provide the classic expectation calculation for goal conditions. Results carry those estimates in log_Z_mean and log_Z_uncert. Calling sample_evidence leaves them unchanged and returns a separate ensemble with its own Monte Carlo mean and uncertainty.