diff --git a/jaccpot/nearfield/_fast_lane.py b/jaccpot/nearfield/_fast_lane.py index 14c0615d..61d846b8 100644 --- a/jaccpot/nearfield/_fast_lane.py +++ b/jaccpot/nearfield/_fast_lane.py @@ -719,6 +719,11 @@ def _radix_fast_lane_prepacked_pallas( ) source_valid_flat = source_valid_mask_padded.reshape((num_leaves, num_source_slots)) + # Source-chunk grid axis (see nearfield_leafpair_pallas): 64 slots per program + # by default; 0 restores the single-pass grid. The longest neighbour row of a + # centrally concentrated distribution is num_leaves-1, and one serial warp + # over it bounded the whole kernel from below (19 ms at 200k, 38 at 400k). + source_chunk = _env_int("JACCPOT_NEARFIELD_LEAFPAIR_SOURCE_CHUNK", 64) out = nearfield_leafpair_pallas( leaf_positions, leaf_masses, @@ -731,6 +736,7 @@ def _radix_fast_lane_prepacked_pallas( num_stages=num_stages, target_subtile=target_subtile, interpret=interpret, + source_chunk=(None if source_chunk <= 0 else int(source_chunk)), ) pair_acc = _scatter_contributions( diff --git a/jaccpot/nearfield/_kernels.py b/jaccpot/nearfield/_kernels.py index 7c1f8484..432853e1 100644 --- a/jaccpot/nearfield/_kernels.py +++ b/jaccpot/nearfield/_kernels.py @@ -31,6 +31,7 @@ from jax import lax from jaxtyping import Array, Bool, Float, jaxtyped +from jaccpot._env import env_int from jaccpot.runtime.grad_options import analytic_p2p_vjp_enabled from .grad import _pair_accel_cvjp @@ -167,6 +168,51 @@ def compute_single(args: tuple[Array, Array, Array]) -> tuple[Array, Array]: # cross-leaf pair blocks, while intra-leaf self interaction is computed here. _compute_single_remat = jax.checkpoint(compute_single) + # Batch the scan: one leaf per step is 5 launches of ~2 us each per leaf + # (7.5 ms at N=200k/leaf 256, 15 ms at leaf 128 where the GPU sits 41 % + # idle -- measured 2026-09-06). Vmapping ``batch`` leaves per step cuts the + # launch count by ``batch`` for ``batch * W^2 * 3`` words of live + # intermediates (25 MB at 32 x 256), so the remat argument above still + # holds per step. 1 restores the historical one-leaf-per-step scan. + batch = env_int("JACCPOT_NEARFIELD_SELF_BATCH", 32, minimum=1) + num_leaves = int(leaf_positions.shape[0]) + if batch > 1 and num_leaves > 0: + pad = (-num_leaves) % batch + if pad: + leaf_positions_b = jnp.pad(leaf_positions, ((0, pad), (0, 0), (0, 0))) + leaf_masses_b = jnp.pad(leaf_masses, ((0, pad), (0, 0))) + mask_b = jnp.pad(mask, ((0, pad), (0, 0))) + else: + leaf_positions_b, leaf_masses_b, mask_b = leaf_positions, leaf_masses, mask + n_steps = (num_leaves + pad) // batch + leaf_positions_b = leaf_positions_b.reshape( + (n_steps, batch) + tuple(leaf_positions.shape[1:]) + ) + leaf_masses_b = leaf_masses_b.reshape( + (n_steps, batch) + tuple(leaf_masses.shape[1:]) + ) + mask_b = mask_b.reshape((n_steps, batch) + tuple(mask.shape[1:])) + _compute_batch = jax.vmap(_compute_single_remat) + + def scan_step_batched( + carry: Any, args: tuple[Array, Array, Array] + ) -> tuple[Any, tuple[Array, Array]]: + accel_b, pot_b = _compute_batch(args) + return carry, (accel_b, pot_b) + + _, (accels_b, potentials_b) = lax.scan( + scan_step_batched, None, (leaf_positions_b, leaf_masses_b, mask_b) + ) + accels = accels_b.reshape((n_steps * batch,) + tuple(accels_b.shape[2:]))[ + :num_leaves + ] + potentials = potentials_b.reshape( + (n_steps * batch,) + tuple(potentials_b.shape[2:]) + )[:num_leaves] + if compute_potential: + return accels, potentials + return accels, None + def scan_step( carry: Any, args: tuple[Array, Array, Array] ) -> tuple[Any, tuple[Array, Array]]: diff --git a/jaccpot/pallas/nearfield_fused_leaf.py b/jaccpot/pallas/nearfield_fused_leaf.py index db4ad65f..0b409102 100644 --- a/jaccpot/pallas/nearfield_fused_leaf.py +++ b/jaccpot/pallas/nearfield_fused_leaf.py @@ -708,9 +708,14 @@ def _nearfield_leafpair_kernel( num_source_slots: int, leaf_width: int, accum_dtype: Any = None, + out_dtype: Any = None, ) -> None: """Leaf-pair near-field update for one target subtile (vector of Bt targets). + ``out_dtype`` (chunked grid only) writes the accumulator without the final + downcast when it equals ``accum_dtype``, so per-chunk partials keep the wide + precision until the caller's reduce. + The production lane. Sources arrive as leaf *ids* and are gathered inside the kernel from the full particle tables, which is what avoids materialising the dense ``(num_leaves, num_source_slots, W_s)`` tensor that is ~99% padding and OOMs at @@ -768,6 +773,11 @@ def _nearfield_leafpair_kernel( cuts the wide adds by ``leaf_width`` (512x at the production leaf) while still removing the dominant ``sqrt(N)`` term. + out_dtype : Any + Dtype of ``out_ref``. ``None`` means the input dtype (one final downcast + of the wide accumulator). The chunked grid passes ``accum_dtype`` so each + chunk's partial keeps the wide precision until the caller's reduce. + Returns ------- None @@ -828,7 +838,8 @@ def _lane_body(j, acc): return lax.cond(slot_valid, _apply, lambda acc: acc, acc) acc_x, acc_y, acc_z, acc_p = lax.fori_loop(0, num_source_slots, _slot_body, acc0) - if wide: + target_out_dtype = zero.dtype if out_dtype is None else out_dtype + if wide and target_out_dtype != accum_dtype: # One downcast, on the FINAL value rather than on the sum being accumulated: # it costs one eps of the result (~6e-8), which is 170x below the 1.1e-05 # target, and it keeps every dtype outside this kernel unchanged. @@ -836,11 +847,12 @@ def _lane_body(j, acc): acc_y = acc_y.astype(zero.dtype) acc_z = acc_z.astype(zero.dtype) acc_p = acc_p.astype(zero.dtype) + zero_out = jnp.zeros_like(acc_x) - out_ref[0, :, 0] = jnp.where(tvalid, acc_x, zero) - out_ref[0, :, 1] = jnp.where(tvalid, acc_y, zero) - out_ref[0, :, 2] = jnp.where(tvalid, acc_z, zero) - out_ref[0, :, 3] = jnp.where(tvalid, acc_p, zero) + out_ref[0, :, 0] = jnp.where(tvalid, acc_x, zero_out) + out_ref[0, :, 1] = jnp.where(tvalid, acc_y, zero_out) + out_ref[0, :, 2] = jnp.where(tvalid, acc_z, zero_out) + out_ref[0, :, 3] = jnp.where(tvalid, acc_p, zero_out) @jaxtyped(typechecker=beartype) @@ -858,9 +870,24 @@ def nearfield_leafpair_pallas( target_subtile: int | None = None, interpret: bool = False, accum: str = "input", + source_chunk: int | None = None, ) -> Array: """Leaf-pair near-field update with Pallas. + ``source_chunk`` splits each target's ``S`` source slots across + ``ceil(S / source_chunk)`` programs that each write a partial sum, reduced + afterwards. Why (measured 2026-09-06, N=200k Plummer, leaf 256, A100): the + grid is one single-warp program per (leaf, subtile) looping over that + leaf's whole neighbour row, and in a centrally concentrated distribution + the longest row is ``num_leaves - 1`` (a halo leaf so extended that the + mutual MAC makes it near to every leaf). That one warp serially sums all N + sources at ~100 ns per lane-step -- 19 ms at 200k, 38 ms at 400k, theta- and + order-independent -- and the kernel's wall time is bounded below by it: + truncating rows 781 -> 150 slots removed 18 ms while dropping 21 % of the + entries. With chunk 64 the longest program is 64 x W lane-steps and the + ~100k programs balance. ``None`` (or ``>= S``) keeps the historical single + pass, bit-for-bit. + See :func:`nearfield_leafpair_jax` for the argument/return contract. Source leaves are gathered by id from ``leaf_positions`` inside the kernel; invalid source slots are skipped with ``lax.cond`` so heavily-padded slot tensors @@ -903,6 +930,10 @@ def nearfield_leafpair_pallas( the whole fix: measured 439x in force accuracy for 1.8 % in time on the distributed lane at 10^7 particles. + source_chunk : int | None + Source slots per program on the chunked grid (see above). ``None`` or a + value of at least ``S`` keeps the single-pass grid. + Returns ------- Array @@ -960,40 +991,114 @@ def nearfield_leafpair_pallas( accum_dtype = _resolve_accum_dtype(accum, dtype) - def _kernel(*refs): + chunk = None + if source_chunk is not None and 0 < int(source_chunk) < num_source_slots: + chunk = int(source_chunk) + if chunk is None: + + def _kernel(*refs): + return _nearfield_leafpair_kernel( + *refs, + num_source_slots=num_source_slots, + leaf_width=leaf_width, + accum_dtype=accum_dtype, + ) + + kernel = pl.pallas_call( + _kernel, + out_shape=jax.ShapeDtypeStruct((num_leaves, width_pad, _OUT_WIDTH), dtype), + in_specs=[ + pl.BlockSpec((1, bt, _POS_WIDTH), lambda leaf, sub: (leaf, sub, 0)), + pl.BlockSpec((1, bt), lambda leaf, sub: (leaf, sub)), + # Full gather tables (indexed by data-dependent source leaf id). + pl.BlockSpec( + (num_leaves, leaf_width, _POS_WIDTH), lambda leaf, sub: (0, 0, 0) + ), + pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub: (0, 0)), + pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub: (0, 0)), + pl.BlockSpec((1, num_source_slots), lambda leaf, sub: (leaf, 0)), + pl.BlockSpec((1, num_source_slots), lambda leaf, sub: (leaf, 0)), + pl.BlockSpec((1,), lambda leaf, sub: (0,)), + pl.BlockSpec((1,), lambda leaf, sub: (0,)), + ], + out_specs=pl.BlockSpec( + (1, bt, _OUT_WIDTH), lambda leaf, sub: (leaf, sub, 0) + ), + grid=(num_leaves, n_sub), + compiler_params=plgpu.CompilerParams( + num_warps=int(num_warps), num_stages=int(num_stages) + ), + interpret=bool(interpret), + name=f"nearfield_leafpair_t{bt}_s{num_source_slots}_w{leaf_width}_a{accum}", + ) + out = kernel( + target_positions_padded, + target_mask_padded, + leaf_positions_padded, + leaf_masses, + leaf_mask, + source_leaf_ids, + source_valid, + softening_sq_arr, + g_arr, + ) + if pad_t: + out = out[:, :leaf_width, :] + return out + + # Chunked grid: (leaf, subtile, source chunk). Each program sums its chunk of + # source slots into a partial; the partials are reduced below. The reduce is + # over at most ceil(S / chunk) terms per target, so it costs nothing against + # the sums inside the programs; with the wide accumulator the partials are + # emitted in the wide dtype so nothing is lost before the reduce. + n_chunks = (num_source_slots + chunk - 1) // chunk + slots_pad = n_chunks * chunk - num_source_slots + if slots_pad: + source_leaf_ids = jnp.pad(source_leaf_ids, ((0, 0), (0, slots_pad))) + source_valid = jnp.pad(source_valid, ((0, 0), (0, slots_pad))) + partial_dtype = accum_dtype if accum_dtype is not None else dtype + + def _kernel_chunk(*refs): return _nearfield_leafpair_kernel( *refs, - num_source_slots=num_source_slots, + num_source_slots=chunk, leaf_width=leaf_width, accum_dtype=accum_dtype, + out_dtype=partial_dtype, ) kernel = pl.pallas_call( - _kernel, - out_shape=jax.ShapeDtypeStruct((num_leaves, width_pad, _OUT_WIDTH), dtype), + _kernel_chunk, + out_shape=jax.ShapeDtypeStruct( + (num_leaves, n_chunks, width_pad, _OUT_WIDTH), partial_dtype + ), in_specs=[ - pl.BlockSpec((1, bt, _POS_WIDTH), lambda leaf, sub: (leaf, sub, 0)), - pl.BlockSpec((1, bt), lambda leaf, sub: (leaf, sub)), - # Full gather tables (indexed by data-dependent source leaf id). + pl.BlockSpec((1, bt, _POS_WIDTH), lambda leaf, sub, c: (leaf, sub, 0)), + pl.BlockSpec((1, bt), lambda leaf, sub, c: (leaf, sub)), pl.BlockSpec( - (num_leaves, leaf_width, _POS_WIDTH), lambda leaf, sub: (0, 0, 0) + (num_leaves, leaf_width, _POS_WIDTH), lambda leaf, sub, c: (0, 0, 0) ), - pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub: (0, 0)), - pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub: (0, 0)), - pl.BlockSpec((1, num_source_slots), lambda leaf, sub: (leaf, 0)), - pl.BlockSpec((1, num_source_slots), lambda leaf, sub: (leaf, 0)), - pl.BlockSpec((1,), lambda leaf, sub: (0,)), - pl.BlockSpec((1,), lambda leaf, sub: (0,)), + pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub, c: (0, 0)), + pl.BlockSpec((num_leaves, leaf_width), lambda leaf, sub, c: (0, 0)), + pl.BlockSpec((1, chunk), lambda leaf, sub, c: (leaf, c)), + pl.BlockSpec((1, chunk), lambda leaf, sub, c: (leaf, c)), + pl.BlockSpec((1,), lambda leaf, sub, c: (0,)), + pl.BlockSpec((1,), lambda leaf, sub, c: (0,)), ], - out_specs=pl.BlockSpec((1, bt, _OUT_WIDTH), lambda leaf, sub: (leaf, sub, 0)), - grid=(num_leaves, n_sub), + out_specs=pl.BlockSpec( + (1, None, bt, _OUT_WIDTH), lambda leaf, sub, c: (leaf, c, sub, 0) + ), + grid=(num_leaves, n_sub, n_chunks), compiler_params=plgpu.CompilerParams( num_warps=int(num_warps), num_stages=int(num_stages) ), interpret=bool(interpret), - name=f"nearfield_leafpair_t{bt}_s{num_source_slots}_w{leaf_width}_a{accum}", + name=( + f"nearfield_leafpair_t{bt}_s{num_source_slots}_c{chunk}" + f"_w{leaf_width}_a{accum}" + ), ) - out = kernel( + partials = kernel( target_positions_padded, target_mask_padded, leaf_positions_padded, @@ -1004,6 +1109,7 @@ def _kernel(*refs): softening_sq_arr, g_arr, ) + out = jnp.sum(partials, axis=1).astype(dtype) if pad_t: out = out[:, :leaf_width, :] return out