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.
- torchlight-0.1.0/.gitignore +6 -0
- torchlight-0.1.0/LICENSE +21 -0
- torchlight-0.1.0/PKG-INFO +202 -0
- torchlight-0.1.0/README.md +174 -0
- torchlight-0.1.0/pyproject.toml +39 -0
- torchlight-0.1.0/src/torchlight/__init__.py +28 -0
- torchlight-0.1.0/src/torchlight/callbacks.py +183 -0
- torchlight-0.1.0/src/torchlight/checkpoint.py +70 -0
- torchlight-0.1.0/src/torchlight/metrics.py +109 -0
- torchlight-0.1.0/src/torchlight/py.typed +0 -0
- torchlight-0.1.0/src/torchlight/trainer.py +195 -0
- torchlight-0.1.0/src/torchlight/utils.py +87 -0
- torchlight-0.1.0/tests/conftest.py +43 -0
- torchlight-0.1.0/tests/test_callbacks.py +103 -0
- torchlight-0.1.0/tests/test_metrics.py +44 -0
- torchlight-0.1.0/tests/test_trainer.py +137 -0
- torchlight-0.1.0/tests/test_utils_and_checkpoint.py +68 -0
torchlight-0.1.0/LICENSE
ADDED
|
@@ -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}")
|