Commit d6aa133
select: keep the solver's Newton carry varying under shard_map
The solver adaptive-cutoff path seeded its fori_loop bracket from
constants (jnp.zeros / jnp.full), while the loop body's outputs derive
from r_ij. Under shard_map r_ij varies over the data-parallel mesh axis,
so carry-in and carry-out disagreed on manual axis types and scan's
carry type check rejected the loop at trace time:
TypeError: scan body function carry input and carry output must have
equal types [...] float32[1024] vs float32[1024]{V:dp}
Any data-parallel run using adaptive_cutoff_method="solver" failed to
trace; single-device runs were unaffected, since callers only apply
shard_map when there is more than one device.
pcast the bracket to varying, reading the axes off the data so the fix
is agnostic to the caller's mesh axis names and a no-op outside
shard_map. Runtime behaviour is unchanged: each shard already solves on
its own atoms, and the annotation carries no computation.
Adds tests/test_shard_map.py, which runs both adaptive-cutoff methods
under shard_map on forced CPU devices and checks values and gradients
against the single-device result -- multi-device coverage the suite
lacked. The device count is forced in conftest, before any test module
imports jax.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>1 parent b3dc553 commit d6aa133
3 files changed
Lines changed: 95 additions & 3 deletions
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
298 | 298 | | |
299 | 299 | | |
300 | 300 | | |
301 | | - | |
302 | | - | |
303 | | - | |
| 301 | + | |
| 302 | + | |
| 303 | + | |
| 304 | + | |
| 305 | + | |
| 306 | + | |
| 307 | + | |
| 308 | + | |
304 | 309 | | |
305 | 310 | | |
306 | 311 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
12 | 12 | | |
13 | 13 | | |
14 | 14 | | |
| 15 | + | |
15 | 16 | | |
16 | 17 | | |
17 | 18 | | |
18 | 19 | | |
| 20 | + | |
| 21 | + | |
| 22 | + | |
| 23 | + | |
| 24 | + | |
19 | 25 | | |
20 | 26 | | |
21 | 27 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
| 1 | + | |
| 2 | + | |
| 3 | + | |
| 4 | + | |
| 5 | + | |
| 6 | + | |
| 7 | + | |
| 8 | + | |
| 9 | + | |
| 10 | + | |
| 11 | + | |
| 12 | + | |
| 13 | + | |
| 14 | + | |
| 15 | + | |
| 16 | + | |
| 17 | + | |
| 18 | + | |
| 19 | + | |
| 20 | + | |
| 21 | + | |
| 22 | + | |
| 23 | + | |
| 24 | + | |
| 25 | + | |
| 26 | + | |
| 27 | + | |
| 28 | + | |
| 29 | + | |
| 30 | + | |
| 31 | + | |
| 32 | + | |
| 33 | + | |
| 34 | + | |
| 35 | + | |
| 36 | + | |
| 37 | + | |
| 38 | + | |
| 39 | + | |
| 40 | + | |
| 41 | + | |
| 42 | + | |
| 43 | + | |
| 44 | + | |
| 45 | + | |
| 46 | + | |
| 47 | + | |
| 48 | + | |
| 49 | + | |
| 50 | + | |
| 51 | + | |
| 52 | + | |
| 53 | + | |
| 54 | + | |
| 55 | + | |
| 56 | + | |
| 57 | + | |
| 58 | + | |
| 59 | + | |
| 60 | + | |
| 61 | + | |
| 62 | + | |
| 63 | + | |
| 64 | + | |
| 65 | + | |
| 66 | + | |
| 67 | + | |
| 68 | + | |
| 69 | + | |
| 70 | + | |
| 71 | + | |
| 72 | + | |
| 73 | + | |
| 74 | + | |
| 75 | + | |
| 76 | + | |
| 77 | + | |
| 78 | + | |
| 79 | + | |
| 80 | + | |
| 81 | + | |
0 commit comments