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 scalar where x maps each variable to its current value array.

  • checkpoint_every (int or None, optional) – Gradient checkpointing segment length.

Return type:

Problem

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:

Objective

parameter(entity, field)

Get a ParameterRef for an entity’s physical quantity.

Parameters:
  • entity (Link, Node, or Demand (or a link/node name)) – Target entity.

  • field (str) – Quantity name, e.g. “free_flow_speed”, “capacity”, “merge_priority”, “toll”, “flow_capacity”, “flow”.

Return type:

ParameterRef

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.

Parameters:

entity (Link, Node, or Demand (or a link/node name)) – Target entity.

Return type:

list of str

path(links)

Build a PathRef from an ordered list of links (objects or names).

Parameters:

links (list of Link or str) – Ordered links from origin to destination.

Return type:

PathRef

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:

SimResult

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:

TollVariable

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) -> Params where theta has shape shape.

  • 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:

CustomVariable

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.

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

class unsim.api.PathView(result, path_ref)

Differentiable path-level queries on a simulation result.

travel_time(departure_time)

Travel time of a virtual vehicle along the path (s). Differentiable.

Parameters:

departure_time (float) – Departure time from the path’s first link (s).

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

class unsim.api.Gradients(mapping)

Read-only mapping from ParameterRef to gradient arrays.

items()

Iterate over (ref, gradient) pairs.

keys()

Iterate over refs.

values()

Iterate over gradient arrays.

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

class unsim.api.PathRef(link_ids: tuple, link_names: tuple = ())

Semantic reference to an ordered path of links.

Ordered link indices from origin to destination.

Type:

tuple of int

Corresponding link names.

Type:

tuple of str

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().

Selected link indices.

Type:

tuple of int

n_steps

Number of toll discretization steps.

Type:

int

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) -> Params where theta has shape shape.

  • 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 scalar where x maps 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

Solution

class unsim.api.Solution(value_map, theta, loss_history)

Result of Problem.solve.

value

Mapping from each variable to its optimized value array.

Type:

Gradients

theta

Final flat variable vector.

Type:

jnp.ndarray

loss_history

Loss at the start of each optimizer step.

Type:

np.ndarray, (steps,)

property final_loss

Loss at the start of the last optimizer step.

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]).