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 withtasksetor a SLURM allocation.preallocate (
bool, default:False) – IfFalse(default), JAX allocates GPU memory as needed instead of reserving 75% of it at start-up. Only applied ifXLA_PYTHON_CLIENT_PREALLOCATEis not already set.
- Return type:
Notes
Matrix products run at full float32 precision (
jax_default_matmul_precision = "highest"), unlessJAX_DEFAULT_MATMUL_PRECISIONis 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=), unlessXLA_FLAGSalready sets that flag. They miscompileanri.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.