Aryan Sethi Claude Opus 5 (1M context) commited on
Commit
be3d1cc
·
1 Parent(s): e4f085d

Make the gradient-norm probe safe to leave on

Browse files

No architectural change was involved in getting the gradient norm -- the model,
loss, optimiser, data pipeline and resolved training config are untouched, and
the logging attaches through the public add_callback API. But the probe patches
torch.nn.utils.clip_grad_norm_, and reviewing that for what goes wrong in
practice turned up three things worth fixing.

Ultralytics fires on_train_end only on success, so a crash left the patch
installed and the second model's probe then wrapped the wrapper, making
restoration hand back a wrapper instead of the real function. install() now
marks the function it creates and refuses to wrap a marked one, and train.py
calls close() from a finally block.

float() on a CUDA tensor synchronises the device, and clip_grad_norm_ does not
otherwise sync since error_if_nonfinite defaults to False. Converting every step
added a stall to every step of the run. Norms are kept as detached tensors and
converted once per epoch instead.

Gradient norm is now its own flag, --log-grad-norm, so loss and metric curves no
longer require patching anything in torch.

Fixes a NameError the refactor introduced: _run_training referenced model_name
from the enclosing scope.

tests/test_grad_norm_probe.py pins the behaviour against a stub torch, needing
neither a GPU nor a 2 GB install: restoration, refusal to nest, pass-through of
the caller's value, and that an uninstalled probe leaves torch alone.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

README.md CHANGED
@@ -419,18 +419,33 @@ the process is changed.
419
 
420
  ### Getting at the gradient norm
421
 
 
 
 
 
422
  `BaseTrainer.optimizer_step` calls `clip_grad_norm_`, whose return value is the
423
  total pre-clip gradient norm, and discards it. It then calls `zero_grad()` in the
424
  same method, and **no callback fires between the two** — the nearest,
425
  `on_train_batch_end`, runs after the gradients are already gone, so it cannot be
426
- recomputed either.
427
-
428
- Wrapping `clip_grad_norm_` for the duration of training is the only hook
429
- available. It is also the most accurate one: the captured value is taken after
430
- `scaler.unscale_()`, so it is a true unscaled norm rather than an AMP-scaled one,
431
- and it is the same number the optimiser acted on. Summarised per epoch rather
432
- than per step, since per-step would be thousands of points that do not line up
433
- with the losses.
 
 
 
 
 
 
 
 
 
 
 
434
 
435
  `grad_norm/clipped_fraction` is the series to read first: the share of steps
436
  whose norm exceeded the clip of 10. Near 1.0 during warmup is normal. If it stays
 
419
 
420
  ### Getting at the gradient norm
421
 
422
+ This is the one part of the logging that is not a callback, and it is opt-in
423
+ separately as `--log-grad-norm`. Losses and metrics come from callbacks that only
424
+ read `trainer.*`; without this flag nothing in torch is touched.
425
+
426
  `BaseTrainer.optimizer_step` calls `clip_grad_norm_`, whose return value is the
427
  total pre-clip gradient norm, and discards it. It then calls `zero_grad()` in the
428
  same method, and **no callback fires between the two** — the nearest,
429
  `on_train_batch_end`, runs after the gradients are already gone, so it cannot be
430
+ recomputed either. Wrapping `clip_grad_norm_` for the duration of training is the
431
+ only hook available, and the most accurate one: the value is captured after
432
+ `scaler.unscale_()`, so it is a true unscaled norm and the same number the
433
+ optimiser acted on.
434
+
435
+ The wrapper only reads a return value. It never touches gradients, the optimiser
436
+ or the loss, so it cannot change what the model learns. Three things it could
437
+ still have cost you, all handled:
438
+
439
+ | Risk | Handling |
440
+ | --- | --- |
441
+ | `on_train_end` does not fire when training raises, so the patch survives and the second model's probe wraps the wrapper | `install()` marks the function it creates and refuses to wrap a marked one; `train.py` calls `close()` from a `finally` |
442
+ | `float()` on a CUDA tensor synchronises the device, and `clip_grad_norm_` does not otherwise sync (`error_if_nonfinite` is False), so converting per step would stall every step | norms are kept as detached tensors and converted once per epoch |
443
+ | Coupling it to `--wandb` would force the patch on anyone wanting loss curves | separate flag |
444
+
445
+ [`tests/test_grad_norm_probe.py`](tests/test_grad_norm_probe.py) pins all of it
446
+ against a stub torch — no GPU, no 2 GB install — covering restoration, refusal to
447
+ nest, pass-through of the caller's value, and that constructing a probe without
448
+ installing it leaves torch untouched.
449
 
