Skip to content

Commit 5f89cf6

Browse files
committed
perf(ifu): add chunked particle accumulation controls
1 parent b52f9a8 commit 5f89cf6

3 files changed

Lines changed: 166 additions & 10 deletions

File tree

‎docs/development/gradient_production_pr_plan.md‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,10 +31,10 @@ proof-of-concept notebooks to production full-IFU workflows.
3131
4. `feat(optimizer-loop)` (completed)
3232
- Add reusable Optax optimization loop with histories/checkpoints.
3333

34-
5. `test(gradient-vs-fd)` (current)
34+
5. `test(gradient-vs-fd)` (completed)
3535
- Add finite-difference validation suite for gradient correctness.
3636

37-
6. `perf(full-ifu-scaling)`
37+
6. `perf(full-ifu-scaling)` (current)
3838
- Add chunking/checkpointing controls and consistency tests.
3939

4040
7. `feat(variational-inference)`

‎rubix/core/ifu.py‎

Lines changed: 98 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,84 @@
2020
from .telescope import get_telescope
2121

2222

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+
23101
@jaxtyped(typechecker=typechecker)
24102
def get_calculate_datacube_particlewise(config: dict) -> Callable:
25103
"""Prepare a per-particle datacube builder for the star component.
@@ -57,6 +135,7 @@ def get_calculate_datacube_particlewise(config: dict) -> Callable:
57135
ssp_wave0 = cosmological_doppler_shift(
58136
z=z_obs, wavelength=ssp_model.wavelength
59137
) # (n_wave_ssp,)
138+
chunk_size, use_remat = _get_performance_options(config)
60139

61140
@jaxtyped(typechecker=typechecker)
62141
def calculate_datacube_particlewise(rubixdata: RubixData) -> RubixData:
@@ -109,10 +188,15 @@ def body(cube, i):
109188
cube = cube.at[pix_i].add(spec_tel)
110189
return cube, None
111190

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,
116200
)
117201

118202
cube_3d = cube_flat.reshape(ns, ns, -1)
@@ -160,6 +244,7 @@ def get_calculate_dusty_datacube_particlewise(config: dict) -> Callable:
160244
ssp_wave0 = cosmological_doppler_shift(
161245
z=z_obs, wavelength=ssp_model.wavelength
162246
) # (n_wave_ssp,)
247+
chunk_size, use_remat = _get_performance_options(config)
163248

164249
@jaxtyped(typechecker=typechecker)
165250
def calculate_dusty_datacube_particlewise(
@@ -237,10 +322,15 @@ def body(cube, i):
237322
cube = cube.at[pix_i].add(spec_extincted)
238323
return cube, None
239324

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,
244334
)
245335

246336
cube_3d = cube_flat.reshape(ns, ns, -1)

‎tests/test_core_ifu_scaling.py‎

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,66 @@
1+
import jax.numpy as jnp
2+
3+
from rubix.core.ifu import _get_performance_options, _scan_particles
4+
5+
6+
def test_get_performance_options_defaults():
7+
chunk_size, use_remat = _get_performance_options({})
8+
9+
assert chunk_size == 0
10+
assert use_remat is False
11+
12+
13+
def test_get_performance_options_reads_valid_values():
14+
chunk_size, use_remat = _get_performance_options(
15+
{"performance": {"particle_chunk_size": 16, "remat_particlewise": True}}
16+
)
17+
18+
assert chunk_size == 16
19+
assert use_remat is True
20+
21+
22+
def test_get_performance_options_rejects_invalid_chunk_size():
23+
chunk_size, use_remat = _get_performance_options(
24+
{"performance": {"particle_chunk_size": -5, "remat_particlewise": False}}
25+
)
26+
27+
assert chunk_size == 0
28+
assert use_remat is False
29+
30+
31+
def _simple_step(cube, i):
32+
update = jnp.full_like(cube, i + 1, dtype=cube.dtype)
33+
return cube + update, None
34+
35+
36+
def test_scan_particles_chunked_matches_unchunked():
37+
init_cube = jnp.zeros((3, 4), dtype=jnp.float32)
38+
nstar = 11
39+
40+
unchunked = _scan_particles(
41+
init_cube=init_cube,
42+
nstar=nstar,
43+
step_fn=_simple_step,
44+
chunk_size=0,
45+
)
46+
chunked = _scan_particles(
47+
init_cube=init_cube,
48+
nstar=nstar,
49+
step_fn=_simple_step,
50+
chunk_size=4,
51+
)
52+
53+
assert jnp.allclose(unchunked, chunked)
54+
55+
56+
def test_scan_particles_returns_init_for_empty_input():
57+
init_cube = jnp.ones((2, 2), dtype=jnp.float32)
58+
59+
result = _scan_particles(
60+
init_cube=init_cube,
61+
nstar=0,
62+
step_fn=_simple_step,
63+
chunk_size=8,
64+
)
65+
66+
assert jnp.allclose(result, init_cube)

0 commit comments

Comments
 (0)