Skip to content

Commit f16ed4e

Browse files
Merge pull request AI-Hypercomputer#4208 from AI-Hypercomputer:fix/deepseek-batchsplit-context-reshard
PiperOrigin-RevId: 936252440
2 parents 0b9f604 + cd08cff commit f16ed4e

2 files changed

Lines changed: 18 additions & 2 deletions

File tree

src/maxtext/trainers/pre_train/train.py

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -221,8 +221,20 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr
221221
one_hot_targets = jax.nn.one_hot(data["targets"], config.vocab_size)
222222
xent, z_loss = max_utils.cross_entropy_with_logits(logits, one_hot_targets, z_loss=config.z_loss_multiplier)
223223

224-
xent = nn.with_logical_constraint(xent, ("activation_embed_and_logits_batch", "activation_length"))
225-
z_loss = nn.with_logical_constraint(z_loss, ("activation_embed_and_logits_batch", "activation_length"))
224+
xent = sharding.maybe_shard_with_logical(
225+
xent,
226+
("activation_embed_and_logits_batch", "activation_length"),
227+
model.mesh,
228+
config.shard_mode,
229+
debug_sharding=config.debug_sharding,
230+
)
231+
z_loss = sharding.maybe_shard_with_logical(
232+
z_loss,
233+
("activation_embed_and_logits_batch", "activation_length"),
234+
model.mesh,
235+
config.shard_mode,
236+
debug_sharding=config.debug_sharding,
237+
)
226238

227239
# Mask out paddings at the end of each example.
228240
xent = xent * (data["targets_segmentation"] != 0)

tests/unit/train_nnx_test.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323
import unittest
2424

2525
from flax import nnx
26+
import jax
2627
import jax.numpy as jnp
2728
from maxtext.common import train_state_nnx
2829
from maxtext.trainers.pre_train import train as pre_train
@@ -59,6 +60,7 @@ class _Cfg:
5960
record_internal_nn_metrics: bool = False
6061
skip_step_on_spikes: bool = False
6162
shard_mode: int = 0 # ShardMode.AUTO
63+
debug_sharding: bool = False
6264
weight_sparsity_n: int = 0
6365
weight_sparsity_m: int = 0
6466

@@ -73,6 +75,8 @@ class _TinyDecoder(nnx.Module):
7375
def __init__(self, vocab_size: int, hidden: int, rngs: nnx.Rngs):
7476
self.embed = nnx.Embed(vocab_size, hidden, rngs=rngs)
7577
self.proj = nnx.Linear(hidden, vocab_size, rngs=rngs)
78+
# loss_fn shards activations against model.mesh, so the stub needs one.
79+
self.mesh = jax.make_mesh((1, 1, 1, 1), ("data", "fsdp", "expert", "context"))
7680

7781
def __call__(
7882
self,

0 commit comments

Comments
 (0)