diff --git a/src/maxtext/utils/mllog_utils.py b/src/maxtext/utils/mllog_utils.py index 7433b9bf05..80bb3f2165 100644 --- a/src/maxtext/utils/mllog_utils.py +++ b/src/maxtext/utils/mllog_utils.py @@ -14,6 +14,7 @@ """Utils for MLPerf submission compliance.""" import atexit +import math import os import threading import time @@ -198,6 +199,18 @@ def _axis_product(config, *names) -> int: return product +def _rule_parallelism(config, logical_axis) -> int: + """Returns the mesh size (ICI x DCN) that `logical_axis` is sharded over.""" + sizes = { + axis: max(ici, 1) * max(dcn, 1) + for axis, ici, dcn in zip(config.mesh_axes, config.ici_parallelism, config.dcn_parallelism) + } + axes = dict(config.logical_axis_rules).get(logical_axis) or () + if isinstance(axes, str): + axes = (axes,) + return math.prod(sizes.get(axis, 1) for axis in axes) + + # Bits per element of the MLPerf pre-approved numerical formats, used to pick the lowest of several. _PRECISION_BITS = { "fp64": 64, @@ -258,30 +271,37 @@ def init_print(config): mllogger.event("target_accuracy", config.target_eval_loss) # MLPerf v6.1 mandatory precision, parallelism, micro-batch size, and config filename disclosure. + use_lineage = getattr(config, "use_lineage", False) + tensor_parallelism = _axis_product( + config, + "ici_tensor_parallelism", + "dcn_tensor_parallelism", + "ici_tensor_sequence_parallelism", + "dcn_tensor_sequence_parallelism", + ) + expert_parallelism = _axis_product(config, "ici_expert_parallelism", "dcn_expert_parallelism") + quantization = getattr(config, "quantization", None) + token_all_gather_quantized = quantization and getattr(config, "moe_quantize_token_all_gather", False) + if use_lineage: + # Lineage runs on a physical mesh where every ici_*_parallelism field is 1. + # Its MLA projections are head-sharded across the `activation_length` + # (TensorCore) axis while attention sees the full sequence: TP + SP, not + # CP. fp8_full also quantizes the EP token all-gather. + expert_parallelism = _rule_parallelism(config, "exp") + tensor_parallelism = _rule_parallelism(config, "activation_length") + quantization = quantization or config.lineage_quantization + token_all_gather_quantized = True dtype_str = _mllog_precision(getattr(config, "dtype", "bfloat16")) - linear_prec = _mllog_precision(getattr(config, "quantization", None), fallback=dtype_str) + linear_prec = _mllog_precision(quantization, fallback=dtype_str) comm_precisions = [_mllog_precision(getattr(config, "grad_dtype", None), fallback=dtype_str)] - if ( - getattr(config, "moe_quantize_token_all_gather", False) - and getattr(config, "quantization", "") - and _axis_product(config, "ici_expert_parallelism", "dcn_expert_parallelism") > 1 - ): - # The ring-of-experts EP all-gather sends tokens quantized with the GMM activation qtype (fp8 for fp8_full). + if token_all_gather_quantized and expert_parallelism > 1: + # The EP token all-gather sends tokens in the GMM activation qtype. comm_precisions.append(linear_prec) comm_prec = _lowest_precision(*comm_precisions) mllogger.event(mllog.constants.LOWEST_NUMERICAL_PRECISION_IN_LINEAR, linear_prec) mllogger.event(mllog.constants.LOWEST_NUMERICAL_PRECISION_IN_ATTN, dtype_str) mllogger.event(mllog.constants.LOWEST_NUMERICAL_PRECISION_IN_COMM, comm_prec) - mllogger.event( - mllog.constants.TENSOR_PARALLELISM, - _axis_product( - config, - "ici_tensor_parallelism", - "dcn_tensor_parallelism", - "ici_tensor_sequence_parallelism", - "dcn_tensor_sequence_parallelism", - ), - ) + mllogger.event(mllog.constants.TENSOR_PARALLELISM, tensor_parallelism) mllogger.event( mllog.constants.PIPELINE_PARALLELISM, _axis_product(config, "ici_pipeline_parallelism", "dcn_pipeline_parallelism"), @@ -290,19 +310,13 @@ def init_print(config): mllog.constants.CONTEXT_PARALLELISM, _axis_product(config, "ici_context_parallelism", "dcn_context_parallelism"), ) - mllogger.event( - mllog.constants.EXPERT_PARALLELISM, - _axis_product(config, "ici_expert_parallelism", "dcn_expert_parallelism"), - ) + mllogger.event(mllog.constants.EXPERT_PARALLELISM, expert_parallelism) # TPU v7x exposes 2 JAX devices (TensorCores) per chip, while system descriptions count chips. mllogger.event( mllog.constants.MICRO_BATCH_SIZE, max(1, int(round(getattr(config, "per_device_batch_size", 1) * 2))), ) - mllogger.event( - mllog.constants.CONFIG_FILENAME, - getattr(config, "mllog_config_filename", "") or "config.yml", - ) + mllogger.event(mllog.constants.CONFIG_FILENAME, f"{config.model_name}.yml") def init_stop(): diff --git a/tests/unit/mllog_utils_test.py b/tests/unit/mllog_utils_test.py index bdb5f2f229..a324d0905b 100644 --- a/tests/unit/mllog_utils_test.py +++ b/tests/unit/mllog_utils_test.py @@ -118,6 +118,7 @@ def make_config(**overrides): "enable_mllog": True, "mllog_file": "", "run_name": "unit-test-run", + "model_name": "deepseek3-671b", "data_shuffle_seed": 1234, "steps": 12000, "global_batch_size_to_train_on": 16384, @@ -296,7 +297,38 @@ def test_init_print_emits_required_keys(self): self.assertEqual(self.mllogger.value_of(_CONSTANTS.LOWEST_NUMERICAL_PRECISION_IN_COMM), "bfloat16") self.assertEqual(self.mllogger.value_of(_CONSTANTS.EXPERT_PARALLELISM), 8) self.assertEqual(self.mllogger.value_of(_CONSTANTS.MICRO_BATCH_SIZE), 2) - self.assertEqual(self.mllogger.value_of(_CONSTANTS.CONFIG_FILENAME), "config.yml") + self.assertEqual(self.mllogger.value_of(_CONSTANTS.CONFIG_FILENAME), "deepseek3-671b.yml") + + def test_init_print_lineage(self): + """Lineage EP/TP come from its axis rules, fp8 from lineage_quantization.""" + config = self.setup_local( + model_name="deepseek3-671b-lineage", + use_lineage=True, + quantization="", + lineage_quantization="fp8_full", + grad_dtype="bfloat16", + ici_expert_parallelism=1, + # The 8192-chip mesh of deepseek3-671b-lineage.yml + the launcher. + mesh_axes=["dcn", "x", "y", "z", "core"], + ici_parallelism=[1, 4, 4, 64, 2], + dcn_parallelism=[8, 1, 1, 1, 1], + logical_axis_rules=( + ("exp", ("x", "y", "core")), + ("activation_length", ("core",)), + ), + ) + mllog_utils.init_print(config) + + expected = { + _CONSTANTS.EXPERT_PARALLELISM: 32, + _CONSTANTS.TENSOR_PARALLELISM: 2, + _CONSTANTS.CONTEXT_PARALLELISM: 1, + _CONSTANTS.LOWEST_NUMERICAL_PRECISION_IN_LINEAR: "fp8", + _CONSTANTS.LOWEST_NUMERICAL_PRECISION_IN_COMM: "fp8", + _CONSTANTS.CONFIG_FILENAME: "deepseek3-671b-lineage.yml", + } + for key, value in expected.items(): + self.assertEqual(self.mllogger.value_of(key), value, key) def test_init_print_comm_precision_includes_quantized_token_all_gather(self): config = self.setup_local(moe_quantize_token_all_gather=True)