diff --git a/src/maxtext/training_engine/abstract_engine.py b/src/maxtext/training_engine/abstract_engine.py index bc85e849e4..8b9144725c 100644 --- a/src/maxtext/training_engine/abstract_engine.py +++ b/src/maxtext/training_engine/abstract_engine.py @@ -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. diff --git a/src/maxtext/training_engine/maxtext_engine.py b/src/maxtext/training_engine/maxtext_engine.py index 2a03a6df0d..c7a9bcef0b 100644 --- a/src/maxtext/training_engine/maxtext_engine.py +++ b/src/maxtext/training_engine/maxtext_engine.py @@ -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 @@ -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)) @@ -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.""" @@ -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. @@ -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 diff --git a/tests/integration/maxtext_engine_e2e_test.py b/tests/integration/maxtext_engine_e2e_test.py index 4d832f14a6..98af648e87 100644 --- a/tests/integration/maxtext_engine_e2e_test.py +++ b/tests/integration/maxtext_engine_e2e_test.py @@ -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)), {}, ) ) diff --git a/tests/maxtext_engine_test.py b/tests/maxtext_engine_test.py index f19efd3ab9..b11bb5493c 100644 --- a/tests/maxtext_engine_test.py +++ b/tests/maxtext_engine_test.py @@ -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() @@ -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() @@ -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) @@ -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()