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
13 changes: 13 additions & 0 deletions src/maxtext/training_engine/abstract_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,19 @@ def compute(self) -> jax.Array:
return self.unreduced_sum * self.compute_scale()


@flax.struct.dataclass
class LossOutput:
"""Output of a loss function containing unreduced primary loss and aux metrics.

Attributes:
primary_loss: The main loss to be optimized.
aux_metrics: A dictionary of auxiliary metrics.
"""

primary_loss: WeightedMetric
aux_metrics: dict[str, Any] = flax.struct.field(default_factory=dict)


@flax.struct.dataclass
class MetricsBuffer:
"""A buffer for storing and aggregating unreduced metrics on-device.
Expand Down
63 changes: 54 additions & 9 deletions src/maxtext/training_engine/maxtext_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,7 @@ def __init__(
self._state: Any = None
self._accumulated_grads: Any = None
self._micro_step_count = 0
self._cached_losses: list[jax.Array] = []
self._cached_losses: list[abstract_engine.WeightedMetric | jax.Array] = []
self._learning_rate_schedule, self._optimizer = train_utils.create_training_optimizer(self._config, self._model)
self._train_step: int = 0

Expand Down Expand Up @@ -177,12 +177,51 @@ def _fwd_bwd_kernel(self, params, rest, batch):

def diff_wrapper(p, r, b):
mdl = nnx.merge(self._model_graphdef, p, r, copy=True)
loss, aux = loss_callable(mdl, self._config, b, None, None, is_train=True)
out = loss_callable(mdl, self._config, b, None, None, is_train=True)
_, _, new_r = nnx.split(mdl, nnx.Param, ...)
return loss, (aux, new_r)

if isinstance(out, abstract_engine.LossOutput):
return out.primary_loss.unreduced_sum, (out, new_r)
elif isinstance(out, abstract_engine.WeightedMetric):
loss_out = abstract_engine.LossOutput(
primary_loss=out,
aux_metrics={},
)
return out.unreduced_sum, (loss_out, new_r)
elif isinstance(out, (tuple, list)) and len(out) == 2:
loss_val, aux = out
if isinstance(loss_val, abstract_engine.WeightedMetric):
primary_loss = loss_val
elif isinstance(aux, dict) and "xent_sum" in aux and "total_weights" in aux:
primary_loss = abstract_engine.WeightedMetric(
unreduced_sum=aux["xent_sum"],
denominator=aux["total_weights"],
)
else:
raise TypeError(
f"Cannot construct WeightedMetric from 2-tuple loss return with elements "
f"of type ({type(loss_val).__name__}, {type(aux).__name__}). Expected first element to be a "
"WeightedMetric, or second element to be a dict containing 'xent_sum' and 'total_weights'."
)

loss_out = abstract_engine.LossOutput(
primary_loss=primary_loss,
aux_metrics=aux if isinstance(aux, dict) else {},
)
return primary_loss.unreduced_sum, (loss_out, new_r)
else:
raise TypeError(
f"Unsupported return type from loss function: {type(out)}. "
"Expected abstract_engine.LossOutput, abstract_engine.WeightedMetric, "
"or a 2-element tuple/list: (loss, aux_metrics)."
)

grad_func = jax.value_and_grad(diff_wrapper, argnums=0, has_aux=True)
(loss, (aux, new_rest)), micro_grads = grad_func(params, rest, batch)
(loss_val, (loss_out, new_rest)), micro_grads = grad_func(params, rest, batch)
if isinstance(loss_out, abstract_engine.LossOutput):
scale = loss_out.primary_loss.compute_scale()
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))
Expand All @@ -191,7 +230,11 @@ def diff_wrapper(p, r, b):
),
micro_grads,
)
return loss, aux, new_rest, micro_grads

if isinstance(loss_out, abstract_engine.LossOutput):
return loss_out.primary_loss, loss_out.aux_metrics, new_rest, micro_grads
else:
return loss_val, {}, new_rest, micro_grads

def _update_kernel(self, state_pure, accumulated_grads, micro_step_count, mean_loss):
"""Applies accumulated gradients to update the NNX model state."""
Expand Down Expand Up @@ -326,9 +369,7 @@ def fwd_bwd(self, payload: abstract_engine.TrainerPayload) -> None:
# the update step.
self._throttler.add_computation(computation=loss, metrics=None)

if loss is not None:
# TODO(mazumdera): This needs to be modified to become
# if isinstance(loss, abstract_engine.WeightedMetric):
if isinstance(loss, abstract_engine.WeightedMetric):
self.record_metrics("loss", loss)

# Record auxiliary metrics.
Expand Down Expand Up @@ -364,7 +405,11 @@ def update(self) -> None:
self._state = train_state_nnx.TrainStateNNX(self._model, self._optimizer)
self._state_graphdef, state_pure = nnx.split(self._state)

mean_loss = jnp.mean(jnp.array(self._cached_losses)) if self._cached_losses else jnp.array(0.0)
if self._cached_losses:
loss_values = [l.compute() if isinstance(l, abstract_engine.WeightedMetric) else l for l in self._cached_losses]
mean_loss = jnp.mean(jnp.stack(loss_values)) if len(loss_values) > 1 else loss_values[0]
else:
mean_loss = jnp.array(0.0)
if self._compiled and hasattr(self, "_compiled_update"):
new_state_pure, grad_norm, is_skipped = self._compiled_update(
state_pure, self._accumulated_grads, self._micro_step_count, mean_loss
Expand Down
2 changes: 1 addition & 1 deletion tests/integration/maxtext_engine_e2e_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,7 +133,7 @@ def test_e2e_training_loop_exercises_all_trainer_apis(self, mock_from_pretrained

trainer_instance.with_loss_fn(
lambda *args, **kwargs: (
jnp.array(0.25),
abstract_engine.WeightedMetric(unreduced_sum=jnp.array(0.25), denominator=jnp.array(1.0)),
{},
)
)
Expand Down
79 changes: 75 additions & 4 deletions tests/maxtext_engine_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,12 @@ def test_raises_value_error_for_missing_model_name(self):

def test_max_text_trainer_instantiation_with_pyconfig(self):
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(lambda *args, **kwargs: (jnp.array(0.5), {}))
t.with_loss_fn(
lambda *args, **kwargs: (
abstract_engine.WeightedMetric(unreduced_sum=jnp.array(0.5), denominator=jnp.array(1.0)),
{},
)
)
self.assertIsInstance(t, abstract_engine.AbstractTrainingEngine)
self.mock_from_pretrained.assert_called_once()

Expand All @@ -119,8 +124,6 @@ def test_max_text_trainer_instantiation_with_pyconfig(self):
)
t.compile(payload)
self.assertTrue(t._compiled)
t.with_loss_fn(lambda *args, **kwargs: (jnp.array(0.5), {}))
self.assertFalse(t._compiled)
t.fwd_bwd(payload)
self.assertEqual(t._micro_step_count, 1)
t.update()
Expand Down Expand Up @@ -342,7 +345,12 @@ def test_record_and_get_metrics(self):

def test_update_with_inflight_throttling(self):
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
t.with_loss_fn(lambda *args, **kwargs: (jnp.array(0.5), {}))
t.with_loss_fn(
lambda *args, **kwargs: (
abstract_engine.WeightedMetric(unreduced_sum=jnp.array(0.5), denominator=jnp.array(1.0)),
{},
)
)

payload = DummyPayload()
t.compile(payload)
Expand Down Expand Up @@ -388,6 +396,69 @@ def test_update_with_inflight_throttling(self):
t.close()
self.assertTrue(t._throttler._inflight_queue.empty())

def test_fwd_bwd_with_loss_output_and_aux_metrics(self):
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
payload = DummyPayload()

def _loss_fn(model, *args, **kwargs):
return abstract_engine.LossOutput(
primary_loss=abstract_engine.WeightedMetric(
unreduced_sum=jnp.sum(model.weights[...]) * 8.0, denominator=jnp.array(4.0)
),
aux_metrics={
"metric_a": abstract_engine.WeightedMetric(unreduced_sum=jnp.array(12.0), denominator=jnp.array(3.0)),
"metric_b": jnp.array(0.42),
},
)

t.with_loss_fn(_loss_fn)
t.compile(payload)
t.fwd_bwd(payload)

self.assertEqual(t._micro_step_count, 1)
self.assertIsNotNone(t._accumulated_grads)

# Check that grad is scaled by 1/4.0
np.testing.assert_allclose(t._accumulated_grads["weights"], jnp.array([2.0, 2.0]), rtol=1e-5)

metrics = t.get_metrics(clear_cache=True)
self.assertLen(metrics, 1)
self.assertIn("loss", metrics[0].weighted_metrics)
self.assertIn("metric_a", metrics[0].weighted_metrics)
self.assertIn("metric_b", metrics[0].scalar_metrics)
self.assertAlmostEqual(
float(metrics[0].weighted_metrics["loss"].compute().item()),
6.0,
places=4,
)
self.assertAlmostEqual(
float(metrics[0].weighted_metrics["metric_a"].compute().item()),
4.0,
places=4,
)

def test_fwd_bwd_with_loss_and_aux_dict_tuple(self):
t = maxtext_engine.MaxTextTrainingEngine(self.mock_config)
payload = DummyPayload()

def custom_loss(model, *args, **kwargs):
unreduced_sum = jnp.sum(model.weights[...]) * 8.0
denominator = jnp.array(4.0)
return unreduced_sum / denominator, {
"aux_stat": jnp.array(1.23),
"xent_sum": unreduced_sum,
"total_weights": denominator,
}

t.with_loss_fn(custom_loss)
t.fwd_bwd(payload)

# Check that grad is scaled by 1/4.0
np.testing.assert_allclose(t._accumulated_grads["weights"], jnp.array([2.0, 2.0]), rtol=1e-5)
metrics = t.get_metrics(clear_cache=True)
self.assertIn("loss", metrics[0].weighted_metrics)
self.assertIn("aux_stat", metrics[0].scalar_metrics)


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