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
39 changes: 39 additions & 0 deletions docs/models.md
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,45 @@ model, model_path = AutoModel.from_pretrained(
)
```

##### Pipeline parallel MaxText models

Tunix can use MaxText's pipeline runtime for trainer forward and backward
passes. Use `MaxTextPipelineConfig` to keep the MaxText config and JAX mesh in
sync. The helper builds a mesh with MaxText's `stage` and `tensor` axis names;
a native Tunix mesh using an axis such as `tp` does not enable MaxText pipeline
parallelism.

For example, this layout assigns the 36 Qwen3-8B decoder layers to two pipeline
stages, with tensor parallelism degree four inside every stage:

```python
from tunix.models import automodel
from tunix.models import maxtext_parallelism

parallelism = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=2,
tensor_parallelism=4,
num_layers_per_pipeline_stage=18,
num_pipeline_microbatches=4,
pipeline_parallel_layers=36,
)
mesh = parallelism.create_mesh()
parallelism.validate_batch_size(global_batch_size=16)

model, model_path = automodel.AutoModel.from_pretrained(
model_id="Qwen/Qwen3-8B",
mesh=mesh,
model_source=automodel.ModelSource.MAXTEXT,
model_path="gs://my-bucket/qwen3-8b-maxtext-checkpoint",
maxtext_pipeline_config=parallelism,
per_device_batch_size=2,
)
```

On eight devices, another valid Qwen3-8B layout is PP4 x TP2 with nine layers
per stage and eight pipeline microbatches. Pipeline and tensor degrees multiply:
PP8 x TP8 therefore requires 64 devices rather than eight.

To load a Maxtext model when launching training via a shell script, append the corresponding override arguments directly to the script execution:

```bash
Expand Down
87 changes: 87 additions & 0 deletions tests/models/automodel_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from absl.testing import parameterized
import jax
from tunix.models import automodel
from tunix.models import maxtext_parallelism
from tunix.models import naming


Expand Down Expand Up @@ -212,6 +213,91 @@ class MockMaxTextConfig:
mock_config, mesh=mock_mesh, wrap_with_tunix_adapter=True
)

@mock.patch(
"tunix.models.automodel.download_model",
return_value="qwen3-8b",
)
def test_from_pretrained_maxtext_pipeline_config(self, mock_download):
del mock_download
m_maxtext = types.ModuleType("maxtext")
m_configs = types.ModuleType("maxtext.configs")
m_pyconfig = types.ModuleType("maxtext.configs.pyconfig")
m_types = types.ModuleType("maxtext.configs.types")
m_utils = types.ModuleType("maxtext.utils")
m_creation = types.ModuleType("maxtext.utils.model_creation_utils")

config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=2,
tensor_parallelism=4,
num_layers_per_pipeline_stage=18,
num_pipeline_microbatches=4,
pipeline_parallel_layers=36,
)
mesh = types.SimpleNamespace(
shape=dict(
zip(
maxtext_parallelism.MAXTEXT_MESH_AXIS_NAMES,
config.mesh_axis_shapes,
)
)
)

class MockMaxTextConfig:
model_fields = {
"skip_jax_distributed_system": True,
**{key: True for key in config.as_maxtext_kwargs()},
}

mock_config = mock.MagicMock()
m_pyconfig.initialize = mock.MagicMock(return_value=mock_config)
m_creation.from_pretrained = mock.MagicMock()
m_types.MaxTextConfig = MockMaxTextConfig
m_pyconfig.__file__ = "/opt/maxtext/configs/pyconfig.py"
m_configs.pyconfig = m_pyconfig
m_configs.types = m_types
m_utils.model_creation_utils = m_creation
m_maxtext.configs = m_configs
m_maxtext.utils = m_utils

with mock.patch.dict(
"sys.modules",
{
"maxtext": m_maxtext,
"maxtext.configs": m_configs,
"maxtext.configs.pyconfig": m_pyconfig,
"maxtext.configs.types": m_types,
"maxtext.utils": m_utils,
"maxtext.utils.model_creation_utils": m_creation,
},
):
automodel.AutoModel.from_pretrained(
"Qwen/Qwen3-8B",
mesh=mesh,
model_source=automodel.ModelSource.MAXTEXT,
maxtext_pipeline_config=config,
)

called_argv = m_pyconfig.initialize.call_args[0][0]
self.assertEqual(called_argv[1], "/opt/maxtext/configs/base.yml")
self.assertIn("ici_pipeline_parallelism=2", called_argv)
self.assertIn("ici_tensor_parallelism=4", called_argv)
self.assertIn("num_layers_per_pipeline_stage=18", called_argv)
self.assertIn("num_pipeline_microbatches=4", called_argv)

def test_rejects_maxtext_pipeline_config_for_native_model(self):
config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=2,
tensor_parallelism=4,
)

