Basic tutorial of UNsim

In this notebook, the basic feature of UNsim is demonstrated.

This tutorial is currently under development. The current explanation content remains minimal. For the basic usage of UNsim, the tutorial of UXsim (a mesoscopic traffic flow simulator with very similar features and syntax) might be useful.

Python-based forward simulation

The most basic feature of UNsim is a macroscopic traffic simulation in Python.

Simple Y-shaped network

First, import the necessary modules.

[1]:
from unsim import World
import matplotlib.pyplot as plt

Define the main simulation environment World. The most important args are tmax (simulation duration) and deltat (time step width).

[2]:
# Units are standardized to seconds (s) and meters (m)
W = World(
    name="merge", deltat=5, tmax=1200,
    print_mode=1, save_mode=0, show_mode=1,
)

Then, setup the network by defining Node and Link. The Node’s x and y coordinates are for visualization only. The actual network size is determined by length of `Link.

[3]:
W.addNode("orig1", x=0, y=0)
W.addNode("orig2", x=0, y=2)
W.addNode("merge", x=1, y=1)
W.addNode("dest", x=2, y=1)
W.addLink("link1", "orig1", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link2", "orig2", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link3", "merge", "dest", length=1000, free_flow_speed=20, capacity=0.8)

[3]:
<Link 'link3'>

And add the traffic demand. flow means demand flow-rate in vehicle per second unit.

[4]:
W.adddemand("orig1", "dest", t_start=0, t_end=1000, flow=0.45)
W.adddemand("orig2", "dest", t_start=400, t_end=1000, flow=0.6)
[4]:
<Demand orig2->dest [400, 1000) flow=0.6>

Now we have set up a Y-shaped merging network. The capacity of the downstream link is 0.8 veh/s, and the total demand is 0.45 for until 400 s and 1.05 s after 400 s. We can expect that congestion will happen at the merging node some time after 400 s.

The simulation can be executed by

[5]:
W.exec_simulation()
 Simulation completed. merge

Let’s see some statistics and visualizations.

[6]:
#stats
W.analyzer.print_simple_stats()

#network plots
W.analyzer.network(t=200)
W.analyzer.network(t=800)

#link time-space diagrams
W.analyzer.time_space_diagram(mode="k_norm", links="link1")
W.analyzer.time_space_diagram(mode="k_norm", links="link2")
W.analyzer.time_space_diagram(mode="k_norm", links="link3")

plt.show()
  Simulation Results:
    Total trips:     810.0
    Completed trips: 740.0
    Total travel time: 136825.0 s
    Avg travel time: 184.9 s
    Avg delay:       84.9 s
_images/tutorial_12_1.png
_images/tutorial_12_2.png
_images/tutorial_12_3.png
_images/tutorial_12_4.png
_images/tutorial_12_5.png

You can see the congestion as expected.

Large-scale network with program-based generation

to be added

JAX-based acceleration and differentiation

By using JAX, a simulation World of UNsim can be compiled into a differentiable model that runs fast on CPU or GPU and can be differentiated by Automatic Differentiation.

Let’s define the same scenario as the previous Y network. Note that this scenario is very small and not suitable for JAX-based high performance computation; the compilation overhead becomes significant enough to actually slow down the process. But since this is a tutorial, let’s proceed.

We keep the returned Link and Demand objects in variables, because we will refer to them when specifying differentiation targets.

[7]:
from unsim import World

W = World(
    name="merge", deltat=5, tmax=1200,
    print_mode=1, save_mode=0, show_mode=1,
)
W.addNode("orig1", x=0, y=0)
W.addNode("orig2", x=0, y=2)
W.addNode("merge", x=1, y=1)
W.addNode("dest", x=2, y=1)
link1 = W.addLink("link1", "orig1", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
link2 = W.addLink("link2", "orig2", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
link3 = W.addLink("link3", "merge", "dest", length=1000, free_flow_speed=20, capacity=0.8)
demand1 = W.adddemand("orig1", "dest", t_start=0, t_end=1000, flow=0.45)
demand2 = W.adddemand("orig2", "dest", t_start=400, t_end=1000, flow=0.6)

Compile the World into a differentiable model M and run it. M is an immutable snapshot: later modifications to W do not affect it.

[8]:
M = W.compile(backend="jax")
R = M.run()

The result object R provides the same traffic metrics as the Python-based simulation. R.metrics.* returns JAX scalars usable inside objective functions, and R.link(...) / R.node(...) give per-element queries. R.analyzer gives the same host-side analyzer (plots and stats) as the Python version.

[9]:
print(f"Total travel time: {R.metrics.total_travel_time():.1f} s")
print(f"Completed trips: {R.metrics.completed_trips():.1f} veh")
print(f"Vehicles on link3 at t=600 s: {R.link(link3).vehicle_count(600):.1f} veh")
Total travel time: 136824.9 s
Completed trips: 740.0 veh
Vehicles on link3 at t=600 s: 40.0 veh

Automatic Differentiation

The great feature of JAX is the Automatic Differentiation (AD): we can compute the gradient of any simulation output with respect to any input parameter.

A differentiation target is specified by the traffic object itself: M.parameter(link, "free_flow_speed") or M.parameter(demand, "flow").

Let’s compute \(\partial TTT/\partial u_l\) for each link and \(\partial TTT/\partial q_1\) for the first demand. The speed gradients should be negative, because increasing free flow speeds should decrease total travel time. The demand gradient should be positive, because increasing demand should increase total travel time.

[10]:
u1 = M.parameter(link1, "free_flow_speed")
u2 = M.parameter(link2, "free_flow_speed")
u3 = M.parameter(link3, "free_flow_speed")
q1 = M.parameter(demand1, "flow")

objective = M.objective(lambda R: R.metrics.total_travel_time())
value, grad = objective.value_and_gradient(wrt=[u1, u2, u3, q1])

print(f"TTT: {value:.1f} s")
print("Gradient of TTT w.r.t. free-flow speed:")
for ref in [u1, u2, u3]:
    print(f"  {ref.name}: {float(grad[ref]):.2f}")
print(f"Gradient of TTT w.r.t. demand1 flow: {float(grad[q1]):.1f}")
TTT: 136824.9 s
Gradient of TTT w.r.t. free-flow speed:
  link1: -1226.95
  link2: -618.75
  link3: -1847.19
Gradient of TTT w.r.t. demand1 flow: 335125.2

We get expected results. Notice that \(\partial TTT/\partial u_3\) has the largest absolute value, since all traffic must use link3 after the merge.

You can list the quantities accepted by M.parameter(entity, field) with M.parameter_fields(entity) as follows. For a link, the accepted fundamental-diagram quantities depend on the link’s FD parameterization (how the link was defined in addLink); the other quantities are always accepted.

[11]:
print(link1, M.parameter_fields(link1))
print(demand1, M.parameter_fields(demand1))
<Link 'link1'> ['free_flow_speed', 'backward_wave_speed', 'capacity', 'merge_priority', 'capacity_out', 'capacity_in', 'toll']
<Demand orig1->dest [0, 1000) flow=0.45> ['flow']

You can ask the objective what exactly is being differentiated:

[12]:
print(objective.explain(wrt=[u3, q1]))
Objective:
    custom function

Variables:
    link3.free_flow_speed [m/s]
        FD parameterization: (u, w, capacity)
        derived quantities: jam_density
    orig1->dest[0,1000).flow [veh/s]
        active interval: [0, 1000)

Route choice:
    fix
    gradient through discrete shortest-path index: no

Static quantities:
    topology, time grid, link length

Custom variables for advanced use

The objective can be any JAX expression of the result R, not only the built-in metrics; for example lambda R: R.link(link2).vehicle_count(600.0) (computed internally as \(N_U - N_D\)) is a valid objective. Similarly, differentiation variables are not limited to the built-in parameter refs: M.variable() defines a variable with a user-supplied injection function into the parameters.

Here, let’s differentiate the model with respect to jam_density \(\kappa\). Under the (u, w, capacity) parameterization, \(\kappa\) is a derived quantity, so M.parameter refuses it because the gradient would be ambiguous. If you nevertheless want \(\partial TTT/\partial \kappa\), a custom variable makes the choice explicit: the injection below writes \(\kappa\) of link2 directly, which holds \(u\) and \(q^*\) fixed and lets \(w = q^* u/(u\kappa - q^*)\) vary.

[13]:
kappa2 = M.variable(
    shape=(1,), initial=float(M.raw.params.kappa[1]), name="kappa_link2",
    inject=lambda p, th: p._replace(kappa=p.kappa.at[1].set(th[0])),
)
value, grad = objective.value_and_gradient(wrt=[kappa2])
print(f"TTT: {value:.1f} s")
print(f"Gradient of TTT w.r.t. jam_density of link2: {float(grad[kappa2][0]):.1f}")
TTT: 136824.9 s
Gradient of TTT w.r.t. jam_density of link2: 37500.0

Low-level JAX API (advanced use / backward compatibility)

For further advanced use, the low-level JAX structures remain accessible via M.raw.params, M.raw.config, and M.raw.simulate, and all functions in unsim.unsim_diff keep working. The object-oriented API above is a facade over a functional JAX core. This section demonstrates the raw interface of that core, which remains available for backward compatibility and fine-grained control. This requires knowledge on JAX itself.

Let’s define the same scenario to the previous Y network again.

[14]:
W = World(
    name="merge", deltat=5, tmax=1200,
    print_mode=1, save_mode=0, show_mode=1,
)
W.addNode("orig1", x=0, y=0)
W.addNode("orig2", x=0, y=2)
W.addNode("merge", x=1, y=1)
W.addNode("dest", x=2, y=1)
W.addLink("link1", "orig1", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link2", "orig2", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link3", "merge", "dest", length=1000, free_flow_speed=20, capacity=0.8)
W.adddemand("orig1", "dest", t_start=0, t_end=1000, flow=0.45)
W.adddemand("orig2", "dest", t_start=400, t_end=1000, flow=0.6)
[14]:
<Demand orig2->dest [400, 1000) flow=0.6>

Now we import the necessary module, convert the World to JAX object, and run JAX-based simulation.

[15]:
from unsim.unsim_diff import *

params, config = world_to_jax(W)
state = simulate(params, config)

You can access the results like this.

[16]:
ttt = total_travel_time(state, config)
print(f"Total travel time: {ttt:.1f} s")
Total travel time: 136824.9 s

Automatic Differentiation (low-level)

With the raw interface, gradients are computed by jax.grad over a loss function that rebuilds Params via _replace and indexes arrays by integer link IDs.

Let’s reuse the previous Y scenario again.

[17]:
from unsim import World
from unsim.unsim_diff import *

W = World(
    name="merge", deltat=5, tmax=1200,
    print_mode=1, save_mode=0, show_mode=1,
)
W.addNode("orig1", x=0, y=0)
W.addNode("orig2", x=0, y=2)
W.addNode("merge", x=1, y=1)
W.addNode("dest", x=2, y=1)
W.addLink("link1", "orig1", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link2", "orig2", "merge", length=1000, free_flow_speed=20, capacity=0.8, merge_priority=1)
W.addLink("link3", "merge", "dest", length=1000, free_flow_speed=20, capacity=0.8)
W.adddemand("orig1", "dest", t_start=0, t_end=1000, flow=0.45)
W.adddemand("orig2", "dest", t_start=400, t_end=1000, flow=0.6)

params, config = world_to_jax(W)
state = simulate(params, config)

ttt = total_travel_time(state, config)
print(f"Total travel time: {ttt:.1f} s")

Total travel time: 136824.9 s

By using this ttt object, you can compute gradient of \(TTT\) (total travel time) with respect to the input parameters contained in params.

For example, lets compute \(\partial TTT/\partial u_l\) where \(u_l\) denotes free-flow speed of link \(l\). It is expected that these values should be negative, as increasing the maximum speed should decrease the total travel time.

[18]:
def ttt_wrt_u(u):
    p = params._replace(u=u)
    s = simulate(p, config)
    return total_travel_time(s, config)

grad_u = jax.grad(ttt_wrt_u)(params.u)

print(f"\nGradient of TTT w.r.t. free-flow speed:")
for i, link in enumerate(W.LINKS):
    print(f"  {link.name}: {float(grad_u[i]):.2f}")

Gradient of TTT w.r.t. free-flow speed:
  link1: -1226.95
  link2: -543.75
  link3: -1847.19

We got an expected result. Furthermore, notice that \(\partial TTT/\partial u_3\) has largest absolute value. This indicates that the free-flow speed of link3 has the largest impact to TTT; this is also reasonable, as all traffic must use link3 after the merge.

FYI, the param contains the following elements. These parameters can be used to differentiate output variables. On the other hand, an output variable to be differentiated should be constructed from unsim_diff’s function. These definition can be tricky.

[19]:
[param for param in params.__dir__() if not param.startswith("_")]
[19]:
['u',
 'kappa',
 'q_star',
 'capacity_out',
 'capacity_in',
 'flow_capacity',
 'absorption_ratio',
 'diverge_ratios',
 'merge_priority',
 'demand_rate',
 'turning_fractions',
 'od_demand_rate',
 'toll',
 'route_bias',
 'index',
 'count']