Skip to content
Merged
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
64 changes: 39 additions & 25 deletions src/maxtext/utils/mllog_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
"""Utils for MLPerf submission compliance."""

import atexit
import math
import os
import threading
import time
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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"),
Expand All @@ -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():
Expand Down
34 changes: 33 additions & 1 deletion tests/unit/mllog_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
Loading