Skip to content

Prediction: Forward Simulation

Forward-simulate trained predictors against a dataset. predict_bucket is the single-bucket primitive (JIT-compiled, vmapped over the bucket's N axis); predict_dataset walks every bucket and returns one [N, T, D] array per bucket.


predict_bucket()

from jaxhybridmodels.prediction import predict_bucket  ·  also re-exported as jaxhybridmodels.predict_bucket

python
predict_bucket(
    predictors: 'Any',
    bp: 'BucketPayload',
    simulate_fn: 'Callable[..., Array]',
    state_to_output: 'Callable[[Array], Array]',
    solver: 'SolverConfig',
) -> Float[Array, 'N T D']

Vmap simulate_fn over the bucket's N axis and project to observed channels.

predict_bucket_obs runs simulate_fn for one experiment, giving a state trajectory [T, S], and maps it to observed channels [T, D]. jax.vmap lifts that over (ts, covariates, y0) along N. predictors and solver are closed over with no vmap axis, being the same for every experiment in the bucket.

One compiled kernel per bucket shape. Python dispatch over buckets lives in predict_dataset, never inside the compiled region.

Parameters

ParameterTypeDescription
predictorsPyTree[eqx.Module]The trainable part of the model, typically a tuple of BoundedPredictor leaves. Forwarded to simulate_fn unchanged.
bpBucketPayloadOne bucket. Its ts, covariates, and y0 are vmapped along N.
simulate_fnUser-supplied integrator with signature (predictors, ts, covariates, y0, solver) -> [T, S].
state_to_outputPure [T, S] -> [T, D] map from full state to observed channels. A property of the model, passed explicitly.
solverSolverConfig. All its fields are static, so it enters the compiled kernel as configuration rather than as data.

Returns

ItemTypeDescription
Float[Array, "N T D"]Predicted output channels for every experiment in the bucket.

predict_dataset()

from jaxhybridmodels.prediction import predict_dataset  ·  also re-exported as jaxhybridmodels.predict_dataset

python
predict_dataset(
    predictors: 'Any',
    dataset: 'Dataset',
    simulate_fn: 'Callable[..., Array]',
    state_to_output: 'Callable[[Array], Array]',
    solver: 'SolverConfig',
) -> tuple[Float[Array, 'N T D'], ...]

Run predict_bucket over every bucket in dataset and return the stack tuple.

Each bucket shape compiles predict_bucket exactly once. The result is a tuple rather than one array, because buckets differ precisely in T and cannot be stacked.

Parameters

ParameterTypeDescription
state_to_outputPure mapping [T, S] -> [T, D] from full simulator state to the observed channels. A property of the model, passed here rather than stored on the Dataset.

Returns

ItemTypeDescription
tuple[Float[Array, "N T D"], ...]One [N_b, T_b, D] array per bucket, in bucket-payload order.

Source


predict_dense()

from jaxhybridmodels.prediction import predict_dense  ·  also re-exported as jaxhybridmodels.predict_dense

python
predict_dense(
    predictors: 'Any',
    dataset: 'Dataset',
    simulate_fn: 'Callable[..., Array]',
    state_to_output: 'Callable[[Array], Array]',
    solver: 'SolverConfig',
    ts_grid: 'Array | None' = None,
    n_points: 'int' = 100,
) -> tuple[Float[Array, 'N T_d D'], ...]

Evaluate a trained model on a dense time grid, one array per bucket.

predict_dataset returns predictions only at the measured timestamps. For smooth trajectory plots or dense evaluation you usually want more points than that. This builds a fine grid per experiment and reuses the same compiled forward pass, so no new kernel or dependency is needed (the diffraxtra VectorizedDenseInterpolation equivalent, folded in ~20 lines).

Parameters

ParameterTypeDescription
ts_gridOptional shared grid [T_d] to evaluate every experiment on. If None, each experiment gets its own linspace from its first to its last measured time with n_points points.
n_pointsPoints per experiment when ts_grid is None. Ignored otherwise.

Returns

ItemTypeDescription
tuple[Float[Array, "N T_d D"], ...]One [N, T_d, D] array per bucket, in bucket-payload order.

Source


ensemble_predictions()

from jaxhybridmodels.prediction import ensemble_predictions  ·  also re-exported as jaxhybridmodels.ensemble_predictions

python
ensemble_predictions(
    members: 'Sequence[Any]',
    dataset: 'Dataset',
    simulate_fn: 'Callable[..., Array]',
    state_to_output: 'Callable[[Array], Array]',
    solver: 'SolverConfig',
) -> tuple[Float[Array, 'N T D'], ...]

Average per-bucket predictions across an ensemble of models.

members is a sequence of predictor pytrees — each the predictors argument you would pass to predict_dataset alone. Each member is run forward and the per-bucket predictions are averaged, so the result is the same shape as a single predict_dataset return.

Parameters

ParameterTypeDescription
membersNon-empty sequence of predictor pytrees. Every member must be compatible with the same simulate_fn.

Returns

ItemTypeDescription
tuple[Float[Array, "N T D"], ...]The member-mean prediction per bucket, in bucket-payload order.

Source


evaluate_predictor()

from jaxhybridmodels.prediction import evaluate_predictor  ·  also re-exported as jaxhybridmodels.evaluate_predictor

python
evaluate_predictor(predictor: 'Any', covariates: 'dict[str, float]') -> 'float'

Evaluate a scalar-valued predictor at named inputs, as a Python float.

Shortcut for the recovered-physics readout every example writes by hand (float(predictor({"k": jnp.asarray(v)}).reshape(()))). Takes the predictor's named inputs as plain Python floats, calls it, and flattens the scalar result to a float.

Source

Released under the BSD-3-Clause License.