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 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
+ ]
@@ -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}")
@@ -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,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.4
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -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.