with self.assertRaisesRegex(ValueError, "only supported"):
automodel.AutoModel.from_pretrained(
"Qwen/Qwen3-8B",
mesh=mock.MagicMock(),
model_source=automodel.ModelSource.HUGGINGFACE,
maxtext_pipeline_config=config,
)

@parameterized.named_parameters(*_get_all_models_test_parameters())
def test_obtain_model_params_valid(self, model_name: str):
automodel.call_model_config(model_name)
Expand Down Expand Up @@ -433,5 +519,6 @@ class FakeConfig:
self.assertEqual(called_config.flash_attention_block_size, 512)
self.assertFalse(hasattr(called_config, "invalid_param"))


if __name__ == "__main__":
absltest.main()
168 changes: 168 additions & 0 deletions tests/models/maxtext_parallelism_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,168 @@
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

import collections
import types

from absl.testing import absltest
from absl.testing import parameterized
from tunix.models import maxtext_parallelism


class MaxTextPipelineConfigTest(parameterized.TestCase):

@parameterized.named_parameters(
dict(
testcase_name="pp2_tp4_qwen3_8b",
pipeline_parallelism=2,
tensor_parallelism=4,
layers_per_stage=18,
microbatches=4,
expected_shapes=(1, 1, 2, 1, 1, 1, 1, 4, 1, 1, 1, 1),
),
dict(
testcase_name="pp4_tp2_qwen3_8b",
pipeline_parallelism=4,
tensor_parallelism=2,
layers_per_stage=9,
microbatches=8,
expected_shapes=(1, 1, 4, 1, 1, 1, 1, 2, 1, 1, 1, 1),
),
)
def test_hybrid_layout(
self,
pipeline_parallelism,
tensor_parallelism,
layers_per_stage,
microbatches,
expected_shapes,
):
config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=pipeline_parallelism,
tensor_parallelism=tensor_parallelism,
num_layers_per_pipeline_stage=layers_per_stage,
num_pipeline_microbatches=microbatches,
pipeline_parallel_layers=36,
)

self.assertEqual(config.required_device_count, 8)
self.assertEqual(config.mesh_axis_shapes, expected_shapes)
self.assertEqual(
config.as_maxtext_kwargs()["ici_pipeline_parallelism"],
pipeline_parallelism,
)
self.assertEqual(
config.as_maxtext_kwargs()["ici_tensor_parallelism"],
tensor_parallelism,
)
config.validate_batch_size(16)

def test_mesh_axes_match_maxtext_default_order(self):
self.assertEqual(
maxtext_parallelism.MAXTEXT_MESH_AXIS_NAMES,
(
"diloco",
"data",
"stage",
"fsdp",
"fsdp_transpose",
"context",
"context_autoregressive",
"tensor",
"tensor_transpose",
"tensor_sequence",
"expert",
"autoregressive",
),
)

def test_validate_exact_mesh(self):
config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=2,
tensor_parallelism=4,
num_layers_per_pipeline_stage=18,
num_pipeline_microbatches=4,
)
mesh = types.SimpleNamespace(
shape=collections.OrderedDict(
zip(
maxtext_parallelism.MAXTEXT_MESH_AXIS_NAMES,
config.mesh_axis_shapes,
)
)
)

config.validate_mesh(mesh)

def test_rejects_tunix_tp_axis_names(self):
config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=2,
tensor_parallelism=4,
)
mesh = types.SimpleNamespace(
shape=collections.OrderedDict((("fsdp", 1), ("tp", 8)))
)

with self.assertRaisesRegex(ValueError, "Missing axes"):
config.validate_mesh(mesh)

@parameterized.named_parameters(
dict(
testcase_name="one_pipeline_stage",
kwargs={"pipeline_parallelism": 1},
error="at least 2",
),
dict(
testcase_name="microbatches_not_divisible",
kwargs={
"pipeline_parallelism": 4,
"num_pipeline_microbatches": 6,
},
error="must be divisible",
),
dict(
testcase_name="layers_not_divisible",
kwargs={
"pipeline_parallelism": 4,
"num_layers_per_pipeline_stage": 3,
"pipeline_parallel_layers": 28,
},
error="must be divisible",
),
dict(
testcase_name="conflicting_fsdp_modes",
kwargs={
"pipeline_parallelism": 2,
"pipeline_fsdp_ag_once": True,
"pipeline_fsdp_ag_per_repeat": True,
},
error="mutually exclusive",
),
)
def test_invalid_config(self, kwargs, error):
with self.assertRaisesRegex(ValueError, error):
maxtext_parallelism.MaxTextPipelineConfig(**kwargs)

def test_rejects_incompatible_batch_size(self):
config = maxtext_parallelism.MaxTextPipelineConfig(
pipeline_parallelism=4,
num_pipeline_microbatches=8,
)

with self.assertRaisesRegex(ValueError, "global_batch_size"):
config.validate_batch_size(12)


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