@nnx.jit def forward(model, x): return model(x) # This is JIT-compiled, just like @jax.jit output = forward(model, dummy_input) __ __