Commit cd08cff
committed
Reshard loss under explicit mesh on the NNX path
loss_fn's NNX branch constrained xent/z_loss with raw
nn.with_logical_constraint, which keeps the size-1 context axis, while
the model shards activations via create_sharding (remove_size_one_mesh_axis
drops context). Under an explicit mesh the mismatch is a hard assert.
Use sharding.maybe_shard_with_logical, as the Linen branch already does:
it builds the sharding via create_sharding (dropping size-1 context, so it
matches the array) and reshards instead of asserting under explicit mode.1 parent 85690da commit cd08cff
2 files changed
Lines changed: 18 additions & 2 deletions
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
221 | 221 | | |
222 | 222 | | |
223 | 223 | | |
224 | | - | |
225 | | - | |
| 224 | + | |
| 225 | + | |
| 226 | + | |
| 227 | + | |
| 228 | + | |
| 229 | + | |
| 230 | + | |
| 231 | + | |
| 232 | + | |
| 233 | + | |
| 234 | + | |
| 235 | + | |
| 236 | + | |
| 237 | + | |
226 | 238 | | |
227 | 239 | | |
228 | 240 | | |
| |||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
| |||
23 | 23 | | |
24 | 24 | | |
25 | 25 | | |
| 26 | + | |
26 | 27 | | |
27 | 28 | | |
28 | 29 | | |
| |||
59 | 60 | | |
60 | 61 | | |
61 | 62 | | |
| 63 | + | |
62 | 64 | | |
63 | 65 | | |
64 | 66 | | |
| |||
73 | 75 | | |
74 | 76 | | |
75 | 77 | | |
| 78 | + | |
| 79 | + | |
76 | 80 | | |
77 | 81 | | |
78 | 82 | | |
| |||
0 commit comments