Skip to content

Simulator Model Gallery

Download notebook

This notebook sweeps through the built-in diffusion models. Each section shows how to instantiate from physical parameters or sample a random model via from_theta(theta) where theta ~ Normal(0, I) with dimension theta_dim.

Models: Ball, Stick, Zeppelin, DTI, Sphere, Cylinder, NODDI (Watson/Bingham), SANDI (Watson/Bingham).

%reload_ext autoreload
%autoreload 2
import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt
import numpy as np

from dmri.simulators import (
    Ball,
    Cylinder,
    Dti,
    NoddiB,
    NoddiW,
    SandiB,
    SandiW,
    Sphere,
    Stick,
    Zeppelin,
)
from dmri.simulators.acquisition_scheme import acquisition_scheme

Acquisition setup

We use 60 random directions and b-values from 0 to 3000 s/mm².

bvals = jnp.linspace(0, 3000, 60)
bvecs = jax.random.normal(jax.random.key(0), (60, 3))
bvecs = bvecs / jnp.linalg.norm(bvecs, axis=-1, keepdims=True)
acq = acquisition_scheme(bvals, bvecs)

Sampling via from_theta

All simulators expose from_theta(theta) where theta is a standard normal vector of length theta_dim. This mirrors how priors are parameterized during inference.

Each simulator defines its own transform from a standard normal vector to physical parameters. These transforms encode bounded, directional, or log-normal priors as needed.

import numpy as np
from dmri.simulators import Ball

theta = np.random.randn(Ball.theta_dim)
ball = Ball.from_theta(theta)

This can be used to easily simulate synthetic MRI signals for a large number of settings.

Ball (isotropic)

The Ball models unrestricted isotropic Gaussian diffusion with one diffusivity parameter. Its signal depends on the b-value but not the gradient direction.

ball = Ball(lam=0.0018)
ball_sig = ball.signal(acq)
plt.plot(bvals, ball_sig)
plt.title("Ball signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

Ball (isotropic)

Stick (oriented)

stick = Stick(mu=jnp.array([jnp.pi / 4, 0.0]), lam_par=0.0015)
stick_sig = stick.signal(acq)
plt.plot(bvals, stick_sig)
plt.title("Stick signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

Stick (oriented)

Zeppelin (axially symmetric tensor)

zepp = Zeppelin(mu=jnp.array([jnp.pi / 4, 0.0]), lambda_par=0.0015, lambda_perp=0.0005)
zepp_sig = zepp.signal(acq)
plt.plot(bvals, zepp_sig)
plt.title("Zeppelin signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

Zeppelin (axially symmetric tensor)

DTI (full tensor)

D = jnp.diag(jnp.array([0.0015, 0.0007, 0.0003]))
dti = Dti(D)
dti_sig = dti.signal(acq)
plt.plot(bvals, dti_sig)
plt.title("DTI signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

DTI (full tensor)

Sphere (restricted)

sphere = Sphere(radius=0.01)
sphere_sig = sphere.signal(acq)
plt.plot(bvals, sphere_sig)
plt.title("Sphere signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

Sphere (restricted)

Cylinder (finite radius)

cyl = Cylinder(mu=jnp.array([jnp.pi / 6, 0.0]), lam_par=0.0012, radius=0.01)
cyl_sig = cyl.signal(acq)
plt.plot(bvals, cyl_sig)
plt.title("Cylinder signal")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.show()

Cylinder (finite radius)

NODDI (Watson/Bingham variants)

noddi_w = NoddiW.from_theta(np.random.randn(NoddiW.theta_dim))
noddi_b = NoddiB.from_theta(np.random.randn(NoddiB.theta_dim))
# Parameters are implicit in kernels/FODs; just simulate.
noddi_w_sig = noddi_w.signal(acq)
noddi_b_sig = noddi_b.signal(acq)
plt.plot(bvals, noddi_w_sig, label="Watson")
plt.plot(bvals, noddi_b_sig, label="Bingham")
plt.title("NODDI signals")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.legend()
plt.show()

NODDI (Watson/Bingham variants)

SANDI (Watson/Bingham variants)

sandi_w = SandiW.from_theta(np.random.randn(SandiW.theta_dim))
sandi_b = SandiB.from_theta(np.random.randn(SandiB.theta_dim))
sandi_w_sig = sandi_w.signal(acq)
sandi_b_sig = sandi_b.signal(acq)
plt.plot(bvals, sandi_w_sig, label="Watson")
plt.plot(bvals, sandi_b_sig, label="Bingham")
plt.title("SANDI signals")
plt.xlabel("b-value [s/mm²]")
plt.ylabel("Signal")
plt.legend()
plt.show()

SANDI (Watson/Bingham variants)

Notes

  • Adjust mu angles to explore orientation effects.
  • Swap lambda_par, lambda_perp, radius, or lam to see sensitivity.
  • Add noise compartments or multi-compartment mixtures as needed.