diff --git a/simpletuner/helpers/training/trainer.py b/simpletuner/helpers/training/trainer.py index 483cd4d68..cbdc7c7dd 100644 --- a/simpletuner/helpers/training/trainer.py +++ b/simpletuner/helpers/training/trainer.py @@ -598,7 +598,10 @@ def _update_grad_metrics( elif ( not require_value_method or self.config.grad_clip_method == "value" ) and not self.config.use_deepspeed_optimizer: - target_logs[f"{prefix}grad_absmax"] = self.grad_norm + grad_value = self.grad_norm + if clone_norm_value: + grad_value = float(self.grad_norm.clone().detach()) + target_logs[f"{prefix}grad_absmax"] = grad_value def _config_uses_bitsandbytes(self) -> bool: if not getattr(self, "config", None): diff --git a/tests/test_trainer.py b/tests/test_trainer.py index e2a93cd77..7f741bf88 100644 --- a/tests/test_trainer.py +++ b/tests/test_trainer.py @@ -1249,6 +1249,19 @@ def test_update_grad_metrics_logs_absmax_without_deepspeed(self): self.assertIn("grad_absmax", logs) self.assertIs(logs["grad_absmax"], trainer.grad_norm) + def test_update_grad_metrics_clones_absmax_value_when_requested(self): + trainer = self._build_trainer_for_grad_logging( + grad_clip_method="value", + use_deepspeed=False, + grad_value=torch.tensor(1.2), + ) + logs = {} + trainer._update_grad_metrics(logs, clone_norm_value=True) + self.assertIn("grad_absmax", logs) + self.assertEqual( + logs["grad_absmax"], float(trainer.grad_norm.clone().detach()) + ) + def test_update_grad_metrics_clones_norm_value_when_requested(self): trainer = self._build_trainer_for_grad_logging( grad_clip_method="norm", @@ -1309,7 +1322,9 @@ def test_compose_training_progress_metrics_includes_grad_absmax(self): ) metrics = trainer._compose_training_progress_metrics(epoch=1) self.assertIn("grad_absmax", metrics) - self.assertIs(metrics["grad_absmax"], trainer.grad_norm) + self.assertEqual( + metrics["grad_absmax"], float(trainer.grad_norm.clone().detach()) + ) self.assertNotIn("grad_norm", metrics) def test_compose_training_progress_metrics_excludes_grad_with_deepspeed(self):