setup#

anri.utils.setup(n_cpu=4, preallocate=False)[source]#

Configure JAX for this machine. Must run before JAX has computed anything.

Parameters:
  • n_cpu (int, default: 4) – Number of XLA CPU devices. Each device also uses several threads, so this is not the number of cores used. Rendering one row of the quartz phantom on 16 cores took 9.0 s with 1 device, 6.3 s with 4 and 8.2 s with 16: more devices than that compete for the same cores. Limit the cores themselves with taskset or a SLURM allocation.

  • preallocate (bool, default: False) – If False (default), JAX allocates GPU memory as needed instead of reserving 75% of it at start-up. Only applied if XLA_PYTHON_CLIENT_PREALLOCATE is not already set.

Return type:

None

Notes

Matrix products run at full float32 precision (jax_default_matmul_precision = "highest"), unless JAX_DEFAULT_MATMUL_PRECISION is set. JAX’s default on NVIDIA GPUs from Ampere on uses TF32 for float32 matrix products, with a 10-bit mantissa: on an L40S it moved rendered spots ~1000 px from the beam centre by up to 0.3 px, a strain error of ~1e-5. The products in anri are small (3x3), so this costs nothing measurable.

With jaxlib >= 0.11, this also turns off XLA:CPU’s YNNPACK fusions (--xla_cpu_experimental_ynn_fusion_type=), unless XLA_FLAGS already sets that flag. They miscompile anri.fwd.render_peaks() for batches of more than a few thousand peaks: in jaxlib 0.11.1 and 0.11.2 most peaks were squeezed into a single pixel (float64) or came out as NaN (float32, from 0.11.0). Older jaxlib does not have the flag, and XLA aborts on unknown flags.