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
5 changes: 4 additions & 1 deletion simpletuner/helpers/training/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
17 changes: 16 additions & 1 deletion tests/test_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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):
Expand Down
Loading