torchlight 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,6 @@
1
+ .venv/
2
+ __pycache__/
3
+ *.egg-info/
4
+ .pytest_cache/
5
+ build/
6
+ dist/
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 nehz
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,202 @@
1
+ Metadata-Version: 2.5
2
+ Name: torchlight
3
+ Version: 0.1.0
4
+ Summary: Torch Light: a lightweight PyTorch companion with a minimal Trainer, callbacks, metrics and checkpointing.
5
+ Author: nehz
6
+ License: MIT
7
+ License-File: LICENSE
8
+ Keywords: deep-learning,pytorch,torch,trainer,training
9
+ Classifier: Development Status :: 3 - Alpha
10
+ Classifier: Intended Audience :: Developers
11
+ Classifier: Intended Audience :: Science/Research
12
+ Classifier: License :: OSI Approved :: MIT License
13
+ Classifier: Operating System :: OS Independent
14
+ Classifier: Programming Language :: Python :: 3
15
+ Classifier: Programming Language :: Python :: 3 :: Only
16
+ Classifier: Programming Language :: Python :: 3.9
17
+ Classifier: Programming Language :: Python :: 3.10
18
+ Classifier: Programming Language :: Python :: 3.11
19
+ Classifier: Programming Language :: Python :: 3.12
20
+ Classifier: Programming Language :: Python :: 3.13
21
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
22
+ Classifier: Typing :: Typed
23
+ Requires-Python: >=3.9
24
+ Requires-Dist: torch>=2.0
25
+ Provides-Extra: test
26
+ Requires-Dist: pytest>=7; extra == 'test'
27
+ Description-Content-Type: text/markdown
28
+
29
+ # torchlight
30
+
31
+ **Torch Light**: a lightweight companion for PyTorch.
32
+
33
+ `torchlight` gives you a small, readable training loop plus the extras you end
34
+ up rewriting in every project: callbacks, metrics, early stopping,
35
+ checkpointing and seeding/device helpers. It is a few hundred lines of plain
36
+ PyTorch with no other dependencies, so you can read all of it in one sitting.
37
+ When you outgrow it, copy the parts you need.
38
+
39
+ ## Features
40
+
41
+ - **`Trainer`**: `fit` / `evaluate` / `predict` over any iterable of
42
+ `(inputs, targets)` batches, with automatic device placement (including
43
+ nested dicts, lists and tuples), optional gradient-norm clipping, and a
44
+ per-epoch `history`.
45
+ - **Callbacks**: hook into fit start, epoch start, batch end, epoch end and
46
+ fit end. Built in: `EarlyStopping` (with `restore_best_weights`),
47
+ `ModelCheckpoint` (best-only or every epoch, templated filenames), and
48
+ `PrintLogger`.
49
+ - **Metrics**: accumulate-then-compute `Accuracy` (multi-class logits or
50
+ binary thresholds) and `MeanAbsoluteError`. Subclass `Metric` to add your own.
51
+ - **Checkpoints**: `save_checkpoint` / `load_checkpoint` for model and
52
+ optimizer state plus any extra values.
53
+ - **Utilities**: `seed_everything`, `get_device`, `to_device`,
54
+ `count_parameters`.
55
+
56
+ ## Install
57
+
58
+ ```bash
59
+ pip install torchlight
60
+ ```
61
+
62
+ `torch>=2.0` is the only dependency. To use the CPU-only PyTorch build, install
63
+ it first, for example
64
+ `pip install torch --index-url https://download.pytorch.org/whl/cpu`.
65
+
66
+ From a source checkout:
67
+
68
+ ```bash
69
+ python -m venv .venv
70
+ .venv/bin/pip install torch --index-url https://download.pytorch.org/whl/cpu
71
+ .venv/bin/pip install -e ".[test]"
72
+ .venv/bin/python -m pytest
73
+ ```
74
+
75
+ ## Quickstart
76
+
77
+ ```python
78
+ import torch
79
+ from torch import nn
80
+ from torch.utils.data import DataLoader, TensorDataset
81
+
82
+ import torchlight as tl
83
+
84
+ tl.seed_everything(42)
85
+
86
+ # Synthetic two-class data
87
+ x = torch.randn(512, 2)
88
+ y = (x[:, 0] + x[:, 1] > 0).long()
89
+ train = DataLoader(TensorDataset(x[:400], y[:400]), batch_size=32, shuffle=True)
90
+ val = DataLoader(TensorDataset(x[400:], y[400:]), batch_size=64)
91
+
92
+ model = nn.Sequential(nn.Linear(2, 16), nn.ReLU(), nn.Linear(16, 2))
93
+ trainer = tl.Trainer(
94
+ model,
95
+ optimizer=torch.optim.Adam(model.parameters(), lr=1e-2),
96
+ loss_fn=nn.CrossEntropyLoss(),
97
+ metrics={"accuracy": tl.Accuracy()},
98
+ callbacks=[
99
+ tl.EarlyStopping(monitor="val_loss", patience=3, restore_best_weights=True),
100
+ tl.ModelCheckpoint("checkpoints/best.pt", monitor="val_accuracy", mode="max"),
101
+ tl.PrintLogger(),
102
+ ],
103
+ )
104
+
105
+ history = trainer.fit(train, val, epochs=20)
106
+ print(history["val_accuracy"][-1])
107
+
108
+ print(trainer.evaluate(val)) # {'loss': ..., 'accuracy': ...}
109
+ probs = trainer.predict(val).softmax(dim=1)
110
+
111
+ # Later: restore the best checkpoint
112
+ ckpt = tl.load_checkpoint("checkpoints/best.pt", model)
113
+ print(ckpt["epoch"], ckpt["logs"])
114
+ ```
115
+
116
+ ## API overview
117
+
118
+ Everything below is importable from the top-level `torchlight` package.
119
+
120
+ ### `Trainer`
121
+
122
+ ```python
123
+ Trainer(model, optimizer, loss_fn, metrics=None, callbacks=None, device=None, grad_clip=None)
124
+ ```
125
+
126
+ | Member | Description |
127
+ | --- | --- |
128
+ | `fit(train_loader, val_loader=None, epochs=1) -> dict[str, list[float]]` | Train for up to `epochs` epochs and return `history`. Logs `loss` and each metric name; with `val_loader`, also `val_loss` and `val_<metric>`. |
129
+ | `evaluate(loader) -> dict[str, float]` | Loss and metrics under `torch.no_grad()`, in eval mode, with no `val_` prefix. |
130
+ | `predict(loader) -> torch.Tensor` | Model outputs concatenated on the CPU. Batches may be `(inputs, targets)` or bare `inputs`. |
131
+ | `history` | `dict[str, list[float]]` of per-epoch values. |
132
+ | `should_stop` | Set to `True` (usually from a callback) to stop after the current epoch. |
133
+ | `model`, `optimizer`, `loss_fn`, `metrics`, `callbacks`, `device`, `grad_clip`, `epochs` | The configured objects. `epochs` is the value passed to the latest `fit`. |
134
+
135
+ Batches must be `(inputs, targets)` pairs. `inputs` goes straight to
136
+ `model(inputs)`, and both halves are moved to `trainer.device`. The loss is
137
+ averaged over samples, with each batch weighted by its size. `device`
138
+ defaults to `get_device()`.
139
+
140
+ ### Callbacks
141
+
142
+ ```python
143
+ class Callback:
144
+ def on_fit_start(self, trainer): ...
145
+ def on_epoch_start(self, trainer, epoch): ... # epoch is 0-based
146
+ def on_batch_end(self, trainer, batch_idx, loss): ... # loss is a float
147
+ def on_epoch_end(self, trainer, epoch, logs): ... # logs: dict[str, float]
148
+ def on_fit_end(self, trainer): ... # also runs after an early stop
149
+ ```
150
+
151
+ - `EarlyStopping(monitor="val_loss", patience=3, mode="min", min_delta=0.0, restore_best_weights=False)`
152
+ sets `trainer.should_stop` after `patience` epochs without improvement.
153
+ Attributes: `best`, `best_epoch`, `wait`, `stopped_epoch`.
154
+ - `ModelCheckpoint(path, monitor="val_loss", mode="min", save_best_only=True)`
155
+ writes `save_checkpoint(path, model, optimizer, epoch=..., logs=...)`.
156
+ `path` may use `{epoch}` and any log key as format fields, e.g.
157
+ `"ckpt/ep{epoch:02d}-{val_loss:.3f}.pt"`. Attributes: `best`, `best_path`.
158
+ - `PrintLogger(precision=4)` prints `epoch 3/10 loss=0.1234 val_loss=...`.
159
+
160
+ If a monitored key is missing from the epoch logs (for example `val_loss`
161
+ without a `val_loader`), the callback raises `KeyError`.
162
+
163
+ ### Metrics
164
+
165
+ ```python
166
+ class Metric:
167
+ def update(self, preds, targets) -> None: ...
168
+ def compute(self) -> float: ...
169
+ def reset(self) -> None: ...
170
+ ```
171
+
172
+ - `Accuracy(threshold=0.5)`: uses `argmax(dim=1)` when `preds` has one more
173
+ dimension than `targets`. Otherwise `preds >= threshold`; use
174
+ `threshold=0.0` for raw binary logits.
175
+ - `MeanAbsoluteError()`: `preds` and `targets` must have the same shape.
176
+ - `Mean()`: weighted running mean, `update(value, weight=1)`. The trainer uses
177
+ it for the loss.
178
+
179
+ `compute()` returns `nan` if nothing has been accumulated.
180
+
181
+ ### Checkpoints
182
+
183
+ - `save_checkpoint(path, model, optimizer=None, **extra) -> pathlib.Path`
184
+ saves `{"model": ..., "optimizer": ... or None, **extra}` and creates
185
+ parent directories.
186
+ - `load_checkpoint(path, model, optimizer=None, map_location="cpu") -> dict`
187
+ loads weights (and optimizer state, if both are present) in place and
188
+ returns the full dict. It uses `torch.load(..., weights_only=True)`.
189
+
190
+ ### Utilities
191
+
192
+ - `seed_everything(seed, deterministic=False) -> int` seeds `random`, NumPy
193
+ (if installed) and torch (CPU and CUDA).
194
+ - `get_device(prefer=None) -> torch.device` returns `prefer` if given,
195
+ otherwise the first available of CUDA, MPS and CPU.
196
+ - `to_device(data, device)` recursively moves tensors in dicts, lists, tuples
197
+ and named tuples.
198
+ - `count_parameters(model, trainable_only=True) -> int`.
199
+
200
+ ## License
201
+
202
+ MIT
@@ -0,0 +1,174 @@
1
+ # torchlight
2
+
3
+ **Torch Light**: a lightweight companion for PyTorch.
4
+
5
+ `torchlight` gives you a small, readable training loop plus the extras you end
6
+ up rewriting in every project: callbacks, metrics, early stopping,
7
+ checkpointing and seeding/device helpers. It is a few hundred lines of plain
8
+ PyTorch with no other dependencies, so you can read all of it in one sitting.
9
+ When you outgrow it, copy the parts you need.
10
+
11
+ ## Features
12
+
13
+ - **`Trainer`**: `fit` / `evaluate` / `predict` over any iterable of
14
+ `(inputs, targets)` batches, with automatic device placement (including
15
+ nested dicts, lists and tuples), optional gradient-norm clipping, and a
16
+ per-epoch `history`.
17
+ - **Callbacks**: hook into fit start, epoch start, batch end, epoch end and
18
+ fit end. Built in: `EarlyStopping` (with `restore_best_weights`),
19
+ `ModelCheckpoint` (best-only or every epoch, templated filenames), and
20
+ `PrintLogger`.
21
+ - **Metrics**: accumulate-then-compute `Accuracy` (multi-class logits or
22
+ binary thresholds) and `MeanAbsoluteError`. Subclass `Metric` to add your own.
23
+ - **Checkpoints**: `save_checkpoint` / `load_checkpoint` for model and
24
+ optimizer state plus any extra values.
25
+ - **Utilities**: `seed_everything`, `get_device`, `to_device`,
26
+ `count_parameters`.
27
+
28
+ ## Install
29
+
30
+ ```bash
31
+ pip install torchlight
32
+ ```
33
+
34
+ `torch>=2.0` is the only dependency. To use the CPU-only PyTorch build, install
35
+ it first, for example
36
+ `pip install torch --index-url https://download.pytorch.org/whl/cpu`.
37
+
38
+ From a source checkout:
39
+
40
+ ```bash
41
+ python -m venv .venv
42
+ .venv/bin/pip install torch --index-url https://download.pytorch.org/whl/cpu
43
+ .venv/bin/pip install -e ".[test]"
44
+ .venv/bin/python -m pytest
45
+ ```
46
+
47
+ ## Quickstart
48
+
49
+ ```python
50
+ import torch
51
+ from torch import nn
52
+ from torch.utils.data import DataLoader, TensorDataset
53
+
54
+ import torchlight as tl
55
+
56
+ tl.seed_everything(42)
57
+
58
+ # Synthetic two-class data
59
+ x = torch.randn(512, 2)
60
+ y = (x[:, 0] + x[:, 1] > 0).long()
61
+ train = DataLoader(TensorDataset(x[:400], y[:400]), batch_size=32, shuffle=True)
62
+ val = DataLoader(TensorDataset(x[400:], y[400:]), batch_size=64)
63
+
64
+ model = nn.Sequential(nn.Linear(2, 16), nn.ReLU(), nn.Linear(16, 2))
65
+ trainer = tl.Trainer(
66
+ model,
67
+ optimizer=torch.optim.Adam(model.parameters(), lr=1e-2),
68
+ loss_fn=nn.CrossEntropyLoss(),
69
+ metrics={"accuracy": tl.Accuracy()},
70
+ callbacks=[
71
+ tl.EarlyStopping(monitor="val_loss", patience=3, restore_best_weights=True),
72
+ tl.ModelCheckpoint("checkpoints/best.pt", monitor="val_accuracy", mode="max"),
73
+ tl.PrintLogger(),
74
+ ],
75
+ )
76
+
77
+ history = trainer.fit(train, val, epochs=20)
78
+ print(history["val_accuracy"][-1])
79
+
80
+ print(trainer.evaluate(val)) # {'loss': ..., 'accuracy': ...}
81
+ probs = trainer.predict(val).softmax(dim=1)
82
+
83
+ # Later: restore the best checkpoint
84
+ ckpt = tl.load_checkpoint("checkpoints/best.pt", model)
85
+ print(ckpt["epoch"], ckpt["logs"])
86
+ ```
87
+
88
+ ## API overview
89
+
90
+ Everything below is importable from the top-level `torchlight` package.
91
+
92
+ ### `Trainer`
93
+
94
+ ```python
95
+ Trainer(model, optimizer, loss_fn, metrics=None, callbacks=None, device=None, grad_clip=None)
96
+ ```
97
+
98
+ | Member | Description |
99
+ | --- | --- |
100
+ | `fit(train_loader, val_loader=None, epochs=1) -> dict[str, list[float]]` | Train for up to `epochs` epochs and return `history`. Logs `loss` and each metric name; with `val_loader`, also `val_loss` and `val_<metric>`. |
101
+ | `evaluate(loader) -> dict[str, float]` | Loss and metrics under `torch.no_grad()`, in eval mode, with no `val_` prefix. |
102
+ | `predict(loader) -> torch.Tensor` | Model outputs concatenated on the CPU. Batches may be `(inputs, targets)` or bare `inputs`. |
103
+ | `history` | `dict[str, list[float]]` of per-epoch values. |
104
+ | `should_stop` | Set to `True` (usually from a callback) to stop after the current epoch. |
105
+ | `model`, `optimizer`, `loss_fn`, `metrics`, `callbacks`, `device`, `grad_clip`, `epochs` | The configured objects. `epochs` is the value passed to the latest `fit`. |
106
+
107
+ Batches must be `(inputs, targets)` pairs. `inputs` goes straight to
108
+ `model(inputs)`, and both halves are moved to `trainer.device`. The loss is
109
+ averaged over samples, with each batch weighted by its size. `device`
110
+ defaults to `get_device()`.
111
+
112
+ ### Callbacks
113
+
114
+ ```python
115
+ class Callback:
116
+ def on_fit_start(self, trainer): ...
117
+ def on_epoch_start(self, trainer, epoch): ... # epoch is 0-based
118
+ def on_batch_end(self, trainer, batch_idx, loss): ... # loss is a float
119
+ def on_epoch_end(self, trainer, epoch, logs): ... # logs: dict[str, float]
120
+ def on_fit_end(self, trainer): ... # also runs after an early stop
121
+ ```
122
+
123
+ - `EarlyStopping(monitor="val_loss", patience=3, mode="min", min_delta=0.0, restore_best_weights=False)`
124
+ sets `trainer.should_stop` after `patience` epochs without improvement.
125
+ Attributes: `best`, `best_epoch`, `wait`, `stopped_epoch`.
126
+ - `ModelCheckpoint(path, monitor="val_loss", mode="min", save_best_only=True)`
127
+ writes `save_checkpoint(path, model, optimizer, epoch=..., logs=...)`.
128
+ `path` may use `{epoch}` and any log key as format fields, e.g.
129
+ `"ckpt/ep{epoch:02d}-{val_loss:.3f}.pt"`. Attributes: `best`, `best_path`.
130
+ - `PrintLogger(precision=4)` prints `epoch 3/10 loss=0.1234 val_loss=...`.
131
+
132
+ If a monitored key is missing from the epoch logs (for example `val_loss`
133
+ without a `val_loader`), the callback raises `KeyError`.
134
+
135
+ ### Metrics
136
+
137
+ ```python
138
+ class Metric:
139
+ def update(self, preds, targets) -> None: ...
140
+ def compute(self) -> float: ...
141
+ def reset(self) -> None: ...
142
+ ```
143
+
144
+ - `Accuracy(threshold=0.5)`: uses `argmax(dim=1)` when `preds` has one more
145
+ dimension than `targets`. Otherwise `preds >= threshold`; use
146
+ `threshold=0.0` for raw binary logits.
147
+ - `MeanAbsoluteError()`: `preds` and `targets` must have the same shape.
148
+ - `Mean()`: weighted running mean, `update(value, weight=1)`. The trainer uses
149
+ it for the loss.
150
+
151
+ `compute()` returns `nan` if nothing has been accumulated.
152
+
153
+ ### Checkpoints
154
+
155
+ - `save_checkpoint(path, model, optimizer=None, **extra) -> pathlib.Path`
156
+ saves `{"model": ..., "optimizer": ... or None, **extra}` and creates
157
+ parent directories.
158
+ - `load_checkpoint(path, model, optimizer=None, map_location="cpu") -> dict`
159
+ loads weights (and optimizer state, if both are present) in place and
160
+ returns the full dict. It uses `torch.load(..., weights_only=True)`.
161
+
162
+ ### Utilities
163
+
164
+ - `seed_everything(seed, deterministic=False) -> int` seeds `random`, NumPy
165
+ (if installed) and torch (CPU and CUDA).
166
+ - `get_device(prefer=None) -> torch.device` returns `prefer` if given,
167
+ otherwise the first available of CUDA, MPS and CPU.
168
+ - `to_device(data, device)` recursively moves tensors in dicts, lists, tuples
169
+ and named tuples.
170
+ - `count_parameters(model, trainable_only=True) -> int`.
171
+
172
+ ## License
173
+
174
+ MIT
@@ -0,0 +1,39 @@
1
+ [build-system]
2
+ requires = ["hatchling>=1.18"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "torchlight"
7
+ version = "0.1.0"
8
+ description = "Torch Light: a lightweight PyTorch companion with a minimal Trainer, callbacks, metrics and checkpointing."
9
+ readme = "README.md"
10
+ requires-python = ">=3.9"
11
+ license = { text = "MIT" }
12
+ authors = [{ name = "nehz" }]
13
+ keywords = ["pytorch", "torch", "training", "trainer", "deep-learning"]
14
+ classifiers = [
15
+ "Development Status :: 3 - Alpha",
16
+ "Intended Audience :: Developers",
17
+ "Intended Audience :: Science/Research",
18
+ "License :: OSI Approved :: MIT License",
19
+ "Operating System :: OS Independent",
20
+ "Programming Language :: Python :: 3",
21
+ "Programming Language :: Python :: 3 :: Only",
22
+ "Programming Language :: Python :: 3.9",
23
+ "Programming Language :: Python :: 3.10",
24
+ "Programming Language :: Python :: 3.11",
25
+ "Programming Language :: Python :: 3.12",
26
+ "Programming Language :: Python :: 3.13",
27
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
28
+ "Typing :: Typed",
29
+ ]
30
+ dependencies = ["torch>=2.0"]
31
+
32
+ [project.optional-dependencies]
33
+ test = ["pytest>=7"]
34
+
35
+ [tool.hatch.build.targets.wheel]
36
+ packages = ["src/torchlight"]
37
+
38
+ [tool.pytest.ini_options]
39
+ testpaths = ["tests"]
@@ -0,0 +1,28 @@
1
+ """torchlight: a lightweight companion for PyTorch training loops."""
2
+
3
+ from .callbacks import Callback, EarlyStopping, ModelCheckpoint, PrintLogger
4
+ from .checkpoint import load_checkpoint, save_checkpoint
5
+ from .metrics import Accuracy, Mean, MeanAbsoluteError, Metric
6
+ from .trainer import Trainer
7
+ from .utils import count_parameters, get_device, seed_everything, to_device
8
+
9
+ __version__ = "0.1.0"
10
+
11
+ __all__ = [
12
+ "Trainer",
13
+ "Callback",
14
+ "EarlyStopping",
15
+ "ModelCheckpoint",
16
+ "PrintLogger",
17
+ "Metric",
18
+ "Mean",
19
+ "Accuracy",
20
+ "MeanAbsoluteError",
21
+ "save_checkpoint",
22
+ "load_checkpoint",
23
+ "seed_everything",
24
+ "get_device",
25
+ "to_device",
26
+ "count_parameters",
27
+ "__version__",
28
+ ]
@@ -0,0 +1,183 @@
1
+ """Callbacks that hook into :meth:`torchlight.Trainer.fit`."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import copy
6
+ import math
7
+ import os
8
+ from pathlib import Path
9
+ from typing import TYPE_CHECKING, Any, Literal
10
+
11
+ from .checkpoint import save_checkpoint
12
+
13
+ if TYPE_CHECKING:
14
+ from .trainer import Trainer
15
+
16
+ __all__ = ["Callback", "EarlyStopping", "ModelCheckpoint", "PrintLogger"]
17
+
18
+ Mode = Literal["min", "max"]
19
+
20
+
21
+ class Callback:
22
+ """Base class for callbacks. Override any subset of the hooks.
23
+
24
+ ``logs`` passed to :meth:`on_epoch_end` is the dict of epoch results, e.g.
25
+ ``{"loss": 0.31, "accuracy": 0.9, "val_loss": 0.35, "val_accuracy": 0.88}``.
26
+ Set ``trainer.should_stop = True`` to end training after the current epoch.
27
+ """
28
+
29
+ def on_fit_start(self, trainer: Trainer) -> None:
30
+ """Called once before the first epoch."""
31
+
32
+ def on_epoch_start(self, trainer: Trainer, epoch: int) -> None:
33
+ """Called before each epoch (``epoch`` is 0-based)."""
34
+
35
+ def on_batch_end(self, trainer: Trainer, batch_idx: int, loss: float) -> None:
36
+ """Called after each training step with that batch's loss."""
37
+
38
+ def on_epoch_end(self, trainer: Trainer, epoch: int, logs: dict[str, float]) -> None:
39
+ """Called after training (and validation, if any) for each epoch."""
40
+
41
+ def on_fit_end(self, trainer: Trainer) -> None:
42
+ """Called once after the last epoch, including after an early stop."""
43
+
44
+
45
+ def _monitored(logs: dict[str, float], monitor: str, owner: str) -> float:
46
+ if monitor not in logs:
47
+ raise KeyError(
48
+ f"{owner}: monitored key {monitor!r} not in epoch logs "
49
+ f"(available: {sorted(logs)})"
50
+ )
51
+ return logs[monitor]
52
+
53
+
54
+ def _improved(current: float, best: float, mode: Mode, min_delta: float) -> bool:
55
+ if mode == "min":
56
+ return current < best - min_delta
57
+ return current > best + min_delta
58
+
59
+
60
+ def _check_mode(mode: str) -> None:
61
+ if mode not in ("min", "max"):
62
+ raise ValueError(f"mode must be 'min' or 'max', got {mode!r}")
63
+
64
+
65
+ class EarlyStopping(Callback):
66
+ """Stop training when a monitored value stops improving.
67
+
68
+ Args:
69
+ monitor: Key in the epoch logs to watch, e.g. ``"val_loss"``.
70
+ patience: Epochs without improvement to tolerate before stopping.
71
+ mode: ``"min"`` if lower is better, ``"max"`` if higher is better.
72
+ min_delta: Minimum change that counts as an improvement.
73
+ restore_best_weights: On stop (or fit end), load the weights from the
74
+ best epoch back into the model.
75
+ """
76
+
77
+ def __init__(
78
+ self,
79
+ monitor: str = "val_loss",
80
+ patience: int = 3,
81
+ mode: Mode = "min",
82
+ min_delta: float = 0.0,
83
+ restore_best_weights: bool = False,
84
+ ) -> None:
85
+ _check_mode(mode)
86
+ self.monitor = monitor
87
+ self.patience = patience
88
+ self.mode: Mode = mode
89
+ self.min_delta = min_delta
90
+ self.restore_best_weights = restore_best_weights
91
+ self.best: float = math.inf if mode == "min" else -math.inf
92
+ self.best_epoch: int = -1
93
+ self.wait = 0
94
+ self.stopped_epoch: int | None = None
95
+ self._best_state: dict[str, Any] | None = None
96
+
97
+ def on_fit_start(self, trainer: Trainer) -> None:
98
+ self.best = math.inf if self.mode == "min" else -math.inf
99
+ self.best_epoch = -1
100
+ self.wait = 0
101
+ self.stopped_epoch = None
102
+ self._best_state = None
103
+
104
+ def on_epoch_end(self, trainer: Trainer, epoch: int, logs: dict[str, float]) -> None:
105
+ current = _monitored(logs, self.monitor, "EarlyStopping")
106
+ if _improved(current, self.best, self.mode, self.min_delta):
107
+ self.best = current
108
+ self.best_epoch = epoch
109
+ self.wait = 0
110
+ if self.restore_best_weights:
111
+ self._best_state = copy.deepcopy(trainer.model.state_dict())
112
+ return
113
+ self.wait += 1
114
+ if self.wait >= self.patience:
115
+ self.stopped_epoch = epoch
116
+ trainer.should_stop = True
117
+
118
+ def on_fit_end(self, trainer: Trainer) -> None:
119
+ if self.restore_best_weights and self._best_state is not None:
120
+ trainer.model.load_state_dict(self._best_state)
121
+
122
+
123
+ class ModelCheckpoint(Callback):
124
+ """Save a checkpoint at the end of epochs.
125
+
126
+ Args:
127
+ path: File to write. May contain ``{epoch}`` and any log key as format
128
+ fields, e.g. ``"ckpt/epoch{epoch:02d}-{val_loss:.3f}.pt"``.
129
+ monitor: Key in the epoch logs used to decide what is "best".
130
+ mode: ``"min"`` or ``"max"``.
131
+ save_best_only: If True, only write when ``monitor`` improves;
132
+ otherwise write every epoch.
133
+
134
+ After fitting, :attr:`best_path` holds the last file written for the best
135
+ value (or the latest file when ``save_best_only`` is False).
136
+ """
137
+
138
+ def __init__(
139
+ self,
140
+ path: str | os.PathLike[str],
141
+ monitor: str = "val_loss",
142
+ mode: Mode = "min",
143
+ save_best_only: bool = True,
144
+ ) -> None:
145
+ _check_mode(mode)
146
+ self.path = str(path)
147
+ self.monitor = monitor
148
+ self.mode: Mode = mode
149
+ self.save_best_only = save_best_only
150
+ self.best: float = math.inf if mode == "min" else -math.inf
151
+ self.best_path: Path | None = None
152
+
153
+ def on_fit_start(self, trainer: Trainer) -> None:
154
+ self.best = math.inf if self.mode == "min" else -math.inf
155
+ self.best_path = None
156
+
157
+ def on_epoch_end(self, trainer: Trainer, epoch: int, logs: dict[str, float]) -> None:
158
+ current = _monitored(logs, self.monitor, "ModelCheckpoint")
159
+ improved = _improved(current, self.best, self.mode, 0.0)
160
+ if self.save_best_only and not improved:
161
+ return
162
+ if improved:
163
+ self.best = current
164
+ target = self.path.format(epoch=epoch, **logs)
165
+ self.best_path = save_checkpoint(
166
+ target,
167
+ trainer.model,
168
+ trainer.optimizer,
169
+ epoch=epoch,
170
+ logs=dict(logs),
171
+ )
172
+
173
+
174
+ class PrintLogger(Callback):
175
+ """Print one line of epoch results, e.g. ``epoch 3/10 loss=0.1234 ...``."""
176
+
177
+ def __init__(self, precision: int = 4) -> None:
178
+ self.precision = precision
179
+
180
+ def on_epoch_end(self, trainer: Trainer, epoch: int, logs: dict[str, float]) -> None:
181
+ parts = " ".join(f"{k}={v:.{self.precision}f}" for k, v in logs.items())
182
+ total = trainer.epochs if trainer.epochs is not None else "?"
183
+ print(f"epoch {epoch + 1}/{total} {parts}")