Differentiable simulation (API)¶
User-friendly object-oriented API for differentiable simulation unsim.api.
A DifferentiableWorld is obtained by World.compile(backend="jax").
DifferentiableWorld¶
- class unsim.api.DifferentiableWorld(W, backend='jax', route_update_interval=None, toll_interval=None)¶
Immutable compiled snapshot of a World for differentiable simulation.
Create via
World.compile(backend="jax"). Holds the JAX parameter and configuration arrays plus the semantic index from Link/Node objects and names to integer indices.- Parameters:
W (World) – Source world. It is finalized if not already.
backend (str, optional) – Only “jax” is supported.
route_update_interval (float or None, optional) – DUO route update interval (s). None uses the built-in default (300 s).
toll_interval (float or None, optional) – Toll discretization interval (s), independent of the route update interval. None couples it to the route update interval (current default behavior).
- minimize(variables, objective, checkpoint_every=None)¶
Build an optimization problem over selected variables.
- Parameters:
variables (list of TollVariable, CustomVariable, or ParameterRef) – Decision variables.
objective (callable) – Function
(result, x) -> JAX scalarwherexmaps each variable to its current value array.checkpoint_every (int or None, optional) – Gradient checkpointing segment length.
- Return type:
- objective(fn, checkpoint_every=None)¶
Build an Objective from a function of a SimResult.
- Parameters:
fn (callable) – Maps a SimResult to a JAX scalar.
checkpoint_every (int or None, optional) – Gradient checkpointing segment length.
- Return type:
- parameter(entity, field)¶
Get a ParameterRef for an entity’s physical quantity.
- Parameters:
- Return type:
- Raises:
ValueError – If the quantity is derived under the entity’s FD parameterization, or the field is unknown.
- parameter_fields(entity)¶
List the field names accepted by
parameter()for an entity.For a link, only the FD quantities that are independent under the link’s FD parameterization are included.
- path(links)¶
Build a PathRef from an ordered list of links (objects or names).
- property raw¶
raw.params,raw.config,raw.simulate.- Type:
Low-level access to the JAX core
- run(differentiable=True, checkpoint_every=None)¶
Run the simulation with the compiled parameters.
- Parameters:
differentiable (bool, optional) – If True (default), use the AD-compatible path. If False, use the faster forward-only path (not compatible with reverse-mode AD).
checkpoint_every (int or None, optional) – Gradient checkpointing segment length (see core simulators).
- Return type:
- toll_variable(links, interval=None, initial=0.0, lower=None, upper=None)¶
Build a toll optimization variable over selected links and all toll steps.
- Parameters:
links (list of Link or str) – Links to toll.
interval (float or None, optional) – Expected toll discretization interval (s). Must match the compiled toll interval (set via
compile(toll_interval=...)); a mismatch raises an error.initial (float, optional) – Initial toll value (s). Default 0.
lower (float or None, optional) – Lower bound (s), applied by projection during optimization.
upper (float or None, optional) – Upper bound (s).
- Return type:
- variable(shape, initial, inject, name='custom', lower=None, upper=None)¶
Build a user-defined variable with a custom injection into Params.
Use this for composite variables that the built-in refs cannot express, e.g. a difference between two parameters, a factor shared across links, or a log-transformed parameter. The inject function must be a pure JAX-differentiable function; FD consistency across derived quantities is the caller’s responsibility.
- Parameters:
shape (tuple of int) – Shape of the variable value.
initial (float or array_like) – Initial value, broadcast to
shape.inject (callable) – Function
(params, theta) -> Paramswherethetahas shapeshape.name (str, optional) – Display name.
lower (float or None, optional) – Lower bound, applied by projection during optimization.
upper (float or None, optional) – Upper bound.
- Return type:
SimResult¶
- class unsim.api.SimResult(state, params, model)¶
Simulation result facade wrapping the raw SimState.
Provides differentiable queries (metrics, link, node, path, od) and host-side analysis (analyzer).
- state¶
Raw JAX simulation state.
- Type:
SimState
- params¶
Parameters used for this run.
- Type:
Params
- config¶
Static network configuration.
- Type:
NetworkConfig
- property analyzer¶
Host-side Analyzer with this result written back into the snapshot World.
Each access rewrites the snapshot World with this result’s arrays, so the returned Analyzer reflects this result at access time. The write-back covers aggregate cumulative curves, origin queues, and absorbed counts; per-destination curves are not restored.
- link(link)¶
Get a LinkView for a link (object, name, or index).
- node(node)¶
Get a NodeView for a node (object, name, or index).
- od(orig, dest)¶
Get an ODView for an origin and destination node.
- path(path)¶
Get a PathView for a PathRef or a list of links.
Metrics¶
- class unsim.api.Metrics(result)¶
Differentiable aggregate traffic metrics of a simulation result.
All methods return JAX scalars and are pure functions usable inside jax.jit and jax.grad.
- average_travel_time()¶
Average travel time per completed trip (s).
- completed_trips()¶
Total completed trips (veh).
- total_travel_time()¶
Total travel time (s).
LinkView¶
- class unsim.api.LinkView(result, link_id, link_name)¶
Differentiable per-link queries on a simulation result.
- cum_arrival()¶
Cumulative arrival curve (veh), shape (tsize+1,).
- cum_departure()¶
Cumulative departure curve (veh), shape (tsize+1,).
- density(t)¶
Average density on the link at time t (veh/m). Differentiable.
- Parameters:
t (float) – Time (s).
- vehicle_count(t)¶
Number of vehicles on the link at time t (veh). Differentiable.
- Parameters:
t (float) – Time (s).
NodeView¶
- class unsim.api.NodeView(result, node_id, node_name)¶
Per-node queries on a simulation result.
- queue(t)¶
Vertical queue length at an origin node at time t (veh). Differentiable in value.
The timestep selection uses a floor index and is not differentiable with respect to t.
- Parameters:
t (float) – Time (s).
PathView¶
ODView¶
- class unsim.api.ODView(result, orig_id, dest_id)¶
Differentiable OD-level queries on a simulation result.
- travel_time(departure_time, method='soft', temperature=None)¶
OD travel time (s).
- Parameters:
departure_time (float) – Departure time (s).
method (str, optional) – “soft” for fully differentiable soft route choice (default). “auto” for shortest-path chaining (route choice itself is not differentiated). “logsum” for the expected perceived cost under logit route choice.
temperature (float or None, optional) – Logit temperature (s) for “soft” and “logsum”. None uses the compiled default.
Objective¶
- class unsim.api.Objective(model, fn, checkpoint_every=None)¶
Objective function over simulation results with gradient support.
- Parameters:
model (DifferentiableWorld) – Compiled model.
fn (callable) – Function mapping a SimResult to a JAX scalar, e.g.
lambda R: R.metrics.total_travel_time().checkpoint_every (int or None, optional) – Gradient checkpointing segment length passed to the core simulator.
- explain(wrt=None)¶
Describe what the gradient computation differentiates and what it holds fixed.
- Parameters:
wrt (list of ParameterRef or None, optional) – Variables to describe. None describes only the objective and model.
- Returns:
Human-readable description.
- Return type:
str
- gradient(wrt, jit=False)¶
Evaluate the gradient only. See
value_and_gradient.
- value()¶
Evaluate the objective at the base parameters.
- Return type:
jnp scalar
- value_and_gradient(wrt, jit=False)¶
Evaluate the objective and its gradient with respect to selected parameters.
- Parameters:
wrt (list of ParameterRef) – Parameters to differentiate with respect to.
jit (bool, optional) – If True, jit-compile the value-and-gradient function.
- Returns:
value (jnp scalar)
gradient (Gradients) – Mapping from each ref to its gradient (same shape as the parameter).
Gradients¶
ParameterRef¶
- class unsim.api.ParameterRef(kind: str, field: str, index: int, name: str = '', unit: str = '', fd_parameterization: str = None, shape: tuple = ())¶
Semantic reference to one differentiable parameter.
Equality and hashing use (kind, field, index) only, so two refs to the same parameter compare equal.
- kind¶
Entity kind: “link” or “node”.
- Type:
str
- field¶
Physical quantity name (e.g. “free_flow_speed”).
- Type:
str
- index¶
Entity index in the compiled arrays.
- Type:
int
- name¶
Entity name (for display).
- Type:
str
- unit¶
Physical unit of the quantity.
- Type:
str
- fd_parameterization¶
FD parameterization of the link, for FD fields only.
- Type:
str or None
- shape¶
Shape of the parameter value; () for scalars.
- Type:
tuple
- property size¶
Number of scalar entries in this parameter.
PathRef¶
TollVariable¶
- class unsim.api.TollVariable(link_ids: tuple, n_steps: int, link_names: tuple = (), initial: float = 0.0, lower: float = None, upper: float = None)¶
Optimization variable block: tolls on selected links over all toll steps.
Created via
DifferentiableWorld.toll_variable().- link_ids¶
Selected link indices.
- Type:
tuple of int
- n_steps¶
Number of toll discretization steps.
- Type:
int
- link_names¶
Selected link names.
- Type:
tuple of str
- initial¶
Initial toll value (s).
- Type:
float
- lower¶
Lower bound (s), applied by projection during optimization.
- Type:
float or None
- upper¶
Upper bound (s).
- Type:
float or None
- property shape¶
Shape of the variable block.
- property size¶
Number of scalar entries.
CustomVariable¶
- class unsim.api.CustomVariable(shape, initial, inject, name='custom', lower=None, upper=None)¶
User-defined variable with a custom injection into Params.
Lets the user define composite variables such as differences, shared factors, or transformed parameters. The inject function maps (params, theta) to a new Params and must be a pure JAX-differentiable function; FD consistency across derived quantities is the user’s responsibility here. Instances compare by identity, so use the same object when building and when reading gradients.
Created via
DifferentiableWorld.variable().- Parameters:
shape (tuple of int) – Shape of the variable value.
initial (float or array_like) – Initial value, broadcast to
shape.inject (callable) – Function
(params, theta) -> Paramswherethetahas shapeshape.name (str, optional) – Display name.
lower (float or None, optional) – Lower bound, applied by projection during optimization.
upper (float or None, optional) – Upper bound.
- property size¶
Number of scalar entries.
Problem¶
- class unsim.api.Problem(model, variables, objective, checkpoint_every=None)¶
Optimization problem over selected variables of a compiled model.
Created via
DifferentiableWorld.minimize().- Parameters:
model (DifferentiableWorld)
variables (list of TollVariable, CustomVariable, or ParameterRef) – Decision variables.
objective (callable) – Function
(result, x) -> JAX scalarwherexmaps each variable to its current value array.checkpoint_every (int or None, optional) – Gradient checkpointing segment length.
- solve(optimizer='adam', steps=200, learning_rate=1.0, jit=True, b1=0.9, b2=0.999, eps=1e-08, verbose=False)¶
Minimize the objective with a first-order optimizer.
Box bounds declared on variables are enforced by projection after each update.
- Parameters:
optimizer (str, optional) – Only “adam” is supported.
steps (int, optional) – Number of optimizer steps.
learning_rate (float, optional) – Adam learning rate.
jit (bool, optional) – If True (default), jit-compile the value-and-gradient function.
b1 (float, optional) – Adam hyperparameters.
b2 (float, optional) – Adam hyperparameters.
eps (float, optional) – Adam hyperparameters.
verbose (bool, optional) – If True, print the loss every 10 steps.
- Return type:
Solution¶
PiecewiseConstant¶
- class unsim.api.PiecewiseConstant(breakpoints, values)¶
Piecewise-constant time profile.
Callable as
profile(t), so it can be used anywhere a time function is expected, e.g.World.set_toll. Returns 0 outside the covered range.- Parameters:
breakpoints (list of float) – Interval boundaries (s), length n+1 for n values.
values (list of float) – Value on each interval [breakpoints[i], breakpoints[i+1]).