Inference API¶
The pyrox.inference subpackage exposes ensemble-of-MAP and ensemble-of-VI runners as a layered surface — pick the level of control that fits your use case.
Layer 1 — Functional primitives¶
Roll your own training loop on top of the vmapped state primitives.
ensemble_init(init_fn: Callable[[PRNGKeyArray], PyTree], optimizer: optax.GradientTransformation, *, ensemble_size: int, seed: PRNGKeyArray) -> EnsembleState
¶
Initialize an ensemble of (params, opt_state) by vmap over keys.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
init_fn
|
Callable[[PRNGKeyArray], PyTree]
|
|
required |
optimizer
|
GradientTransformation
|
|
required |
ensemble_size
|
int
|
Number of independent ensemble members |
required |
seed
|
PRNGKeyArray
|
PRNG key, split into |
required |
Returns:
| Type | Description |
|---|---|
EnsembleState
|
|
EnsembleState
|
Array leaves carry a leading |
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
ensemble_loss(log_joint: Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]], *, prior_weight: float = 1.0, scale: float = 1.0) -> Callable[[PyTree, Array, Array], tuple[Float[Array, ''], PyTree]]
¶
Build the filter_value_and_grad loss from a log-joint.
Returns a function loss_fn(params, x_batch, y_batch) -> (loss, grads)
that computes
prior_weight=0 short-circuits the prior term so the user may
return a placeholder 0.0 for logprior.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
log_joint
|
Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]]
|
|
required |
prior_weight
|
float
|
Weight on the |
1.0
|
scale
|
float
|
Multiplicative weight on the |
1.0
|
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
ensemble_step(state: EnsembleState, x_batch: Array, y_batch: Array, *, log_joint: Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]], optimizer: optax.GradientTransformation, prior_weight: float = 1.0, scale: float = 1.0) -> tuple[EnsembleState, Float[Array, ' E']]
¶
Perform one ensemble update step on a batch.
Each ensemble member computes ∇L(θ_e) independently (vmapped),
advances its optax state, and returns the updated params + state.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state
|
EnsembleState
|
Current |
required |
x_batch
|
Array
|
Inputs. |
required |
y_batch
|
Array
|
Targets. |
required |
log_joint
|
Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]]
|
|
required |
optimizer
|
GradientTransformation
|
Same |
required |
prior_weight
|
float
|
Weight on |
1.0
|
scale
|
float
|
Weight on |
1.0
|
Returns:
| Type | Description |
|---|---|
EnsembleState
|
|
Float[Array, ' E']
|
|
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
Layer 2 — NumPyro-like inference ops¶
init / update / run triplets that mirror numpyro.infer.SVI.
EnsembleMAP
¶
Bases: Module
NumPyro-like ensemble MAP/MLE runner.
Mirrors numpyro.infer.SVI's init / update / run
triplet, but every operation is ensembled by vmap over the
leading (E,) axis.
Per-member objective is the tempered negative log-posterior
where \(w_{\text{prior}}\) is prior_weight.
Tempering. Down-weighting the prior has the same argmax as tempering the likelihood by \(\beta = 1 / w_{\text{prior}}\), since
and rescaling an objective does not move its argmax. The two are
not interchangeable in general: prior_weight scales the
entire log prior — hyperpriors included — while a likelihood-side
\(\beta\) (e.g. lfr_factor, which applies
numpyro.handlers.scale to the likelihood site only) leaves every
prior at unit weight. When the two coincide on which terms they
reweight, their gradients differ by the uniform factor \(\beta\) —
same argmax, but a scale an adaptive optimizer still responds to.
This holds per minibatch as well: the tempered gradient is exactly
\(\beta\) times the prior-weighted one, so minibatching introduces no
further difference between the formulations.
Ensembling over seeds is worth more than usual for multi-modal
objectives such as pyrox_gp.LatentFactorGPPrior's. Note what the
spread does and does not mean: where the latent priors are identical
the objective is rotation-invariant, and seed-to-seed spread along
that gauge is an unidentifiable coordinate artifact, not
uncertainty — averaging raw parameters across members is meaningless
there. Ensemble gauge-invariant predictions instead, or align the
members first. (With distinct latent kernels the rotation changes the
objective, so no such flat manifold exists.)
Examples:
>>> runner = EnsembleMAP(
... log_joint=log_joint,
... init_fn=init_fn,
... optimizer=optax.adam(5e-3),
... ensemble_size=16,
... )
>>> # numpyro-style three-method API
>>> state = runner.init(jr.PRNGKey(0))
>>> for _ in range(2000):
... state, losses = runner.update(state, x, y)
>>> # or one-shot
>>> result = runner.run(jr.PRNGKey(0), 2000, x, y)
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 | |
init(seed: PRNGKeyArray) -> EnsembleState
¶
Initialize the ensemble. Mirrors numpyro.infer.SVI.init.
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
update(state: EnsembleState, x_batch: Array, y_batch: Array, *, scale: float = 1.0) -> tuple[EnsembleState, Float[Array, ' E']]
¶
One ensemble update step. Mirrors numpyro.infer.SVI.update.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
state
|
EnsembleState
|
Current |
required |
x_batch
|
Array
|
Batch inputs. |
required |
y_batch
|
Array
|
Batch targets. |
required |
scale
|
float
|
|
1.0
|
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
run(seed: PRNGKeyArray, num_epochs: int, x: Array, y: Array, *, batch_size: int | None = None) -> EnsembleResult
¶
Fit the ensemble end-to-end. Mirrors numpyro.infer.SVI.run.
Internally drives ensemble_step via lax.scan for
speed; equivalent to a hand-written Python loop over
update.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
seed
|
PRNGKeyArray
|
PRNG key used for both init and (when applicable) mini-batch index permutation. |
required |
num_epochs
|
int
|
Number of optimizer steps per member. |
required |
x
|
Array
|
Inputs. |
required |
y
|
Array
|
Targets. |
required |
batch_size
|
int | None
|
Optional mini-batch size. |
None
|
Returns:
| Type | Description |
|---|---|
EnsembleResult
|
|
EnsembleResult
|
history of shape |
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
EnsembleVI
¶
Bases: Module
NumPyro-like ensemble variational-inference runner.
Wraps numpyro.infer.SVI + numpyro.infer.Trace_ELBO
with the same ensemble surface as EnsembleMAP.
Per-member objective is the tempered ELBO
where \(\beta\) is kl_weight.
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 | |
init(seed: PRNGKeyArray, *args: Any, **kwargs: Any) -> Any
¶
Initialize the ensemble of SVI states.
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
update(state: Any, *args: Any, **kwargs: Any) -> tuple[Any, Float[Array, ' E']]
¶
One ensemble SVI update. Mirrors numpyro.infer.SVI.update.
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
run(seed: PRNGKeyArray, num_epochs: int, *args: Any, **kwargs: Any) -> EnsembleResult
¶
Fit the ensemble end-to-end via vmapped svi.run.
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
Layer 3 — One-shot sugar¶
ensemble_map(log_joint: Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]], init_fn: Callable[[PRNGKeyArray], PyTree], *, ensemble_size: int, num_epochs: int, data: tuple[Array, Array], seed: PRNGKeyArray, batch_size: int | None = None, learning_rate: float = 0.005, prior_weight: float = 1.0, optimizer: optax.GradientTransformation | None = None) -> tuple[PyTree, Float[Array, 'E T']]
¶
One-shot wrapper around EnsembleMAP.
Equivalent to EnsembleMAP(log_joint, init_fn, optimizer,
ensemble_size=E, prior_weight=w).run(seed, num_epochs, *data,
batch_size=B).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
log_joint
|
Callable[[PyTree, Array, Array], tuple[Float[Array, ''], Float[Array, '']]]
|
|
required |
init_fn
|
Callable[[PRNGKeyArray], PyTree]
|
|
required |
ensemble_size
|
int
|
Number of independent MAP fits |
required |
num_epochs
|
int
|
Optimizer steps per member. |
required |
data
|
tuple[Array, Array]
|
|
required |
seed
|
PRNGKeyArray
|
PRNG key. |
required |
batch_size
|
int | None
|
Mini-batch size. |
None
|
learning_rate
|
float
|
Default-Adam learning rate. Ignored if
|
0.005
|
prior_weight
|
float
|
|
1.0
|
optimizer
|
GradientTransformation | None
|
Optional |
None
|
Returns:
| Type | Description |
|---|---|
PyTree
|
|
Float[Array, 'E T']
|
|
Examples:
>>> params, losses = ensemble_map(
... log_joint, init_fn,
... ensemble_size=16, num_epochs=2000,
... data=(X, y), seed=jr.PRNGKey(0),
... )
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
ensemble_vi(model_fn: Callable[..., None], guide_fn: Callable[..., None], *, ensemble_size: int, num_epochs: int, data: tuple[Array, Array], seed: PRNGKeyArray, kl_weight: float = 1.0, learning_rate: float = 0.005, optimizer: Any = None, num_particles: int = 1) -> tuple[PyTree, Float[Array, 'E T']]
¶
One-shot wrapper around EnsembleVI.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model_fn
|
Callable[..., None]
|
NumPyro model |
required |
guide_fn
|
Callable[..., None]
|
NumPyro guide. |
required |
ensemble_size
|
int
|
Number of SVI fits |
required |
num_epochs
|
int
|
Steps per member. |
required |
data
|
tuple[Array, Array]
|
|
required |
seed
|
PRNGKeyArray
|
PRNG key. |
required |
kl_weight
|
float
|
ELBO temper |
1.0
|
learning_rate
|
float
|
Default-Adam learning rate. |
0.005
|
optimizer
|
Any
|
Optional |
None
|
num_particles
|
int
|
MC particles per ELBO estimate. |
1
|
Returns:
| Type | Description |
|---|---|
tuple[PyTree, Float[Array, 'E T']]
|
|
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
ensemble_predict(params_stacked: PyTree, predict_fn: Callable[[PyTree, Array], Array], x_new: Array) -> Array
¶
Vmap predict_fn over the leading ensemble axis of params.
Uses equinox.filter_vmap so it works whether
params_stacked is a pure-array PyTree or an
equinox.Module containing non-array leaves (e.g. captured
jax.nn.tanh). Array leaves are mapped over axis 0; non-array
leaves are broadcast.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
params_stacked
|
PyTree
|
PyTree returned by |
required |
predict_fn
|
Callable[[PyTree, Array], Array]
|
|
required |
x_new
|
Array
|
Inputs to predict at; shared across all members. |
required |
Returns:
| Type | Description |
|---|---|
Array
|
Stacked predictions with leading |
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
Optimizer helpers¶
param_group_optimizer(groups: dict[str, optax.GradientTransformation], label_fn: Callable[[tuple, Any], str]) -> optax.GradientTransformation
¶
Apply a different optimizer to each labelled parameter group.
Models that mix a large block of free latent parameters with a few kernel hyperparameters usually want different step sizes for each: the latents sit on a well-conditioned objective and tolerate a large step, while lengthscales and noise live on a log scale and destabilize under one. A 10x ratio is a common starting point.
A label_fn returning a label that is not a key of groups fails
loudly at init time — optax.multi_transform raises a
ValueError naming the offending labels.
Label by path, never by leaf value
optax.multi_transform evaluates a callable param_labels on
the parameter tree at init and on the update tree at
every update. A label_fn that inspects leaf values (a
sign test, say) can therefore assign one group at init and another
afterwards, silently applying the wrong transform or tripping a
masked-state structure error. Depend only on path and on
update-invariant leaf metadata such as shape / dtype, which
are identical for parameters and their updates.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
groups
|
dict[str, GradientTransformation]
|
Maps a group label to the optimizer for that group. Every
label returned by |
required |
label_fn
|
Callable[[tuple, Any], str]
|
Called as |
required |
Returns:
| Type | Description |
|---|---|
GradientTransformation
|
A single |
GradientTransformation
|
|
Examples:
>>> import optax
>>> def by_name(path, _):
... return "latents" if "Z_T" in str(path) else "globals"
>>> opt = param_group_optimizer(
... {"latents": optax.adam(1e-2), "globals": optax.adam(1e-3)},
... by_name,
... )
Source code in packages/pyrox/src/pyrox/inference/_param_groups.py
Result containers¶
EnsembleState
¶
Bases: NamedTuple
Stacked state for an ensemble of optimizer runs.
Attributes:
| Name | Type | Description |
|---|---|---|
params |
PyTree
|
Stacked parameter PyTree. Array leaves carry a leading
|
opt_state |
PyTree
|
Stacked optax optimizer state with the same axis convention. |
Source code in packages/pyrox/src/pyrox/inference/_ensemble.py
EnsembleResult
¶
Bases: NamedTuple
Output of a full EnsembleMAP.run / EnsembleVI.run.
Attributes:
| Name | Type | Description |
|---|---|---|
params |
PyTree
|
Final stacked parameters with leading |
losses |
Float[Array, 'E T']
|
|