Make the gradient-norm probe safe to leave on
Browse filesNo 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 +23 -8
- notebooks/colab_train.ipynb +7 -2
- src/train/train.py +27 -6
- src/train/wandb_logger.py +66 -17
- tests/test_grad_norm_probe.py +97 -0
|
@@ -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 |
-
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 432 |
-
|
| 433 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
@@ -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 |
{
|
|
@@ -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"):
|
|
@@ -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.
|
| 51 |
|
| 52 |
def install(self) -> None:
|
| 53 |
import torch
|
| 54 |
|
| 55 |
-
|
|
|
|
|
|
|
|
|
|
| 56 |
return
|
| 57 |
-
self._original =
|
| 58 |
|
| 59 |
def recording_clip_grad_norm(parameters, max_norm, *args, **kwargs):
|
| 60 |
total_norm = self._original(parameters, max_norm, *args, **kwargs)
|
| 61 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 77 |
return {}
|
| 78 |
-
norms = np.array(self.
|
| 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.
|
| 90 |
|
| 91 |
|
| 92 |
-
def attach(model, project: str, run_name: str, extra_config: dict | 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
|
|
|
|
| 110 |
|
| 111 |
def on_train_epoch_start(trainer):
|
| 112 |
-
probe
|
|
|
|
| 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 |
-
|
| 126 |
-
|
| 127 |
-
|
|
|
|
| 128 |
|
| 129 |
def on_train_end(trainer):
|
| 130 |
-
|
| 131 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
@@ -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()
|