pypomp.functional.align_params

pypomp.functional.align_params(params: Mapping[str, float | Array], names: list[str], axis: int = -1) Array[source]

Align and stack parameter arrays into the canonical ordering for a model struct.

Builds a single JAX array from a dictionary of named parameter values, reordering them to match the canonical param_names order expected by pypomp.functional.pfilter(), mif(), and train().

Parameters:
  • params (mapping of str to jax.Array or float) – Dictionary mapping parameter names to JAX arrays or float scalars.

  • names (list of str) – Canonical parameter name ordering (e.g. struct.param_names).

  • axis (int, optional) – Axis along which to stack. Defaults to -1 (last axis).

Returns:

Array whose last axis (by default) corresponds to names in order.

Return type:

jax.Array

Examples

>>> import jax.numpy as jnp
>>> import pypomp.functional as F
>>> params = {"beta": jnp.array(0.5), "gamma": jnp.array(0.1)}
>>> F.align_params(params, names=["gamma", "beta"])
Array([0.1, 0.5], dtype=float32)