450
  `grad_norm/clipped_fraction` is the series to read first: the share of steps
451
  whose norm exceeded the clip of 10. Near 1.0 during warmup is normal. If it stays
notebooks/colab_train.ipynb CHANGED
@@ -139,7 +139,7 @@
139
  " --data /content/data/data.yaml \\\n",
140
  " --models yolo11n.pt yolo11s.pt \\\n",
141
  " --epochs 100 --imgsz 640 --batch 16 --device 0 \\\n",
142
- " --wandb --wandb-project cone-distance \\\n",
143
  " --project /content/runs"
144
  ]
145
  },
@@ -163,7 +163,12 @@
163
  "gradients at norm 10, and this is the share of steps that hit the clip. It sits\n",
164
  "near 1.0 during warmup, which is normal; if it stays there once warmup ends, the\n",
165
  "effective step size is being set by the clip rather than by `lr0`, and lowering\n",
166
- "the learning rate will do more than tuning anything else."
 
 
 
 
 
167
  ]
168
  },
169
  {
 
139
  " --data /content/data/data.yaml \\\n",
140
  " --models yolo11n.pt yolo11s.pt \\\n",
141
  " --epochs 100 --imgsz 640 --batch 16 --device 0 \\\n",
142
+ " --wandb --wandb-project cone-distance --log-grad-norm \\\n",
143
  " --project /content/runs"
144
  ]
145
  },
 
163
  "gradients at norm 10, and this is the share of steps that hit the clip. It sits\n",
164
  "near 1.0 during warmup, which is normal; if it stays there once warmup ends, the\n",
165
  "effective step size is being set by the clip rather than by `lr0`, and lowering\n",
166
+ "the learning rate will do more than tuning anything else.\n",
167
+ "\n",
168
+ "Gradient norm is the one thing here that patches a torch function, because\n",
169
+ "Ultralytics discards the value and fires no callback while the gradients are\n",
170
+ "live. Dropping `--log-grad-norm` above removes that patch entirely and keeps\n",
171
+ "every loss and metric curve."
172
  ]
173
  },
