Skip to content
Closed
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
1 change: 1 addition & 0 deletions src/maxtext/configs/base.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
6 changes: 3 additions & 3 deletions src/maxtext/training_engine/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
),
)

Expand Down
3 changes: 1 addition & 2 deletions src/maxtext/training_engine/inflight_throttler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
16 changes: 6 additions & 10 deletions src/maxtext/training_engine/maxtext_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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()
Expand Down Expand Up @@ -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,
)

Expand All @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion src/maxtext/training_engine/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
4 changes: 1 addition & 3 deletions src/maxtext/utils/gradient_accumulation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading