Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 83 additions & 4 deletions src/maxtext/training_engine/maxtext_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -541,6 +541,8 @@ def __init__(
self._compiled_eval: Any = None
self._compiled_eval_signature: Any = None
self._signature_compare_warned: bool = False
# True when the pure mirror has advanced past the live NNX objects; see `_publish_to_live`.
self._live_stale: bool = False
if not training_config.model_name:
raise ValueError("training_config.model_name must be specified")
model_or_model_mesh_pair = model_creation_utils.from_pretrained(
Expand Down Expand Up @@ -608,6 +610,7 @@ def __init__(
@property
def model(self) -> Any:
"""Returns the NNX model instance."""
self._sync_live_objects()
return self._model

@model.setter
Expand All @@ -626,6 +629,7 @@ def model(self, new_model: Any) -> None:
@property
def optimizer(self) -> Any:
"""Returns the NNX optimizer instance."""
self._sync_live_objects()
return self._optimizer

@optimizer.setter
Expand Down Expand Up @@ -654,6 +658,7 @@ def train_step(self, step: int) -> None:
@property
def state(self) -> Any:
"""Returns the current train state, initializing it if necessary."""
self._sync_live_objects()
if self._state is None and self._model is not None and self._optimizer is not None:
self._state = train_state_nnx.TrainStateNNX(self._model, self._optimizer)
return self._state
Expand Down Expand Up @@ -747,11 +752,66 @@ def _invalidate_pure_state(self) -> None:

For the three ways the live NNX variables get replaced behind the engine's back: the
`model`/`optimizer`/`state` setters, and a checkpoint restore.

Flushes a deferred publish before forgetting the mirror it would have been published
from, so dropping the cache never drops a step's results with it.
"""
self._sync_live_objects()
self._params_pure = None
self._rest_pure = None
self._state_pure = None

def _publish_to_live(self, target: Any, pure: Any) -> None:
"""Writes `pure` into the live NNX objects, or defers it when the mirror can do it.

`nnx.update` walks the module graph, and its cost is the graph's size rather than the
state's: 6.2 ms to publish 180 leaves of non-parameter state into an unrolled 28-layer
qwen3-0.6b, once per *micro-batch*, and 16.7 ms for the parameters once per update.
That was the last per-step graph walk left after `_refresh_pure_state` removed the
`nnx.split`s.

Nothing inside the step path reads it back -- `_read_model_pure` and `_read_state_pure`
answer from the pure mirror, and the live objects are used only as containers -- so the
walk is pure publication, for readers outside the engine. It is deferred to
`_sync_live_objects`, which the property getters and every method that reads real
values call first.

Deferral is only possible while the mirror is authoritative. With the pure state
disabled there is nothing to publish from later, so the write happens now, as it always
did.
"""
if self._params_pure is None:
nnx.update(target, pure)
return
self._live_stale = True

def _sync_live_objects(self) -> None:
"""Brings `self._model`/`self._state` up to date with the pure mirror.

Every caller that reads values off the live NNX objects -- the `model`, `optimizer` and
`state` properties, checkpointing, weight sync -- goes through here first. It is also
the reason a deferred publish cannot be observed: the arrays a reader would otherwise
find are not merely stale but *deleted*, since `update()` donates the state it passed
to the update kernel.

One write covers both deferral sites. `fwd_bwd` defers the model's non-parameter state
and `update` the whole train state, but `_publish_model_rest` has already folded the
former into `_state_pure`, and the live model is the same object the state holds under
`_MODEL_STATE_KEY` -- so publishing the state publishes the model with it.
"""
if not self._live_stale:
return
# Cleared first: `nnx.update` cannot re-enter this, but a raising one must not leave the
# flag set and re-run the same failed write on the next read.
self._live_stale = False
# `_state_pure` is what `_publish_to_live` gated its deferral on -- it checks
# `_params_pure`, and the three move together, being set as a group by
# `_refresh_pure_state` and `_publish_state` and cleared as one by
# `_invalidate_pure_state`. Should they ever stop moving together, this drops the write
# the flag promised, so keep them set and cleared as a group.
if self._state is not None and self._state_pure is not None:
nnx.update(self._state, self._state_pure)

def _disable_pure_state(self, reason: str) -> None:
"""Falls back to re-splitting the module graph on every step, saying so once."""
self._invalidate_pure_state()
Expand Down Expand Up @@ -796,9 +856,11 @@ def _refresh_pure_state(self) -> None:
Once per compile rather than once per step, which is the point: the two `nnx.split`
calls the step path used to make walked 1756 graph nodes for 92 ms of a 283 ms step on
an unrolled qwen3-0.6b, against 0.84 ms for `nnx.split_state` over the flat state.
Publication is unchanged -- `fwd_bwd` and `update` still `nnx.update` the live objects
where they always did, so `self.model` and `self.state` are never stale.
Publication out of the step path is deferred rather than eager -- see
`_publish_to_live` -- so this syncs first, since it re-reads the live objects it is
about to split.
"""
self._sync_live_objects()
if self._state is None:
self._state = train_state_nnx.TrainStateNNX(self._model, self._optimizer)
model = getattr(self._state, _MODEL_STATE_KEY, self._model)
Expand Down Expand Up @@ -1307,6 +1369,10 @@ def accum_kernel(params, rest, dynamic, acc_grads, acc_denom):

def _compile_eval_for_batch(self, dynamic_batch: Any, static_batch: dict[str, Any]) -> None:
"""JIT-compiles the forward-only eval kernel for one batch structure."""
# Its caller has synced already, but this splits the live model rather than reading the
# mirror, so it does not get to assume that: the shardings below are taken off the
# leaves it finds, and a deferred publish leaves those holding donated arrays.
self._sync_live_objects()
self._model_graphdef, params_pure, rest_pure = nnx.split(self._model, nnx.Param, ...)

def kernel(params, rest, dynamic):
Expand Down Expand Up @@ -1406,8 +1472,8 @@ def fwd_bwd(self, payload: abstract_engine.TrainerPayload, **kwargs: Any) -> Non
loss, aux, new_rest, acc_grads, acc_denom = self._fwd_bwd_kernel(
params, rest, batch, self._accumulated_grads, self._accumulated_denominator
)
nnx.update(model, new_rest)
self._publish_model_rest(new_rest)
self._publish_to_live(model, new_rest)

# Don't add metrics to the throttler queue because metrics are logged after
# the update step.
Expand Down Expand Up @@ -1474,8 +1540,8 @@ def update(self, **kwargs: Any) -> int:
new_state_pure, grad_norm, is_skipped = self._update_kernel(
state_pure, self._accumulated_grads, self._accumulated_denominator, mean_loss
)
nnx.update(self._state, new_state_pure)
self._publish_state(new_state_pure)
self._publish_to_live(self._state, new_state_pure)

if grad_norm is not None:
self.record_metrics("gradient_norm", grad_norm)
Expand All @@ -1487,6 +1553,10 @@ def update(self, **kwargs: Any) -> int:
# entries until it pops them, which pinned three parameter trees, and once the state is
# donated a late pop raises "Array has been deleted". The norm comes out of the same
# executable, so its readiness still means the update landed. Tunix v2 does the same.
if grad_norm is None:
# The fallback below reads leaves off the live objects, which a deferred publish
# leaves holding the arrays `update()` just donated.
self._sync_live_objects()
self._throttler.add_computation(
grad_norm if grad_norm is not None else (self._state if self._state is not None else self._model),
self._metrics_recorder.get_step_metrics(self.train_step),
Expand Down Expand Up @@ -1540,6 +1610,12 @@ def eval_step(self, payload: abstract_engine.TrainerPayload, **kwargs: Any) -> N
"""
batch = self._prepare_batch(payload)

# Eval is the one step path that reads real values off the live model instead of the
# pure mirror, so it is also the one that has to flush a deferred publish. Skipping it
# does not evaluate stale weights, it raises: `update()` donates the state it passes to
# the update kernel, so the parameters split out below would be deleted arrays.
self._sync_live_objects()

model = getattr(self._state, "model", self._model) if self._state is not None else self._model
if not isinstance(model, nnx.Module):
raise TypeError("MaxTextTrainingEngine requires an NNX model (flax.nnx.Module), got" f" {type(model).__name__}")
Expand Down Expand Up @@ -1594,6 +1670,7 @@ def save_checkpoint(self, metadata: Any, **kwargs: Any) -> None:
metadata: Checkpoint metadata payload from Orchestrator.
**kwargs: Additional checkpoint saving options.
"""
self._sync_live_objects()
# Drain all inflight computations and log pending metrics before checkpointing.
self._throttler.wait_for_all()

Expand Down Expand Up @@ -1651,6 +1728,7 @@ def restore_checkpoint(self, **kwargs: Any) -> Any:
Returns:
The metadata PyTree of the restored checkpoint.
"""
self._sync_live_objects()
step = kwargs.get("step", None)
checkpoint_state = checkpointing.CheckpointState(
model=self.model,
Expand Down Expand Up @@ -1847,6 +1925,7 @@ def get_metrics(self, clear_cache: bool = True) -> abstract_engine.MetricsBuffer

def _get_trainable_params_state(self) -> Any:
"""Extracts pure parameter weights from the model, excluding optimizer and RNG state."""
self._sync_live_objects()
model = getattr(self._state, "model", None) if self._state is not None else self._model
if isinstance(model, nnx.Module):
return nnx.state(model, nnx.Param)
Expand Down
94 changes: 94 additions & 0 deletions tests/post_training/unit/maxtext_engine_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -1349,6 +1349,100 @@ def test_perplexity_is_emitted_alongside_the_loss(self):
self.assertIn("perplexity", processed)
self.assertAlmostEqual(processed["perplexity"], float(np.exp(6.0)), places=3)

def test_live_objects_are_not_walked_on_the_step_path(self):
"""The step path publishes to the pure mirror only; `nnx.update` is deferred.

The saving this buys is the whole point of the deferral -- `nnx.update` walks the
module graph, and its cost scales with the graph rather than the state written -- so a
step that walks it anyway has silently lost it.

Only writes aimed at the engine's own live objects count. `nnx.Optimizer.update` calls
`nnx.update` internally, on the state the kernel merged locally, and that is neither
publication nor something this change can avoid.
"""
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(self._weighted_loss_fn)
t.compile(DummyPayload())
t.fwd_bwd(DummyPayload())
t.update()

live = (t._model, t._state, t._optimizer)
walked = []
real_update = maxtext_engine.nnx.update

def spy(target, *args, **kwargs):
if any(target is obj for obj in live):
walked.append(target)
return real_update(target, *args, **kwargs)

with mock.patch.object(maxtext_engine.nnx, "update", spy):
t.fwd_bwd(DummyPayload())
t.update()

self.assertEmpty(walked, "the step path wrote to the live NNX objects")
self.assertTrue(t._live_stale, "nothing was deferred, so the publish was skipped, not deferred")

def test_reading_the_model_flushes_a_deferred_publish(self):
"""A deferred publish is invisible: every public reader syncs before it answers."""
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(self._weighted_loss_fn)
t.compile(DummyPayload())
t.fwd_bwd(DummyPayload())
t.update()
self.assertTrue(t._live_stale, "nothing was deferred, so this test proves nothing")

published = nnx.state(t.model, nnx.Param)
self.assertFalse(t._live_stale)
jax.tree.map(np.testing.assert_array_equal, jax.tree.leaves(published), jax.tree.leaves(t._params_pure))

def test_eval_after_update_flushes_a_deferred_publish(self):
"""`eval_step` reads the live model rather than the mirror, so it syncs like a getter.

The eval path is the exception to "nothing inside the step path reads the live objects
back", and it is the one place the deferral can be observed. Not as stale weights
either: `update()` donates the state it hands the update kernel, so an unsynced
`eval_step` splits deleted arrays out of the live model and raises inside the kernel.

Ordered after an `update()` on purpose -- the pre-existing eval tests call `eval_step`
on a freshly compiled engine, which never defers anything.
"""
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(self._weighted_loss_fn)
t.compile(DummyPayload())
t.fwd_bwd(DummyPayload())
t.update()
self.assertTrue(t._live_stale, "nothing was deferred, so this test proves nothing")

t.eval_step(DummyPayload())
self.assertFalse(t._live_stale)

def test_dropping_the_pure_state_flushes_rather_than_drops_the_step(self):
"""Invalidating the mirror must not discard an update that was only deferred."""
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(self._weighted_loss_fn)
t.compile(DummyPayload())
t.fwd_bwd(DummyPayload())
t.update()
expected = jax.tree.leaves(t._params_pure)

t._invalidate_pure_state()
self.assertFalse(t._live_stale)
jax.tree.map(np.testing.assert_array_equal, jax.tree.leaves(nnx.state(t._model, nnx.Param)), expected)

def test_publication_stays_eager_without_a_pure_mirror(self):
"""With the pure state disabled there is nothing to publish from later."""
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(self._weighted_loss_fn)
t.compile(DummyPayload())
t.fwd_bwd(DummyPayload())
t.update()
t._disable_pure_state("test")

with mock.patch.object(maxtext_engine.nnx, "update", wraps=maxtext_engine.nnx.update) as spy:
t._publish_to_live(t._model, nnx.state(t._model, nnx.Param))
self.assertEqual(spy.call_count, 1)
self.assertFalse(t._live_stale)


if __name__ == "__main__":
absltest.main()
Loading