diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 227c7e36f1..10ef89d87c 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -809,6 +809,7 @@ olmo_apply_ngram_filter: true # mask instances with repetitive n-grams (OLMo-cor # Training loop steps: 150_001 # If set to -1 then will inherit value from learning_rate_schedule_steps log_period: 100 # The frequency of Tensorboard flush, gcs metrics writing, and managed profiler metrics updating. +max_inflight_computations: 2 # Maximum number of inflight computations on device. jax_distributed_initialization_timeout: 300 # This is the default timeout in https://github.com/jax-ml/jax/blob/main/jax/_src/distributed.py # Note there are two separate initializations - the jax coordination service (aka jax.distributed.initialize) and the backend (e.g. PjRT), the timeout above refers diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 4865557b5a..d63495fa55 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -1690,6 +1690,7 @@ class TrainingLoop(BaseModel): enable_data_shuffling: bool = Field(True, description="Enables shuffling of the training data.") data_shuffle_seed: int = Field(0, description="Seed for data shuffling.") init_weights_seed: int = Field(0, description="Seed for model weight initialization.") + max_inflight_computations: int = Field(2, description="Maximum number of inflight computations on device.") class ManifoldConstrainedHyperConnections(BaseModel): diff --git a/src/maxtext/training_engine/checkpointing.py b/src/maxtext/training_engine/checkpointing.py index 675cd7838e..32493c111a 100644 --- a/src/maxtext/training_engine/checkpointing.py +++ b/src/maxtext/training_engine/checkpointing.py @@ -54,9 +54,9 @@ def __init__( self._checkpoint_manager = ocp.CheckpointManager( directory=checkpoint_dir, options=ocp.CheckpointManagerOptions( - save_interval_steps=getattr(config, "checkpoint_period", 1), - max_to_keep=getattr(config, "max_num_checkpoints_to_keep", None), - enable_async_checkpointing=getattr(config, "async_checkpointing", True), + save_interval_steps=config.checkpoint_period, + max_to_keep=config.max_num_checkpoints_to_keep, + enable_async_checkpointing=config.async_checkpointing, ), ) diff --git a/src/maxtext/training_engine/inflight_throttler.py b/src/maxtext/training_engine/inflight_throttler.py index 1ed1256d5c..dac7063f05 100644 --- a/src/maxtext/training_engine/inflight_throttler.py +++ b/src/maxtext/training_engine/inflight_throttler.py @@ -32,8 +32,7 @@ def __init__(self, config: pyconfig.HyperParameters): Args: config: The training configuration. """ - max_inflight = getattr(config, "max_inflight_computations", 2) - self._inflight_queue = queue.Queue[Any](maxsize=max_inflight) + self._inflight_queue = queue.Queue[Any](maxsize=config.max_inflight_computations) self._metrics_logger = metrics_module.MetricsLogger(config=config) def add_computation(self, computation: Any, metrics: abstract_engine.MetricsBuffer | None) -> None: diff --git a/src/maxtext/training_engine/maxtext_engine.py b/src/maxtext/training_engine/maxtext_engine.py index c7a9bcef0b..9bc9345d5c 100644 --- a/src/maxtext/training_engine/maxtext_engine.py +++ b/src/maxtext/training_engine/maxtext_engine.py @@ -67,11 +67,11 @@ def __init__( ) self._config = training_config self._mesh = mesh - self._init_rng = jax.random.PRNGKey(getattr(training_config, "init_weights_seed", 0)) + self._init_rng = jax.random.PRNGKey(training_config.init_weights_seed) self._loss_fn: Callable[..., Any] | None = None self._gen_model_input_fn: Callable[[Any], dict[str, Any]] | None = None self._compiled = False - if not getattr(training_config, "model_name", None): + if not training_config.model_name: raise ValueError("training_config.model_name must be specified") self._model = model_creation_utils.from_pretrained( config=self._config, @@ -87,7 +87,7 @@ def __init__( self._train_step: int = 0 self._checkpoint_manager = checkpointing.CheckpointManager( - checkpoint_dir=getattr(self._config, "checkpoint_dir", getattr(self._config, "checkpoint_directory", "")), + checkpoint_dir=self._config.checkpoint_dir, config=self._config, ) self._metrics_recorder = metrics_module.MetricsRecorder() @@ -223,11 +223,7 @@ def diff_wrapper(p, r, b): micro_grads = jax.tree.map(lambda g: g * scale, micro_grads) micro_grads = jax.tree.map( - lambda x: ( - x.astype(getattr(self._config, "grad_dtype", jnp.float32)) - if hasattr(x, "dtype") and x.dtype == jnp.float32 - else x - ), + lambda x: (x.astype(self._config.grad_dtype) if hasattr(x, "dtype") and x.dtype == jnp.float32 else x), micro_grads, ) @@ -248,11 +244,11 @@ def _update_kernel(self, state_pure, accumulated_grads, micro_step_count, mean_l lambda g: g / micro_step_count, accumulated_grads, ) - if getattr(self._config, "gradient_clipping_threshold", 0.0) > 0: + if self._config.gradient_clipping_threshold > 0: grads = maxtext_utils.apply_gradient_clipping(grads, None, self._config.gradient_clipping_threshold) local_state = nnx.merge(self._state_graphdef, state_pure, copy=True) if hasattr(local_state, "apply_gradients"): - if getattr(self._config, "skip_step_on_spikes", False): + if self._config.skip_step_on_spikes: grad_norm = max_utils.l2norm_pytree(grads) local_state.apply_gradients(grads, loss=mean_loss, grad_norm=grad_norm) opt_obj = getattr(local_state, "optimizer", self._optimizer) diff --git a/src/maxtext/training_engine/metrics.py b/src/maxtext/training_engine/metrics.py index 91ffb8b54f..0739ce72b7 100644 --- a/src/maxtext/training_engine/metrics.py +++ b/src/maxtext/training_engine/metrics.py @@ -162,7 +162,7 @@ def __init__(self, config: pyconfig.HyperParameters): """ self._tb_writer = None - if getattr(config, "enable_tensorboard", False): + if config.enable_tensorboard: self._tb_writer = max_utils.initialize_summary_writer( config.tensorboard_dir, config.run_name, config.enable_tensorboard ) diff --git a/src/maxtext/utils/gradient_accumulation.py b/src/maxtext/utils/gradient_accumulation.py index 106baee07d..985dfdcb66 100644 --- a/src/maxtext/utils/gradient_accumulation.py +++ b/src/maxtext/utils/gradient_accumulation.py @@ -177,9 +177,7 @@ def reshape_to_microbatch_accumulations(batch_arr): raw_grads = jax.tree.map(_maybe_shard_with_name, raw_grads, unreduced_shardings) raw_grads = jax.tree.map(_maybe_shard_with_name, raw_grads, params_shardings) divisor = ( - config.gradient_accumulation_steps - if getattr(config, "use_tunix_gradient_accumulation", False) - else grad_and_loss["total_weights"] + config.gradient_accumulation_steps if config.use_tunix_gradient_accumulation else grad_and_loss["total_weights"] ) raw_grads = jax.tree_util.tree_map(lambda arr: arr / divisor, raw_grads) aux = jax.tree.map(lambda x: jnp.sum(x, axis=0), aux) # pytype: disable=module-attr diff --git a/tests/end_to_end/tpu/compare_training_engine.py b/tests/end_to_end/tpu/compare_training_engine.py index 9acb39d5f8..3e7c3b5b7a 100644 --- a/tests/end_to_end/tpu/compare_training_engine.py +++ b/tests/end_to_end/tpu/compare_training_engine.py @@ -68,7 +68,7 @@ def get_tpu_mesh( cfg: pyconfig.HyperParameters | None = None, ) -> jax.sharding.Mesh: """Returns SPMD device mesh based on configuration or default 1-device TPU mesh.""" - if cfg is not None and getattr(cfg, "model_name", "simple_mlp") != "simple_mlp": + if cfg is not None and cfg.model_name != "default": return maxtext_utils.get_mesh_from_config(cfg) devices = jax.devices("tpu") if jax.default_backend() == "tpu" else jax.devices() return jax.make_mesh((1, 1, 1, 1), ("data", "fsdp", "expert", "context"), devices=devices[:1]) @@ -87,7 +87,7 @@ def __init__( self.mesh = get_tpu_mesh(cfg) if cfg is None: cfg = setup_config( - "simple_mlp", + "default", emb_dim=hidden, mlp_dim=hidden * 2, vocab_size=vocab_size, @@ -98,24 +98,25 @@ def __init__( def __call__( self, - decoder_input_tokens, - decoder_positions=None, - decoder_segment_ids=None, - encoder_images=None, - encoder_image_masks=None, - enable_dropout=False, - decoder_target_tokens=None, - decoder_target_mask=None, - ): - del ( - encoder_images, - encoder_image_masks, - enable_dropout, - decoder_target_tokens, - decoder_target_mask, - ) + decoder_input_tokens: jax.Array, + decoder_positions: jax.Array | None = None, + decoder_segment_ids: jax.Array | None = None, + deterministic: bool = True, + model_mode: str = "train", + **kwargs: Any, + ) -> jax.Array: + del kwargs x = self.embed(decoder_input_tokens) - x = self.layer(x, decoder_positions, decoder_segment_ids, False, "train") + x = self.layer( + x, + positions=decoder_positions, + segmentation=decoder_segment_ids, + deterministic=deterministic, + model_mode=model_mode, + ) + # SimpleMlpDecoderLayer returns (output, None) when scan_layers=True (default in base.yml). + if isinstance(x, tuple): + x = x[0] return self.proj(x) @@ -141,17 +142,17 @@ class DummyPayload(abstract_engine.TrainerPayload): def make_dummy_data( batch_size: int = 2, - seq_len: int = 4, + seq_len: int = 16, vocab_size: int = 8, seed: int | None = None, cfg: pyconfig.HyperParameters | None = None, mask_prob: float = 0.0, ) -> dict[str, jax.Array]: """Constructs dummy token batch dictionary for loss_fn / train_step.""" - if cfg is not None and getattr(cfg, "model_name", "simple_mlp") != "simple_mlp": - batch_size = getattr(cfg, "micro_batch_size_to_train_on", batch_size) - seq_len = getattr(cfg, "max_target_length", seq_len) - vocab_size = getattr(cfg, "vocab_size", vocab_size) + if cfg is not None and cfg.model_name != "default": + batch_size = cfg.micro_batch_size_to_train_on + seq_len = cfg.max_target_length + vocab_size = cfg.vocab_size rng = np.random.default_rng(seed) if seed is not None else _DUMMY_DATA_RNG tokens = jnp.array(rng.integers(0, vocab_size, size=(batch_size, seq_len)), dtype=jnp.int32) targets = jnp.array(rng.integers(0, vocab_size, size=(batch_size, seq_len)), dtype=jnp.int32) @@ -174,7 +175,7 @@ def make_dummy_data( "decoder_loss_weights": weights, "decoder_positions": positions, } - if cfg is not None and getattr(cfg, "model_name", "simple_mlp") != "simple_mlp": + if cfg is not None and cfg.model_name != "default": mesh = get_tpu_mesh(cfg) data_sharding = sharding.get_input_data_sharding(cfg, mesh) res = {k: jax.device_put(v, data_sharding) for k, v in res.items()} @@ -182,7 +183,7 @@ def make_dummy_data( def setup_config( - model_name: str = "simple_mlp", + model_name: str = "default", gradient_accumulation_steps: int = 1, skip_step_on_spikes: bool = False, gradient_clipping_threshold: float = 0.0, @@ -198,11 +199,10 @@ def setup_config( ): config_cls = types.RLConfig - actual_model = "default" if model_name == "simple_mlp" else model_name argv = [ "compare_training_engine.py", base_yml, - f"model_name={actual_model}", + f"model_name={model_name}", f"gradient_accumulation_steps={gradient_accumulation_steps}", f"skip_step_on_spikes={skip_step_on_spikes}", f"gradient_clipping_threshold={gradient_clipping_threshold}", @@ -217,7 +217,7 @@ def setup_config( "record_internal_nn_metrics=False", "skip_jax_distributed_system=True", ] - if model_name == "simple_mlp": + if model_name == "default": argv.extend( [ "vocab_size=8", @@ -233,34 +233,29 @@ def setup_config( for override in cli_overrides: clean_override = override[2:] if override.startswith("--") else override if clean_override.startswith("model_name="): - override_model = clean_override.split("=", 1)[1] - actual_model = "default" if override_model == "simple_mlp" else override_model - argv[2] = f"model_name={actual_model}" + argv[2] = clean_override elif clean_override.endswith(".yml") or clean_override.endswith(".yaml"): argv[1] = clean_override if "rl" in os.path.basename(clean_override): config_cls = types.RLConfig else: argv.append(clean_override) - cfg = pyconfig.initialize(argv, config_class=config_cls) - if model_name == "simple_mlp": - cfg.model_name = "simple_mlp" - if not getattr(cfg, "compiled_trainstep_file", None): - cfg.compiled_trainstep_file = "" - return cfg + return pyconfig.initialize(argv, config_class=config_cls) def create_identical_models_and_opts( cfg: pyconfig.HyperParameters, + learning_rate_schedule: Any = None, ) -> tuple[Any, Any, Any, Any, Any, Any]: """Creates two identical model/optimizer pairs with identical initial weights.""" mesh = get_tpu_mesh(cfg) - if getattr(cfg, "model_name", "simple_mlp") == "simple_mlp": + if cfg.model_name == "default": + lr = learning_rate_schedule if learning_rate_schedule is not None else cfg.learning_rate model_baseline = TinyDecoder(vocab_size=cfg.vocab_size, hidden=4, rngs=nnx.Rngs(42)) opt_baseline = nnx.Optimizer( model_baseline, optax.adamw( - learning_rate=getattr(cfg, "learning_rate_schedule", 0.01), + learning_rate=lr, b1=0.9, b2=0.999, weight_decay=1e-4, @@ -272,7 +267,7 @@ def create_identical_models_and_opts( opt_engine = nnx.Optimizer( model_engine, optax.adamw( - learning_rate=getattr(cfg, "learning_rate_schedule", 0.01), + learning_rate=lr, b1=0.9, b2=0.999, weight_decay=1e-4, @@ -305,7 +300,7 @@ def create_identical_models_and_opts( config=cfg, mesh=mesh, model_mode=common_types.MODEL_MODE_TRAIN, - rng_key=jax.random.PRNGKey(getattr(cfg, "init_weights_seed", 42)), + rng_key=jax.random.PRNGKey(cfg.init_weights_seed), ) _, tx_b = train_utils.create_training_optimizer(cfg, model_baseline) opt_baseline = nnx.Optimizer(model_baseline, tx_b, wrt=nnx.Param) @@ -314,7 +309,7 @@ def create_identical_models_and_opts( config=cfg, mesh=mesh, model_mode=common_types.MODEL_MODE_TRAIN, - rng_key=jax.random.PRNGKey(getattr(cfg, "init_weights_seed", 42)), + rng_key=jax.random.PRNGKey(cfg.init_weights_seed), ) _, tx_e = train_utils.create_training_optimizer(cfg, model_engine) opt_engine = nnx.Optimizer(model_engine, tx_e, wrt=nnx.Param) @@ -410,9 +405,15 @@ class VerificationContext: @contextlib.contextmanager -def verification_harness(cfg: pyconfig.HyperParameters, compiled: bool = False) -> Any: +def verification_harness( + cfg: pyconfig.HyperParameters, + compiled: bool = False, + learning_rate_schedule: Any = None, +) -> Any: """Context manager establishing synchronized baseline/engine testing scaffold and memory cleanup.""" - model_b, opt_b, model_e, opt_e, state_shardings, params_shardings = create_identical_models_and_opts(cfg) + model_b, opt_b, model_e, opt_e, state_shardings, params_shardings = create_identical_models_and_opts( + cfg, learning_rate_schedule=learning_rate_schedule + ) mesh = get_tpu_mesh(cfg) ts_baseline = train_state_nnx.TrainStateNNX(model_b, opt_b) @@ -554,7 +555,7 @@ def verify_parity_with_train_py( ) -> None: """Verifies numerical and weight parity between standalone train_step and MaxTextTrainingEngine over N steps.""" if cfg is None: - cfg = setup_config("simple_mlp") + cfg = setup_config("default") with verification_harness(cfg, compiled) as ctx: p_train_step = None @@ -596,7 +597,7 @@ def verify_auxiliary_metrics_and_telemetry_parity( ) -> None: """Verifies aux metrics, gradient norm, and spike skipping telemetry parity over N steps (CL 956059885).""" if cfg is None: - cfg = setup_config("simple_mlp", skip_step_on_spikes=True, gradient_clipping_threshold=1.0) + cfg = setup_config("default", skip_step_on_spikes=True, gradient_clipping_threshold=1.0) with verification_harness(cfg, compiled) as ctx: p_train_step = None @@ -648,18 +649,23 @@ def verify_gradient_accumulation_parity( """Verifies multi-step gradient accumulation parity across M microbatches per outer step under dynamic LR.""" if cfg is None: cfg = setup_config( - "simple_mlp", + "default", gradient_accumulation_steps=m_steps, use_tunix_gradient_accumulation=True, ) - - def lr_schedule(step): - return 0.01 * (0.9**step) - - cfg.learning_rate_schedule = lr_schedule - cfg.use_tunix_gradient_accumulation = True - - with verification_harness(cfg, compiled) as ctx: + else: + assert ( + cfg.use_tunix_gradient_accumulation + ), "Gradient accumulation parity verification requires cfg.use_tunix_gradient_accumulation=True" + + def lr_schedule(step: int | float) -> float: + return cfg.learning_rate * (0.9**step) + + with verification_harness( + cfg, + compiled, + learning_rate_schedule=lr_schedule if cfg.model_name == "default" else None, + ) as ctx: p_train_step = None state_pure = ctx.state_pure for step in range(num_steps): @@ -743,7 +749,10 @@ def benchmark_gradient_accumulation_performance( gradient_accumulation_steps=m_steps, use_tunix_gradient_accumulation=True, ) - cfg.use_tunix_gradient_accumulation = True + else: + assert ( + cfg.use_tunix_gradient_accumulation + ), "Gradient accumulation performance benchmarking requires cfg.use_tunix_gradient_accumulation=True" print( "\n=== [BENCHMARK] Initializing Gradient Accumulation Hardware" @@ -896,16 +905,14 @@ def run_all_verifications(cli_overrides: list[str] | None = None) -> None: continue if clean_arg.startswith("model_name="): model_name = clean_arg.split("=", 1)[1] - if model_name != "simple_mlp": + if model_name != "default": overrides.append(arg) continue - elif clean_arg.endswith(".yml") and model_name == "simple_mlp": - model_name = "default" if clean_arg.startswith("max_target_length="): has_target_length = True overrides.append(arg) - if model_name != "simple_mlp" and not has_target_length: + if model_name != "default" and not has_target_length: overrides.append("max_target_length=256") cfg_base = setup_config( @@ -927,7 +934,7 @@ def run_all_verifications(cli_overrides: list[str] | None = None) -> None: cli_overrides=overrides, ) - if model_name == "simple_mlp": + if model_name == "default": print( "=== Running Training Engine Parity Verification Suite (Eager" " Mode) ===", flush=True, diff --git a/tests/end_to_end/tpu/test_training_engine_parity.sh b/tests/end_to_end/tpu/test_training_engine_parity.sh index 08b17ed831..feebf5e3d3 100755 --- a/tests/end_to_end/tpu/test_training_engine_parity.sh +++ b/tests/end_to_end/tpu/test_training_engine_parity.sh @@ -31,9 +31,9 @@ LLAMA_OVERRIDES=( ) echo "==========================================================================" -echo "=== [Steps 1-3] Eager Mode Evaluation on simple_mlp ===" +echo "=== [Steps 1-3] Eager Mode Evaluation on default (Tiny MLP) ===" echo "==========================================================================" -python3 tests/end_to_end/tpu/compare_training_engine.py model_name=simple_mlp test_suite=eager_all +python3 tests/end_to_end/tpu/compare_training_engine.py model_name=default test_suite=eager_all echo "" echo "==========================================================================" diff --git a/tests/integration/maxtext_engine_e2e_test.py b/tests/integration/maxtext_engine_e2e_test.py index 98af648e87..a48679825d 100644 --- a/tests/integration/maxtext_engine_e2e_test.py +++ b/tests/integration/maxtext_engine_e2e_test.py @@ -25,6 +25,7 @@ from maxtext.configs import pyconfig from maxtext.training_engine import abstract_engine from maxtext.training_engine import maxtext_engine +from tests.utils.test_helpers import get_test_config_path import optax import pytest @@ -100,20 +101,33 @@ def run( class MaxTextTrainingEngineE2ETest(absltest.TestCase): """End-to-end MaxText training engine test.""" - def setup_config(self, enable_checkpointing: bool = False): - """Sets up mock configuration for testing.""" - mock_config = mock.MagicMock(spec=pyconfig.HyperParameters) - mock_config.init_weights_seed = 42 - mock_config.model_name = "llama3.1-8b" - mock_config.tensorboard_dir = "/tmp/tb_dir" - mock_config.run_name = "test_run" - mock_config.enable_tensorboard = False + def setup_config(self, enable_checkpointing: bool = False, **kwargs): + """Sets up a MaxText config via pyconfig.initialize.""" + overrides = { + "model_name": "llama3.1-8b", + "run_name": "test_run", + "base_output_directory": self.create_tempdir().full_path, + "init_weights_seed": 42, + "micro_batch_size_to_train_on": 2, + "gradient_accumulation_steps": 1, + "enable_dropout": False, + "record_internal_nn_metrics": False, + "enable_tensorboard": False, + "tensorboard_dir": self.create_tempdir().full_path, + "skip_jax_distributed_system": True, + "enable_checkpointing": enable_checkpointing, + } if enable_checkpointing: - mock_config.checkpoint_directory = "/tmp/test_out/e2e_checkpoints" - mock_config.checkpoint_period = 2 - mock_config.max_num_checkpoints_to_keep = 5 - mock_config.async_checkpointing = True - return mock_config + overrides.update( + { + "checkpoint_dir": self.create_tempdir().full_path, + "checkpoint_period": 2, + "max_num_checkpoints_to_keep": 5, + "async_checkpointing": True, + } + ) + overrides.update(kwargs) + return pyconfig.initialize([None, get_test_config_path()], **overrides) @mock.patch.object(maxtext_engine.train_utils, "create_training_optimizer") @mock.patch.object(maxtext_engine.checkpointing, "CheckpointManager") diff --git a/tests/unit/gradient_accumulation_nnx_test.py b/tests/unit/gradient_accumulation_nnx_test.py index 478fe9a746..6a61304e34 100644 --- a/tests/unit/gradient_accumulation_nnx_test.py +++ b/tests/unit/gradient_accumulation_nnx_test.py @@ -30,6 +30,7 @@ @dataclass class _Cfg: gradient_accumulation_steps: int = 2 + use_tunix_gradient_accumulation: bool = False shard_optimizer_over_data: bool = False shard_mode: int = ShardMode.AUTO ici_data_parallelism: int = 1