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
33 changes: 30 additions & 3 deletions src/maxtext/utils/train_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -204,6 +204,30 @@ def jit_train_and_eval_step(
return p_train_step, p_eval_step


class _ReorderedDataIterator:
"""Applies a reorder function to each batch."""

def __init__(self, reorder_fn, data_iterator):
self.reorder_fn = reorder_fn
self.data_iterator = data_iterator

def __iter__(self):
return self

def __next__(self):
return self.reorder_fn(next(self.data_iterator))

def reset(self):
self.data_iterator.reset()


def _reorder_data_iterator_for_loader(reorder_fn, data_iterator):
"""Wraps the data iterator, or each element of an iterator list, with the reorder view."""
if isinstance(data_iterator, list):
return [_ReorderedDataIterator(reorder_fn, iterator) for iterator in data_iterator]
return _ReorderedDataIterator(reorder_fn, data_iterator)


def setup_train_loop(config, recorder, devices=None):
"""Set up prerequisites for the training loop -

Expand Down Expand Up @@ -278,6 +302,7 @@ def create_train_state_fn():
raise ValueError("Packing is only supported for load balanced ring attention with context parallelism for GPU.")

# Apply reordering wrapper to data iterators if context parallelism is enabled
data_iterator_for_loader = data_iterator
with jax.set_mesh(mesh):
if context_parallel_size > 1 and config.context_parallel_load_balance:

Expand All @@ -294,12 +319,14 @@ def create_train_state_fn():
reorder_fn = maxtext_utils.get_reorder_callable(
context_parallel_size, config.shard_mode, reorder_strategy, config.hardware
)
data_iterator = map(reorder_fn, data_iterator)
# data_iterator itself stays unwrapped because checkpointing dispatches on
# its concrete type; only batch consumers receive the reordered view.
data_iterator_for_loader = _reorder_data_iterator_for_loader(reorder_fn, data_iterator)
if eval_data_iterator:
eval_data_iterator = map(reorder_fn, eval_data_iterator)
eval_data_iterator = _ReorderedDataIterator(reorder_fn, eval_data_iterator)

# Create data_loader AFTER reordering wrapper is applied
data_loader = create_dataloader(config, mesh, data_iterator, recorder, rampup_manager)
data_loader = create_dataloader(config, mesh, data_iterator_for_loader, recorder, rampup_manager)

state, _, state_mesh_shardings, data_iterator, _ = maxtext_utils.setup_training_state(
data_iterator, config, mesh, checkpoint_manager, init_state_fn
Expand Down
23 changes: 20 additions & 3 deletions tests/integration/setup_train_loop_nnx_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@
from maxtext.common import train_state_nnx
from maxtext.configs import pyconfig
from maxtext.utils.globals import MAXTEXT_ASSETS_ROOT
from maxtext.utils.train_utils import setup_train_loop
from maxtext.utils import train_utils
from tests.utils.test_helpers import get_test_config_path
import pytest

Expand Down Expand Up @@ -83,7 +83,7 @@ def test_pure_nnx_setup_returns_train_state_nnx(self):
rampup_manager,
eval_data_iterator,
train_state,
) = setup_train_loop(config, recorder=None)
) = train_utils.setup_train_loop(config, recorder=None)

# The NNX path returns a fully-merged TrainStateNNX (lines 352-354 in train_utils.py).
self.assertIsInstance(train_state, train_state_nnx.TrainStateNNX)
Expand All @@ -109,13 +109,30 @@ def test_pure_nnx_setup_returns_train_state_nnx(self):
# flag them as unused — they're part of the public return contract.
del checkpoint_manager, rampup_manager, eval_data_iterator

def test_load_balanced_cp_keeps_checkpoint_iterator_unwrapped(self):
"""The reordered view goes to the DataLoader and eval; the original iterator goes to checkpointing."""
config = _tiny_nnx_pyconfig(
ici_context_parallelism=2,
context_parallel_load_balance=True,
packing=False,
eval_interval=1,
)

*_, data_iterator, data_loader, _, eval_data_iterator, _ = train_utils.setup_train_loop(config, recorder=None)

# pylint: disable=protected-access
self.assertIsInstance(data_loader.data_iterator, train_utils._ReorderedDataIterator)
self.assertIs(data_loader.data_iterator.data_iterator, data_iterator)
self.assertNotIsInstance(data_iterator, train_utils._ReorderedDataIterator)
self.assertIsInstance(eval_data_iterator, train_utils._ReorderedDataIterator)

def test_pure_nnx_setup_param_only_split_matches_model(self):
"""nnx.split(state.model, nnx.Param, ...) must yield a non-empty Param tree

whose structure matches state_mesh_shardings.model after the same split.
"""
config = _tiny_nnx_pyconfig()
*_, state_mesh_shardings, model, _, _, _, _, _, _, train_state = setup_train_loop(config, recorder=None)
*_, state_mesh_shardings, model, _, _, _, _, _, _, train_state = train_utils.setup_train_loop(config, recorder=None)

_, params, _ = nnx.split(train_state.model, nnx.Param, ...)
_, params_shardings, _ = nnx.split(state_mesh_shardings.model, nnx.Param, ...)
Expand Down
53 changes: 53 additions & 0 deletions tests/unit/train_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,59 @@ class MockConfig:
use_iota_embed: bool = False


class _FakeIterator:
"""Minimal resettable iterator for reorder-view tests."""

def __init__(self, batches):
self._batches = batches
self._index = 0
self.reset_calls = 0

def __next__(self):
batch = self._batches[self._index]
self._index += 1
return batch

def reset(self):
self.reset_calls += 1
self._index = 0


class TestReorderedDataIterator(unittest.TestCase):
"""Tests for the load-balanced CP reorder view."""

# pylint: disable=protected-access

def test_reorders_batches_and_forwards_reset(self):
inner = _FakeIterator([1, 2])
wrapped = train_utils._ReorderedDataIterator(lambda batch: batch * 10, inner)

self.assertEqual(next(wrapped), 10)
self.assertEqual(next(wrapped), 20)
self.assertIs(iter(wrapped), wrapped)
wrapped.reset()
self.assertEqual(inner.reset_calls, 1)
self.assertEqual(next(wrapped), 10)

def test_loader_view_wraps_single_iterator(self):
inner = _FakeIterator([1])
view = train_utils._reorder_data_iterator_for_loader(lambda batch: batch * 10, inner)

self.assertIsInstance(view, train_utils._ReorderedDataIterator)
self.assertIs(view.data_iterator, inner)

def test_loader_view_wraps_each_list_element(self):
inners = [_FakeIterator([1]), _FakeIterator([2])]
view = train_utils._reorder_data_iterator_for_loader(lambda batch: batch * 10, inners)

self.assertIsInstance(view, list)
self.assertEqual(len(view), 2)
for wrapped, inner in zip(view, inners):
self.assertIsInstance(wrapped, train_utils._ReorderedDataIterator)
self.assertIs(wrapped.data_iterator, inner)
self.assertEqual(next(view[1]), 20)


class TestValidateTrainConfig(unittest.TestCase):
"""Tests for validate_train_config."""

Expand Down
Loading