|
20 | 20 | from .telescope import get_telescope |
21 | 21 |
|
22 | 22 |
|
| 23 | +@jaxtyped(typechecker=typechecker) |
| 24 | +def _get_performance_options(config: dict) -> tuple[int, bool]: |
| 25 | + """Read optional particlewise performance settings from the config. |
| 26 | +
|
| 27 | + Args: |
| 28 | + config (dict): Runtime configuration dictionary. |
| 29 | +
|
| 30 | + Returns: |
| 31 | + tuple[int, bool]: |
| 32 | + ``(chunk_size, use_remat)`` where ``chunk_size == 0`` means |
| 33 | + unchunked execution. |
| 34 | + """ |
| 35 | + perf_config = config.get("performance", {}) |
| 36 | + chunk_size = perf_config.get("particle_chunk_size", 0) |
| 37 | + if not isinstance(chunk_size, int) or chunk_size <= 0: |
| 38 | + chunk_size = 0 |
| 39 | + |
| 40 | + use_remat = bool(perf_config.get("remat_particlewise", False)) |
| 41 | + return chunk_size, use_remat |
| 42 | + |
| 43 | + |
| 44 | +@jaxtyped(typechecker=typechecker) |
| 45 | +def _scan_particles( |
| 46 | + init_cube: Float[Array, "n_spaxels n_wave_bins"], |
| 47 | + nstar: int, |
| 48 | + step_fn: Callable, |
| 49 | + chunk_size: int, |
| 50 | +) -> Float[Array, "n_spaxels n_wave_bins"]: |
| 51 | + """Accumulate per-particle contributions with optional chunking. |
| 52 | +
|
| 53 | + Args: |
| 54 | + init_cube (Float[Array, "n_spaxels n_wave_bins"]): Initial cube. |
| 55 | + nstar (int): Number of particles. |
| 56 | + step_fn (Callable): Particle step function of signature |
| 57 | + ``(cube, index) -> (cube, aux)``. |
| 58 | + chunk_size (int): Chunk size; ``0`` disables chunking. |
| 59 | +
|
| 60 | + Returns: |
| 61 | + Float[Array, "n_spaxels n_wave_bins"]: Accumulated flat cube. |
| 62 | + """ |
| 63 | + if nstar == 0: |
| 64 | + return init_cube |
| 65 | + |
| 66 | + if chunk_size <= 0: |
| 67 | + cube_flat, _ = lax.scan(step_fn, init_cube, jnp.arange(nstar, dtype=jnp.int32)) |
| 68 | + return cube_flat |
| 69 | + |
| 70 | + n_chunks = (nstar + chunk_size - 1) // chunk_size |
| 71 | + max_index = nstar - 1 |
| 72 | + |
| 73 | + def chunk_body(cube, chunk_idx): |
| 74 | + start = chunk_idx * chunk_size |
| 75 | + local = jnp.arange(chunk_size, dtype=jnp.int32) |
| 76 | + idx = start + local |
| 77 | + idx_safe = jnp.minimum(idx, max_index) |
| 78 | + valid = idx < nstar |
| 79 | + |
| 80 | + def inner_body(i, cube_inner): |
| 81 | + def do_step(current_cube): |
| 82 | + cube_new, _ = step_fn(current_cube, idx_safe[i]) |
| 83 | + return cube_new |
| 84 | + |
| 85 | + return lax.cond( |
| 86 | + valid[i], |
| 87 | + do_step, |
| 88 | + lambda current_cube: current_cube, |
| 89 | + cube_inner, |
| 90 | + ) |
| 91 | + |
| 92 | + cube = lax.fori_loop(0, chunk_size, inner_body, cube) |
| 93 | + return cube, None |
| 94 | + |
| 95 | + cube_flat, _ = lax.scan( |
| 96 | + chunk_body, init_cube, jnp.arange(n_chunks, dtype=jnp.int32) |
| 97 | + ) |
| 98 | + return cube_flat |
| 99 | + |
| 100 | + |
23 | 101 | @jaxtyped(typechecker=typechecker) |
24 | 102 | def get_calculate_datacube_particlewise(config: dict) -> Callable: |
25 | 103 | """Prepare a per-particle datacube builder for the star component. |
@@ -57,6 +135,7 @@ def get_calculate_datacube_particlewise(config: dict) -> Callable: |
57 | 135 | ssp_wave0 = cosmological_doppler_shift( |
58 | 136 | z=z_obs, wavelength=ssp_model.wavelength |
59 | 137 | ) # (n_wave_ssp,) |
| 138 | + chunk_size, use_remat = _get_performance_options(config) |
60 | 139 |
|
61 | 140 | @jaxtyped(typechecker=typechecker) |
62 | 141 | def calculate_datacube_particlewise(rubixdata: RubixData) -> RubixData: |
@@ -109,10 +188,15 @@ def body(cube, i): |
109 | 188 | cube = cube.at[pix_i].add(spec_tel) |
110 | 189 | return cube, None |
111 | 190 |
|
112 | | - cube_flat, _ = lax.scan( |
113 | | - body, |
114 | | - init_cube, |
115 | | - jnp.arange(nstar, dtype=jnp.int32), |
| 191 | + particle_step = body |
| 192 | + if use_remat: |
| 193 | + particle_step = jax.checkpoint(body) |
| 194 | + |
| 195 | + cube_flat = _scan_particles( |
| 196 | + init_cube=init_cube, |
| 197 | + nstar=nstar, |
| 198 | + step_fn=particle_step, |
| 199 | + chunk_size=chunk_size, |
116 | 200 | ) |
117 | 201 |
|
118 | 202 | cube_3d = cube_flat.reshape(ns, ns, -1) |
@@ -160,6 +244,7 @@ def get_calculate_dusty_datacube_particlewise(config: dict) -> Callable: |
160 | 244 | ssp_wave0 = cosmological_doppler_shift( |
161 | 245 | z=z_obs, wavelength=ssp_model.wavelength |
162 | 246 | ) # (n_wave_ssp,) |
| 247 | + chunk_size, use_remat = _get_performance_options(config) |
163 | 248 |
|
164 | 249 | @jaxtyped(typechecker=typechecker) |
165 | 250 | def calculate_dusty_datacube_particlewise( |
@@ -237,10 +322,15 @@ def body(cube, i): |
237 | 322 | cube = cube.at[pix_i].add(spec_extincted) |
238 | 323 | return cube, None |
239 | 324 |
|
240 | | - cube_flat, _ = lax.scan( |
241 | | - body, |
242 | | - init_cube, |
243 | | - jnp.arange(nstar, dtype=jnp.int32), |
| 325 | + particle_step = body |
| 326 | + if use_remat: |
| 327 | + particle_step = jax.checkpoint(body) |
| 328 | + |
| 329 | + cube_flat = _scan_particles( |
| 330 | + init_cube=init_cube, |
| 331 | + nstar=nstar, |
| 332 | + step_fn=particle_step, |
| 333 | + chunk_size=chunk_size, |
244 | 334 | ) |
245 | 335 |
|
246 | 336 | cube_3d = cube_flat.reshape(ns, ns, -1) |
|
0 commit comments