174
  {
src/train/train.py CHANGED
@@ -118,17 +118,36 @@ def train_one(model_name: str, data_yaml: Path, args) -> None:
118
 
119
  model = YOLO(model_name)
120
 
 
121
  if args.wandb:
122
  from src.train import wandb_logger
123
 
124
- wandb_logger.attach(
125
  model,
126
  project=args.wandb_project,
127
  run_name=f"{Path(model_name).stem}-{args.epochs}ep",
128
  extra_config={"dataset": str(data_yaml)},
 
129
  )
130
- print(f"logging to W&B project {args.wandb_project!r}")
 
131
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
132
  model.train(
133
  data=str(data_yaml),
134
  epochs=args.epochs,
@@ -142,10 +161,6 @@ def train_one(model_name: str, data_yaml: Path, args) -> None:
142
  seed=args.seed,
143
  )
144
 
145
- metrics = model.val(data=str(data_yaml), imgsz=args.imgsz, device=args.device)
146
- print_per_class_metrics(metrics, model.names)
147
- print(f"\nweights: {args.project}/{Path(model_name).stem}/weights/best.pt")
148
-
149
 
150
  def main() -> None:
151
  parser = argparse.ArgumentParser(description=__doc__,
@@ -167,6 +182,12 @@ def main() -> None:
167
  help="log train/val loss and gradient norm to Weights & Biases; "
168
  "authenticate with $WANDB_API_KEY")
169
  parser.add_argument("--wandb-project", default="cone-distance")
 
 
 
 
 
 
170
  args = parser.parse_args()
171
 
172
  if args.wandb and not os.environ.get("WANDB_API_KEY"):
 
118
 
119
  model = YOLO(model_name)
120
 
121
+ close_logger = None
122
  if args.wandb:
123
  from src.train import wandb_logger
124
 
125
+ close_logger = wandb_logger.attach(
126
  model,
127
  project=args.wandb_project,
128
  run_name=f"{Path(model_name).stem}-{args.epochs}ep",
129
  extra_config={"dataset": str(data_yaml)},
130
+ log_grad_norm=args.log_grad_norm,
131
  )
132
+ print(f"logging to W&B project {args.wandb_project!r}"
133
+ + (" (with gradient norm)" if args.log_grad_norm else ""))
134
 
135
+ try:
136
+ _run_training(model, model_name, data_yaml, args)
137
+ finally:
138
+ # Ultralytics fires on_train_end only on success, and the gradient-norm
139
+ # probe patches a torch function that must be put back either way --
140
+ # otherwise a crash training the first model leaves a wrapper in place
141
+ # for the second.
142
+ if close_logger is not None:
143
+ close_logger()
144
+
145
+ metrics = model.val(data=str(data_yaml), imgsz=args.imgsz, device=args.device)
146
+ print_per_class_metrics(metrics, model.names)
147
+ print(f"\nweights: {args.project}/{Path(model_name).stem}/weights/best.pt")
148
+
149
+
150
+ def _run_training(model, model_name: str, data_yaml: Path, args) -> None:
151
  model.train(
152
  data=str(data_yaml),
153
  epochs=args.epochs,
 
161
  seed=args.seed,
162
  )
163
 
 
 
 
 
164
 
165
  def main() -> None:
166
  parser = argparse.ArgumentParser(description=__doc__,
 
182
  help="log train/val loss and gradient norm to Weights & Biases; "
183
  "authenticate with $WANDB_API_KEY")
184
  parser.add_argument("--wandb-project", default="cone-distance")
185
+ parser.add_argument("--log-grad-norm", action="store_true",
186
+ help="also log gradient norm. Needs --wandb. Unlike the loss "
187
+ "callbacks, this patches torch.nn.utils.clip_grad_norm_ "
188
+ "for the duration of training, because Ultralytics "
189
+ "discards the value and fires no callback while the "
190
+ "gradients are still live.")
191
  args = parser.parse_args()
192
 
193
  if args.wandb and not os.environ.get("WANDB_API_KEY"):
src/train/wandb_logger.py CHANGED
@@ -29,6 +29,11 @@ import numpy as np
29
  CLIP_MAX_NORM = 10.0
30
 
31
 
 
 
 
 
 
32
  class _GradNormProbe:
33
  """Capture the gradient norm that Ultralytics computes and throws away.
34
 
@@ -43,24 +48,45 @@ class _GradNormProbe:
43
  only hook available. It is also the most accurate one: the value captured
44
  is taken after `scaler.unscale_()`, so it is a true unscaled norm rather
45
  than an AMP-scaled one, and it is the same number the optimiser acted on.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46
  """
47
 
48
  def __init__(self) -> None:
49
  self._original = None
50
- self.norms: list[float] = []
51
 
52
  def install(self) -> None:
53
  import torch
54
 
55
- if self._original is not None:
 
 
 
56
  return
57
- self._original = torch.nn.utils.clip_grad_norm_
58
 
59
  def recording_clip_grad_norm(parameters, max_norm, *args, **kwargs):
60
  total_norm = self._original(parameters, max_norm, *args, **kwargs)
61
- self.norms.append(float(total_norm))
 
 
 
62
  return total_norm
63
 
 
64
  torch.nn.utils.clip_grad_norm_ = recording_clip_grad_norm
65
 
66
  def remove(self) -> None:
@@ -69,13 +95,15 @@ class _GradNormProbe:
69
  if self._original is not None:
70
  torch.nn.utils.clip_grad_norm_ = self._original
71
  self._original = None
 
72
 
73
  def summarise(self) -> dict[str, float]:
74
  """Per-epoch summary. Logging every step would be thousands of points
75
- per epoch and would not line up with the losses."""
76
- if not self.norms:
 
77
  return {}
78
- norms = np.array(self.norms)
79
  return {
80
  "grad_norm/mean": float(norms.mean()),
81
  "grad_norm/max": float(norms.max()),
@@ -86,18 +114,28 @@ class _GradNormProbe:
86
  }
87
 
88
  def reset(self) -> None:
89
- self.norms.clear()
90
 
91
 
92
- def attach(model, project: str, run_name: str, extra_config: dict | None = None) -> None:
 
93
  """Register W&B callbacks on one Ultralytics model.
94
 
95
  Call once per model. Each model gets its own W&B run, so the YOLO11n and
96
  YOLO11s arms appear as two comparable runs rather than one interleaved mess.
 
 
 
 
 
 
 
 
 
97
  """
98
  import wandb
99
 
100
- probe = _GradNormProbe()
101
 
102
  def on_pretrain_routine_start(trainer):
103
  wandb.init(
@@ -106,10 +144,12 @@ def attach(model, project: str, run_name: str, extra_config: dict | None = None)
106
  config={**vars(trainer.args), **(extra_config or {})},
107
  reinit=True,
108
  )
109
- probe.install()
 
110
 
111
  def on_train_epoch_start(trainer):
112
- probe.reset()
 
113
 
114
  def on_train_epoch_end(trainer):
115
  step = trainer.epoch + 1
@@ -122,13 +162,20 @@ def attach(model, project: str, run_name: str, extra_config: dict | None = None)
122
  # step as the training losses above so the two curves overlay.
123
  step = trainer.epoch + 1
124
  wandb.log(trainer.metrics, step=step)
125
- summary = probe.summarise()
126
- if summary:
127
- wandb.log(summary, step=step)
 
128
 
129
  def on_train_end(trainer):
130
- probe.remove()
131
- wandb.finish()
 
 
 
 
 
 
132
 
133
  for event, handler in (
134
  ("on_pretrain_routine_start", on_pretrain_routine_start),
@@ -138,3 +185,5 @@ def attach(model, project: str, run_name: str, extra_config: dict | None = None)
138
  ("on_train_end", on_train_end),
139
  ):
140
  model.add_callback(event, handler)
 
 
 
29
  CLIP_MAX_NORM = 10.0
30
 
31
 
32
+ # Set on the patched function so a second install can recognise its own work and
33
+ # refuse to wrap a wrapper.
34
+ _PATCH_MARKER = "_cone_distance_grad_norm_probe"
35
+
36
+
37
  class _GradNormProbe:
38
  """Capture the gradient norm that Ultralytics computes and throws away.
39
 
 
48
  only hook available. It is also the most accurate one: the value captured
49
  is taken after `scaler.unscale_()`, so it is a true unscaled norm rather
50
  than an AMP-scaled one, and it is the same number the optimiser acted on.
51
+
52
+ The wrapper only reads a return value -- it never touches gradients, the
53
+ optimiser or the loss -- so it cannot change what the model learns. Two
54
+ things it could still cost you, both handled here:
55
+
56
+ * `float()` on a CUDA tensor forces a device synchronisation, and
57
+ `clip_grad_norm_` does not otherwise sync (`error_if_nonfinite` is False
58
+ by default). Doing that every step would add a stall to every step of
59
+ the run, so norms are kept as detached tensors and converted once per
60
+ epoch instead.
61
+ * If training raises, `on_train_end` never fires. `install()` is therefore
62
+ idempotent, and `train.py` removes the probe from a `finally` block, so
63
+ a crash in the first model cannot leave a wrapper in place for the
64
+ second.
65
  """
66
 
67
  def __init__(self) -> None:
68
  self._original = None
69
+ self._norms: list = []
70
 
71
  def install(self) -> None:
72
  import torch
73
 
74
+ current = torch.nn.utils.clip_grad_norm_
75
+ if self._original is not None or getattr(current, _PATCH_MARKER, False):
76
+ # Already patched, by us or by a probe whose removal was skipped.
77
+ # Wrapping again would nest wrappers and make restoration wrong.
78
  return
79
+ self._original = current
80
 
81
  def recording_clip_grad_norm(parameters, max_norm, *args, **kwargs):
82
  total_norm = self._original(parameters, max_norm, *args, **kwargs)
83
+ # detach() only, no float(): converting here would synchronise the
84
+ # device on every optimiser step.
85
+ self._norms.append(total_norm.detach() if hasattr(total_norm, "detach")
86
+ else total_norm)
87
  return total_norm
88
 
89
+ setattr(recording_clip_grad_norm, _PATCH_MARKER, True)
90
  torch.nn.utils.clip_grad_norm_ = recording_clip_grad_norm
91
 
92
  def remove(self) -> None:
 
95
  if self._original is not None:
96
  torch.nn.utils.clip_grad_norm_ = self._original
97
  self._original = None
98
+ self._norms.clear()
99
 
100
  def summarise(self) -> dict[str, float]:
101
  """Per-epoch summary. Logging every step would be thousands of points
102
+ per epoch and would not line up with the losses. This is also the only
103
+ place the device is synchronised -- once per epoch, not once per step."""
104
+ if not self._norms:
105
  return {}
106
+ norms = np.array([float(n) for n in self._norms])
107
  return {
108
  "grad_norm/mean": float(norms.mean()),
109
  "grad_norm/max": float(norms.max()),
 
114
  }
115
 
116
  def reset(self) -> None:
117
+ self._norms.clear()
118
 
119
 
120
+ def attach(model, project: str, run_name: str, extra_config: dict | None = None,
121
+ log_grad_norm: bool = False):
122
  """Register W&B callbacks on one Ultralytics model.
123
 
124
  Call once per model. Each model gets its own W&B run, so the YOLO11n and
125
  YOLO11s arms appear as two comparable runs rather than one interleaved mess.
126
+
127
+ Returns a `close()` callable. Call it from a `finally` block: Ultralytics
128
+ only fires `on_train_end` on success, and with `log_grad_norm` the probe
129
+ patches a torch function that must be put back even when training raises.
130
+
131
+ `log_grad_norm` is separate from W&B logging on purpose. Losses and metrics
132
+ come from callbacks that only read `trainer.*`; the gradient norm needs a
133
+ patched `torch.nn.utils.clip_grad_norm_`. Leaving it off means nothing in
134
+ torch is touched.
135
  """
136
  import wandb
137
 
138
+ probe = _GradNormProbe() if log_grad_norm else None
139
 
140
  def on_pretrain_routine_start(trainer):
141
  wandb.init(
 
144
  config={**vars(trainer.args), **(extra_config or {})},
145
  reinit=True,
146
  )
147
+ if probe is not None:
148
+ probe.install()
149
 
150
  def on_train_epoch_start(trainer):
151
+ if probe is not None:
152
+ probe.reset()
153
 
154
  def on_train_epoch_end(trainer):
155
  step = trainer.epoch + 1
 
162
  # step as the training losses above so the two curves overlay.
163
  step = trainer.epoch + 1
164
  wandb.log(trainer.metrics, step=step)
165
+ if probe is not None:
166
+ summary = probe.summarise()
167
+ if summary:
168
+ wandb.log(summary, step=step)
169
 
170
  def on_train_end(trainer):
171
+ close()
172
+
173
+ def close():
174
+ """Restore torch and end the run. Safe to call more than once."""
175
+ if probe is not None:
176
+ probe.remove()
177
+ if wandb.run is not None:
178
+ wandb.finish()
179
 
180
  for event, handler in (
181
  ("on_pretrain_routine_start", on_pretrain_routine_start),
 
185
  ("on_train_end", on_train_end),
186
  ):
187
  model.add_callback(event, handler)
188
+
189
+ return close
tests/test_grad_norm_probe.py ADDED
@@ -0,0 +1,97 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """The gradient-norm probe patches a torch function, so it gets its own test.
2
+
3
+ Runs against a stub torch, so it needs neither a GPU nor a 2 GB install. What
4
+ it pins down is the behaviour that would be expensive to discover during a real
5
+ 100-epoch run: that the patch restores cleanly, that it refuses to wrap itself,
6
+ and that it never alters the value the optimiser receives.
7
+
8
+ python -m tests.test_grad_norm_probe
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import sys
14
+ import types
15
+
16
+
17
+ def install_stub_torch() -> types.ModuleType:
18
+ """Minimal stand-in exposing only torch.nn.utils.clip_grad_norm_."""
19
+ torch = types.ModuleType("torch")
20
+ nn = types.ModuleType("torch.nn")
21
+ utils = types.ModuleType("torch.nn.utils")
22
+
23
+ class FakeNorm(float):
24
+ """Stands in for the tensor clip_grad_norm_ returns."""
25
+ def detach(self):
26
+ return self
27
+
28
+ calls = []
29
+
30
+ def clip_grad_norm_(parameters, max_norm, *args, **kwargs):
31
+ calls.append(max_norm)
32
+ return FakeNorm(12.5)
33
+
34
+ clip_grad_norm_.calls = calls
35
+ utils.clip_grad_norm_ = clip_grad_norm_
36
+ nn.utils = utils
37
+ torch.nn = nn
38
+ sys.modules["torch"] = torch
39
+ sys.modules["torch.nn"] = nn
40
+ sys.modules["torch.nn.utils"] = utils
41
+ return torch
42
+
43
+
44
+ def main() -> None:
45
+ torch = install_stub_torch()
46
+ from src.train.wandb_logger import _GradNormProbe
47
+
48
+ pristine = torch.nn.utils.clip_grad_norm_
49
+
50
+ # 1. Install, and the value the caller gets back is untouched.
51
+ a = _GradNormProbe()
52
+ a.install()
53
+ assert torch.nn.utils.clip_grad_norm_ is not pristine, "probe did not patch"
54
+ returned = torch.nn.utils.clip_grad_norm_([], 10.0)
55
+ assert float(returned) == 12.5, returned
56
+ assert pristine.calls == [10.0], "the real function was not called through"
57
+
58
+ # 2. A second probe must not wrap the wrapper. This is the case that arises
59
+ # when training the first model raises and its probe is never removed.
60
+ b = _GradNormProbe()
61
+ b.install()
62
+ assert torch.nn.utils.clip_grad_norm_ is not pristine
63
+ depth_marker = torch.nn.utils.clip_grad_norm_
64
+ b.install()
65
+ assert torch.nn.utils.clip_grad_norm_ is depth_marker, "probe nested on itself"
66
+
67
+ # 3. Removing the probe that actually patched restores the real function.
68
+ b.remove() # b never owned the patch, so this is a no-op
69
+ assert torch.nn.utils.clip_grad_norm_ is not pristine, "non-owner removed the patch"
70
+ a.remove()
71
+ assert torch.nn.utils.clip_grad_norm_ is pristine, "patch was not restored"
72
+ a.remove() # idempotent
73
+
74
+ # 4. Summaries, including the clipped fraction against the clip of 10.
75
+ c = _GradNormProbe()
76
+ c.install()
77
+ for _ in range(4):
78
+ torch.nn.utils.clip_grad_norm_([], 10.0) # each returns 12.5 > 10
79
+ summary = c.summarise()
80
+ assert summary["grad_norm/mean"] == 12.5, summary
81
+ assert summary["grad_norm/max"] == 12.5, summary
82
+ assert summary["grad_norm/clipped_fraction"] == 1.0, summary
83
+ c.reset()
84
+ assert c.summarise() == {}, "reset did not clear the epoch"
85
+ c.remove()
86
+ assert torch.nn.utils.clip_grad_norm_ is pristine
87
+
88
+ # 5. The default path must not touch torch at all.
89
+ before = torch.nn.utils.clip_grad_norm_
90
+ _GradNormProbe() # constructed but never installed
91
+ assert torch.nn.utils.clip_grad_norm_ is before
92
+
93
+ print("GRAD NORM PROBE TEST PASSED")
94
+
95
+
96
+ if __name__ == "__main__":
97
+ main()