torchlight 0.1.0__py3-none-any.whl
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/__init__.py +28 -0
- torchlight/callbacks.py +183 -0
- torchlight/checkpoint.py +70 -0
- torchlight/metrics.py +109 -0
- torchlight/py.typed +0 -0
- torchlight/trainer.py +195 -0
- torchlight/utils.py +87 -0
- torchlight-0.1.0.dist-info/METADATA +202 -0
- torchlight-0.1.0.dist-info/RECORD +11 -0
- torchlight-0.1.0.dist-info/WHEEL +4 -0
- torchlight-0.1.0.dist-info/licenses/LICENSE +21 -0
torchlight/__init__.py
ADDED
|
@@ -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
|
+
]
|
torchlight/callbacks.py
ADDED
|
@@ -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}")
|
torchlight/checkpoint.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
"""Saving and restoring model/optimizer state."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
from torch import nn
|
|
11
|
+
|
|
12
|
+
__all__ = ["save_checkpoint", "load_checkpoint"]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def save_checkpoint(
|
|
16
|
+
path: str | os.PathLike[str],
|
|
17
|
+
model: nn.Module,
|
|
18
|
+
optimizer: torch.optim.Optimizer | None = None,
|
|
19
|
+
**extra: Any,
|
|
20
|
+
) -> Path:
|
|
21
|
+
"""Write a checkpoint dictionary to ``path``.
|
|
22
|
+
|
|
23
|
+
The file holds ``{"model": state_dict, "optimizer": state_dict | None,
|
|
24
|
+
**extra}``. Parent directories are created as needed.
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
path: Destination file.
|
|
28
|
+
model: Model whose ``state_dict`` is saved.
|
|
29
|
+
optimizer: Optional optimizer whose ``state_dict`` is saved.
|
|
30
|
+
**extra: Additional picklable values, e.g. ``epoch=3``.
|
|
31
|
+
|
|
32
|
+
Returns:
|
|
33
|
+
The path written, as a :class:`~pathlib.Path`.
|
|
34
|
+
"""
|
|
35
|
+
path = Path(path)
|
|
36
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
37
|
+
payload: dict[str, Any] = {
|
|
38
|
+
"model": model.state_dict(),
|
|
39
|
+
"optimizer": optimizer.state_dict() if optimizer is not None else None,
|
|
40
|
+
**extra,
|
|
41
|
+
}
|
|
42
|
+
torch.save(payload, path)
|
|
43
|
+
return path
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def load_checkpoint(
|
|
47
|
+
path: str | os.PathLike[str],
|
|
48
|
+
model: nn.Module,
|
|
49
|
+
optimizer: torch.optim.Optimizer | None = None,
|
|
50
|
+
map_location: str | torch.device = "cpu",
|
|
51
|
+
) -> dict[str, Any]:
|
|
52
|
+
"""Restore state written by :func:`save_checkpoint`.
|
|
53
|
+
|
|
54
|
+
Args:
|
|
55
|
+
path: Checkpoint file.
|
|
56
|
+
model: Model to load weights into (in place).
|
|
57
|
+
optimizer: If given and the checkpoint contains optimizer state, it is
|
|
58
|
+
loaded too.
|
|
59
|
+
map_location: Passed to :func:`torch.load`.
|
|
60
|
+
|
|
61
|
+
Returns:
|
|
62
|
+
The full checkpoint dictionary, so callers can read ``extra`` values.
|
|
63
|
+
"""
|
|
64
|
+
checkpoint: dict[str, Any] = torch.load(
|
|
65
|
+
Path(path), map_location=map_location, weights_only=True
|
|
66
|
+
)
|
|
67
|
+
model.load_state_dict(checkpoint["model"])
|
|
68
|
+
if optimizer is not None and checkpoint.get("optimizer") is not None:
|
|
69
|
+
optimizer.load_state_dict(checkpoint["optimizer"])
|
|
70
|
+
return checkpoint
|
torchlight/metrics.py
ADDED
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
"""Stateful, accumulate-then-compute metrics."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
__all__ = ["Metric", "Mean", "Accuracy", "MeanAbsoluteError"]
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Metric:
|
|
11
|
+
"""Base class for metrics accumulated over batches.
|
|
12
|
+
|
|
13
|
+
Subclasses implement :meth:`update`, :meth:`compute` and :meth:`reset`.
|
|
14
|
+
The :class:`~torchlight.Trainer` calls ``reset()`` at the start of each
|
|
15
|
+
epoch, ``update(outputs, targets)`` after each batch, and ``compute()`` at
|
|
16
|
+
the end of the epoch.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
def update(self, preds: torch.Tensor, targets: torch.Tensor) -> None:
|
|
20
|
+
"""Accumulate statistics for one batch."""
|
|
21
|
+
raise NotImplementedError
|
|
22
|
+
|
|
23
|
+
def compute(self) -> float:
|
|
24
|
+
"""Return the metric value over everything seen since the last reset."""
|
|
25
|
+
raise NotImplementedError
|
|
26
|
+
|
|
27
|
+
def reset(self) -> None:
|
|
28
|
+
"""Clear accumulated state."""
|
|
29
|
+
raise NotImplementedError
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class Mean:
|
|
33
|
+
"""Weighted running mean of scalar values (used for the epoch loss)."""
|
|
34
|
+
|
|
35
|
+
def __init__(self) -> None:
|
|
36
|
+
self.reset()
|
|
37
|
+
|
|
38
|
+
def update(self, value: float, weight: int = 1) -> None:
|
|
39
|
+
"""Add ``value`` with the given ``weight`` (e.g. the batch size)."""
|
|
40
|
+
self._total += float(value) * weight
|
|
41
|
+
self._count += weight
|
|
42
|
+
|
|
43
|
+
def compute(self) -> float:
|
|
44
|
+
"""Return the mean, or ``nan`` if nothing has been added."""
|
|
45
|
+
return self._total / self._count if self._count else float("nan")
|
|
46
|
+
|
|
47
|
+
def reset(self) -> None:
|
|
48
|
+
"""Clear accumulated state."""
|
|
49
|
+
self._total = 0.0
|
|
50
|
+
self._count = 0
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class Accuracy(Metric):
|
|
54
|
+
"""Classification accuracy.
|
|
55
|
+
|
|
56
|
+
* If ``preds`` has one more dimension than ``targets`` (e.g. logits of
|
|
57
|
+
shape ``(N, C)`` against labels of shape ``(N,)``), the predicted class
|
|
58
|
+
is ``preds.argmax(dim=1)``.
|
|
59
|
+
* If shapes match (binary case), a prediction is positive when
|
|
60
|
+
``preds >= threshold``. Pass probabilities, or set ``threshold=0.0`` for
|
|
61
|
+
raw logits.
|
|
62
|
+
"""
|
|
63
|
+
|
|
64
|
+
def __init__(self, threshold: float = 0.5) -> None:
|
|
65
|
+
self.threshold = threshold
|
|
66
|
+
self.reset()
|
|
67
|
+
|
|
68
|
+
def update(self, preds: torch.Tensor, targets: torch.Tensor) -> None:
|
|
69
|
+
if preds.dim() == targets.dim() + 1:
|
|
70
|
+
predicted = preds.argmax(dim=1)
|
|
71
|
+
elif preds.shape == targets.shape:
|
|
72
|
+
predicted = (preds >= self.threshold).to(targets.dtype)
|
|
73
|
+
else:
|
|
74
|
+
raise ValueError(
|
|
75
|
+
f"Accuracy: incompatible shapes preds={tuple(preds.shape)} "
|
|
76
|
+
f"targets={tuple(targets.shape)}"
|
|
77
|
+
)
|
|
78
|
+
self._correct += int((predicted == targets).sum().item())
|
|
79
|
+
self._total += targets.numel()
|
|
80
|
+
|
|
81
|
+
def compute(self) -> float:
|
|
82
|
+
return self._correct / self._total if self._total else float("nan")
|
|
83
|
+
|
|
84
|
+
def reset(self) -> None:
|
|
85
|
+
self._correct = 0
|
|
86
|
+
self._total = 0
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
class MeanAbsoluteError(Metric):
|
|
90
|
+
"""Mean absolute error for regression outputs."""
|
|
91
|
+
|
|
92
|
+
def __init__(self) -> None:
|
|
93
|
+
self.reset()
|
|
94
|
+
|
|
95
|
+
def update(self, preds: torch.Tensor, targets: torch.Tensor) -> None:
|
|
96
|
+
if preds.shape != targets.shape:
|
|
97
|
+
raise ValueError(
|
|
98
|
+
f"MeanAbsoluteError: shape mismatch preds={tuple(preds.shape)} "
|
|
99
|
+
f"targets={tuple(targets.shape)}"
|
|
100
|
+
)
|
|
101
|
+
self._abs_sum += float((preds.detach() - targets).abs().sum().item())
|
|
102
|
+
self._count += targets.numel()
|
|
103
|
+
|
|
104
|
+
def compute(self) -> float:
|
|
105
|
+
return self._abs_sum / self._count if self._count else float("nan")
|
|
106
|
+
|
|
107
|
+
def reset(self) -> None:
|
|
108
|
+
self._abs_sum = 0.0
|
|
109
|
+
self._count = 0
|
torchlight/py.typed
ADDED
|
File without changes
|
torchlight/trainer.py
ADDED
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
"""A minimal, explicit training loop."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable, Iterable, Mapping, Sequence
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
import torch
|
|
9
|
+
from torch import nn
|
|
10
|
+
|
|
11
|
+
from .callbacks import Callback
|
|
12
|
+
from .metrics import Mean, Metric
|
|
13
|
+
from .utils import get_device, to_device
|
|
14
|
+
|
|
15
|
+
__all__ = ["Trainer"]
|
|
16
|
+
|
|
17
|
+
Batch = Any
|
|
18
|
+
LossFn = Callable[[torch.Tensor, torch.Tensor], torch.Tensor]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Trainer:
|
|
22
|
+
"""Train, evaluate and predict with a ``torch.nn.Module``.
|
|
23
|
+
|
|
24
|
+
Each batch from a data loader must be an ``(inputs, targets)`` pair, where
|
|
25
|
+
``inputs`` is passed straight to ``model(inputs)``. ``inputs`` and
|
|
26
|
+
``targets`` may be tensors or nested lists/tuples/dicts of tensors; they
|
|
27
|
+
are moved to :attr:`device` automatically. For :meth:`predict`, a batch may
|
|
28
|
+
also be bare ``inputs``.
|
|
29
|
+
|
|
30
|
+
Args:
|
|
31
|
+
model: The model to train. It is moved to ``device``.
|
|
32
|
+
optimizer: Optimizer over the model's parameters.
|
|
33
|
+
loss_fn: Callable ``loss_fn(outputs, targets) -> scalar tensor``.
|
|
34
|
+
metrics: Mapping of name to :class:`~torchlight.metrics.Metric`,
|
|
35
|
+
reported for training and (prefixed ``val_``) validation.
|
|
36
|
+
callbacks: Callbacks invoked during :meth:`fit`.
|
|
37
|
+
device: Device to run on; defaults to :func:`~torchlight.get_device`.
|
|
38
|
+
grad_clip: If set, clip the global gradient norm to this value.
|
|
39
|
+
|
|
40
|
+
Attributes:
|
|
41
|
+
history: Per-epoch values for every log key, e.g.
|
|
42
|
+
``trainer.history["val_loss"]``.
|
|
43
|
+
should_stop: Set to True (typically by a callback) to stop
|
|
44
|
+
:meth:`fit` after the current epoch.
|
|
45
|
+
epochs: Number of epochs requested by the running/last :meth:`fit`.
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
def __init__(
|
|
49
|
+
self,
|
|
50
|
+
model: nn.Module,
|
|
51
|
+
optimizer: torch.optim.Optimizer,
|
|
52
|
+
loss_fn: LossFn,
|
|
53
|
+
metrics: Mapping[str, Metric] | None = None,
|
|
54
|
+
callbacks: Sequence[Callback] | None = None,
|
|
55
|
+
device: str | torch.device | None = None,
|
|
56
|
+
grad_clip: float | None = None,
|
|
57
|
+
) -> None:
|
|
58
|
+
self.device = get_device(device)
|
|
59
|
+
self.model = model.to(self.device)
|
|
60
|
+
self.optimizer = optimizer
|
|
61
|
+
self.loss_fn = loss_fn
|
|
62
|
+
self.metrics: dict[str, Metric] = dict(metrics or {})
|
|
63
|
+
self.callbacks: list[Callback] = list(callbacks or [])
|
|
64
|
+
self.grad_clip = grad_clip
|
|
65
|
+
self.history: dict[str, list[float]] = {}
|
|
66
|
+
self.should_stop = False
|
|
67
|
+
self.epochs: int | None = None
|
|
68
|
+
|
|
69
|
+
# ------------------------------------------------------------------ fit
|
|
70
|
+
def fit(
|
|
71
|
+
self,
|
|
72
|
+
train_loader: Iterable[Batch],
|
|
73
|
+
val_loader: Iterable[Batch] | None = None,
|
|
74
|
+
epochs: int = 1,
|
|
75
|
+
) -> dict[str, list[float]]:
|
|
76
|
+
"""Train for up to ``epochs`` epochs.
|
|
77
|
+
|
|
78
|
+
Logged keys per epoch are ``loss`` and each metric name, plus
|
|
79
|
+
``val_loss`` and ``val_<metric>`` when ``val_loader`` is given.
|
|
80
|
+
|
|
81
|
+
Returns:
|
|
82
|
+
:attr:`history`, mapping each log key to its per-epoch values.
|
|
83
|
+
"""
|
|
84
|
+
if epochs < 1:
|
|
85
|
+
raise ValueError(f"epochs must be >= 1, got {epochs}")
|
|
86
|
+
self.epochs = epochs
|
|
87
|
+
self.should_stop = False
|
|
88
|
+
self.history = {}
|
|
89
|
+
self._call("on_fit_start")
|
|
90
|
+
for epoch in range(epochs):
|
|
91
|
+
self._call("on_epoch_start", epoch)
|
|
92
|
+
logs = self._train_epoch(train_loader)
|
|
93
|
+
if val_loader is not None:
|
|
94
|
+
val = self.evaluate(val_loader)
|
|
95
|
+
logs.update({f"val_{k}": v for k, v in val.items()})
|
|
96
|
+
for key, value in logs.items():
|
|
97
|
+
self.history.setdefault(key, []).append(value)
|
|
98
|
+
self._call("on_epoch_end", epoch, logs)
|
|
99
|
+
if self.should_stop:
|
|
100
|
+
break
|
|
101
|
+
self._call("on_fit_end")
|
|
102
|
+
return self.history
|
|
103
|
+
|
|
104
|
+
def _train_epoch(self, loader: Iterable[Batch]) -> dict[str, float]:
|
|
105
|
+
self.model.train()
|
|
106
|
+
loss_mean = Mean()
|
|
107
|
+
self._reset_metrics()
|
|
108
|
+
for batch_idx, batch in enumerate(loader):
|
|
109
|
+
inputs, targets = self._split(batch)
|
|
110
|
+
self.optimizer.zero_grad(set_to_none=True)
|
|
111
|
+
outputs = self.model(inputs)
|
|
112
|
+
loss = self.loss_fn(outputs, targets)
|
|
113
|
+
loss.backward()
|
|
114
|
+
if self.grad_clip is not None:
|
|
115
|
+
nn.utils.clip_grad_norm_(self.model.parameters(), self.grad_clip)
|
|
116
|
+
self.optimizer.step()
|
|
117
|
+
|
|
118
|
+
loss_value = float(loss.detach().item())
|
|
119
|
+
loss_mean.update(loss_value, _batch_size(targets))
|
|
120
|
+
self._update_metrics(outputs.detach(), targets)
|
|
121
|
+
self._call("on_batch_end", batch_idx, loss_value)
|
|
122
|
+
return {"loss": loss_mean.compute(), **self._compute_metrics()}
|
|
123
|
+
|
|
124
|
+
# ------------------------------------------------------------- evaluate
|
|
125
|
+
@torch.no_grad()
|
|
126
|
+
def evaluate(self, loader: Iterable[Batch]) -> dict[str, float]:
|
|
127
|
+
"""Compute loss and metrics on ``loader`` without updating weights.
|
|
128
|
+
|
|
129
|
+
Returns:
|
|
130
|
+
``{"loss": ..., <metric name>: ...}`` (no ``val_`` prefix).
|
|
131
|
+
"""
|
|
132
|
+
self.model.eval()
|
|
133
|
+
loss_mean = Mean()
|
|
134
|
+
self._reset_metrics()
|
|
135
|
+
for batch in loader:
|
|
136
|
+
inputs, targets = self._split(batch)
|
|
137
|
+
outputs = self.model(inputs)
|
|
138
|
+
loss = self.loss_fn(outputs, targets)
|
|
139
|
+
loss_mean.update(float(loss.item()), _batch_size(targets))
|
|
140
|
+
self._update_metrics(outputs, targets)
|
|
141
|
+
return {"loss": loss_mean.compute(), **self._compute_metrics()}
|
|
142
|
+
|
|
143
|
+
# -------------------------------------------------------------- predict
|
|
144
|
+
@torch.no_grad()
|
|
145
|
+
def predict(self, loader: Iterable[Batch]) -> torch.Tensor:
|
|
146
|
+
"""Run the model over ``loader`` and concatenate outputs on the CPU.
|
|
147
|
+
|
|
148
|
+
Each batch may be ``(inputs, targets)`` (targets are ignored) or bare
|
|
149
|
+
``inputs``. The model must return a tensor.
|
|
150
|
+
"""
|
|
151
|
+
self.model.eval()
|
|
152
|
+
outputs: list[torch.Tensor] = []
|
|
153
|
+
for batch in loader:
|
|
154
|
+
if isinstance(batch, (list, tuple)) and len(batch) == 2:
|
|
155
|
+
batch = batch[0]
|
|
156
|
+
outputs.append(self.model(to_device(batch, self.device)).cpu())
|
|
157
|
+
if not outputs:
|
|
158
|
+
raise ValueError("predict() received an empty loader")
|
|
159
|
+
return torch.cat(outputs, dim=0)
|
|
160
|
+
|
|
161
|
+
# -------------------------------------------------------------- helpers
|
|
162
|
+
def _split(self, batch: Batch) -> tuple[Any, Any]:
|
|
163
|
+
if not (isinstance(batch, (list, tuple)) and len(batch) == 2):
|
|
164
|
+
raise TypeError(
|
|
165
|
+
"Each batch must be an (inputs, targets) pair; "
|
|
166
|
+
f"got {type(batch).__name__}"
|
|
167
|
+
)
|
|
168
|
+
inputs, targets = batch
|
|
169
|
+
return to_device(inputs, self.device), to_device(targets, self.device)
|
|
170
|
+
|
|
171
|
+
def _reset_metrics(self) -> None:
|
|
172
|
+
for metric in self.metrics.values():
|
|
173
|
+
metric.reset()
|
|
174
|
+
|
|
175
|
+
def _update_metrics(self, outputs: torch.Tensor, targets: torch.Tensor) -> None:
|
|
176
|
+
for metric in self.metrics.values():
|
|
177
|
+
metric.update(outputs, targets)
|
|
178
|
+
|
|
179
|
+
def _compute_metrics(self) -> dict[str, float]:
|
|
180
|
+
return {name: metric.compute() for name, metric in self.metrics.items()}
|
|
181
|
+
|
|
182
|
+
def _call(self, hook: str, *args: Any) -> None:
|
|
183
|
+
for callback in self.callbacks:
|
|
184
|
+
getattr(callback, hook)(self, *args)
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
def _batch_size(targets: Any) -> int:
|
|
188
|
+
"""Best-effort batch size used to weight the per-batch loss."""
|
|
189
|
+
if isinstance(targets, torch.Tensor) and targets.dim() > 0:
|
|
190
|
+
return int(targets.shape[0])
|
|
191
|
+
if isinstance(targets, (list, tuple)) and targets:
|
|
192
|
+
return _batch_size(targets[0])
|
|
193
|
+
if isinstance(targets, Mapping) and targets:
|
|
194
|
+
return _batch_size(next(iter(targets.values())))
|
|
195
|
+
return 1
|
torchlight/utils.py
ADDED
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
"""Small helpers for reproducibility, device selection and moving data around."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import os
|
|
6
|
+
import random
|
|
7
|
+
from collections.abc import Mapping, Sequence
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
import torch
|
|
11
|
+
from torch import nn
|
|
12
|
+
|
|
13
|
+
__all__ = ["seed_everything", "get_device", "to_device", "count_parameters"]
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def seed_everything(seed: int, deterministic: bool = False) -> int:
|
|
17
|
+
"""Seed Python, NumPy (if installed) and PyTorch RNGs.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
seed: The seed value.
|
|
21
|
+
deterministic: If True, also ask PyTorch to use deterministic
|
|
22
|
+
algorithms (may be slower, and raises on ops that have none).
|
|
23
|
+
|
|
24
|
+
Returns:
|
|
25
|
+
The seed that was applied, for convenient logging.
|
|
26
|
+
"""
|
|
27
|
+
random.seed(seed)
|
|
28
|
+
os.environ["PYTHONHASHSEED"] = str(seed)
|
|
29
|
+
try:
|
|
30
|
+
import numpy as np
|
|
31
|
+
|
|
32
|
+
np.random.seed(seed)
|
|
33
|
+
except ImportError: # NumPy is optional
|
|
34
|
+
pass
|
|
35
|
+
torch.manual_seed(seed)
|
|
36
|
+
if torch.cuda.is_available():
|
|
37
|
+
torch.cuda.manual_seed_all(seed)
|
|
38
|
+
if deterministic:
|
|
39
|
+
torch.use_deterministic_algorithms(True)
|
|
40
|
+
torch.backends.cudnn.benchmark = False
|
|
41
|
+
return seed
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def get_device(prefer: str | torch.device | None = None) -> torch.device:
|
|
45
|
+
"""Return a ``torch.device``.
|
|
46
|
+
|
|
47
|
+
Args:
|
|
48
|
+
prefer: An explicit device (e.g. ``"cpu"``, ``"cuda:1"``). If omitted,
|
|
49
|
+
picks CUDA, then Apple MPS, then CPU, whichever is available first.
|
|
50
|
+
"""
|
|
51
|
+
if prefer is not None:
|
|
52
|
+
return torch.device(prefer)
|
|
53
|
+
if torch.cuda.is_available():
|
|
54
|
+
return torch.device("cuda")
|
|
55
|
+
mps = getattr(torch.backends, "mps", None)
|
|
56
|
+
if mps is not None and mps.is_available():
|
|
57
|
+
return torch.device("mps")
|
|
58
|
+
return torch.device("cpu")
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def to_device(data: Any, device: str | torch.device) -> Any:
|
|
62
|
+
"""Recursively move tensors in ``data`` to ``device``.
|
|
63
|
+
|
|
64
|
+
Handles tensors, mappings, lists and tuples (including named tuples).
|
|
65
|
+
Any other object is returned unchanged.
|
|
66
|
+
"""
|
|
67
|
+
if isinstance(data, torch.Tensor):
|
|
68
|
+
return data.to(device)
|
|
69
|
+
if isinstance(data, Mapping):
|
|
70
|
+
return type(data)({k: to_device(v, device) for k, v in data.items()}) # type: ignore[call-arg]
|
|
71
|
+
if isinstance(data, tuple) and hasattr(data, "_fields"): # named tuple
|
|
72
|
+
return type(data)(*(to_device(v, device) for v in data))
|
|
73
|
+
if isinstance(data, Sequence) and not isinstance(data, (str, bytes)):
|
|
74
|
+
return type(data)(to_device(v, device) for v in data) # type: ignore[call-arg]
|
|
75
|
+
return data
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def count_parameters(model: nn.Module, trainable_only: bool = True) -> int:
|
|
79
|
+
"""Count the parameters of ``model``.
|
|
80
|
+
|
|
81
|
+
Args:
|
|
82
|
+
model: The module to inspect.
|
|
83
|
+
trainable_only: Count only parameters with ``requires_grad=True``.
|
|
84
|
+
"""
|
|
85
|
+
return sum(
|
|
86
|
+
p.numel() for p in model.parameters() if p.requires_grad or not trainable_only
|
|
87
|
+
)
|
|
@@ -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,11 @@
|
|
|
1
|
+
torchlight/__init__.py,sha256=D82Nv0gGD_nRdNtVCd8g9KkRXPVFqZqiunGdbuK3gc4,719
|
|
2
|
+
torchlight/callbacks.py,sha256=o9w-ylhbe2L4vab8rW5oGWjvyhGpGw5wzba9WsBUL1A,6533
|
|
3
|
+
torchlight/checkpoint.py,sha256=QaC84FlJ3B-SvzPo-EU1wnqyxlC0DlArsLPfcID_GoE,2094
|
|
4
|
+
torchlight/metrics.py,sha256=uLdWqdZBm7YRsjI8P2xp5uHXsh0gm1Ggv_3dDycNOh8,3591
|
|
5
|
+
torchlight/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
|
|
6
|
+
torchlight/trainer.py,sha256=avnNDr1KWjcuvSqeXbpHOV0NUJMIiudbzaCj7tHwmYg,7731
|
|
7
|
+
torchlight/utils.py,sha256=iZJEpS5lIOj2I9Gw5bIbaX3pVJtvgZLchPiJXI3QBP4,2866
|
|
8
|
+
torchlight-0.1.0.dist-info/METADATA,sha256=19w2UwLqJfEAPYMF9Bev7fKUt0yVT0O2uX1h_TTQ6l4,7859
|
|
9
|
+
torchlight-0.1.0.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
|
|
10
|
+
torchlight-0.1.0.dist-info/licenses/LICENSE,sha256=pAfYREEW9GAy7cnK20OXjDn7ofJahYny9GCuIZQTDAA,1061
|
|
11
|
+
torchlight-0.1.0.dist-info/RECORD,,
|
|
@@ -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.
|