eztrain 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.
eztrain/__init__.py ADDED
@@ -0,0 +1,48 @@
1
+ """eztrain: a small, framework-agnostic training-loop library.
2
+
3
+ The trainer companion to EzConfy: every public class is instantiable from
4
+ plain keyword arguments, so it can be built straight from YAML.
5
+ """
6
+
7
+ from importlib.metadata import PackageNotFoundError, version
8
+
9
+ from eztrain.callbacks import (
10
+ Callback,
11
+ CheckpointCallback,
12
+ EarlyStopping,
13
+ MonitorCallback,
14
+ )
15
+ from eztrain.checkpoint import Checkpointer
16
+ from eztrain.loggers import Logger, NullLogger, RecordingLogger, WandbLogger
17
+ from eztrain.media import Image, Video
18
+ from eztrain.metrics import Metric, MetricCollection
19
+ from eztrain.run import RunInfo, RunType, generate_run_id, resolve_run, run_id_base
20
+ from eztrain.trainer import EpochTrainer, Trainer
21
+
22
+ try:
23
+ __version__ = version("eztrain")
24
+ except PackageNotFoundError: # pragma: no cover - package not installed
25
+ __version__ = "0.0.0"
26
+
27
+ __all__ = [
28
+ "Callback",
29
+ "CheckpointCallback",
30
+ "Checkpointer",
31
+ "EarlyStopping",
32
+ "EpochTrainer",
33
+ "Image",
34
+ "Logger",
35
+ "Metric",
36
+ "MetricCollection",
37
+ "MonitorCallback",
38
+ "NullLogger",
39
+ "RecordingLogger",
40
+ "RunInfo",
41
+ "RunType",
42
+ "Trainer",
43
+ "Video",
44
+ "WandbLogger",
45
+ "generate_run_id",
46
+ "resolve_run",
47
+ "run_id_base",
48
+ ]
eztrain/callbacks.py ADDED
@@ -0,0 +1,125 @@
1
+ """Callbacks: observe and steer a trainer through named hooks.
2
+
3
+ ``Callback`` is a nominal base class (so config systems like EzConfy can
4
+ type a polymorphic ``list[Callback]``), but dispatch is dynamic by hook
5
+ name — see ``Trainer.call_hook`` — so callbacks may also implement custom
6
+ hooks emitted by project-specific trainers.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ from typing import TYPE_CHECKING, Literal
12
+
13
+ from loguru import logger as log
14
+
15
+ from eztrain.checkpoint import Checkpointer
16
+
17
+ if TYPE_CHECKING:
18
+ from eztrain.trainer import Trainer
19
+
20
+
21
+ class Callback:
22
+ def on_train_start(self, trainer: Trainer) -> None:
23
+ pass
24
+
25
+ def on_iteration_end(self, trainer: Trainer, iteration: int) -> None:
26
+ pass
27
+
28
+ def on_eval_end(self, trainer: Trainer) -> None:
29
+ pass
30
+
31
+ def on_train_end(self, trainer: Trainer) -> None:
32
+ pass
33
+
34
+
35
+ class MonitorCallback(Callback):
36
+ """Tracks the best value of ``trainer.history[monitor]``.
37
+
38
+ Base for anything that reacts to "the metric improved" (early stopping,
39
+ best-model checkpointing). Subclasses call :meth:`improved`.
40
+ """
41
+
42
+ def __init__(self, *, monitor: str, mode: Literal["min", "max"] = "min") -> None:
43
+ self.monitor = monitor
44
+ self.mode = mode
45
+ self.best: float | None = None
46
+
47
+ def improved(self, trainer: Trainer) -> bool:
48
+ value = trainer.history.get(self.monitor)
49
+ if value is None:
50
+ log.warning("'{}' not found in trainer history; skipping.", self.monitor)
51
+ return False
52
+ try:
53
+ value = float(value)
54
+ except (TypeError, ValueError):
55
+ log.warning(
56
+ "'{}' value {!r} is not a number; skipping.", self.monitor, value
57
+ )
58
+ return False
59
+
60
+ if self.best is None or (
61
+ value < self.best if self.mode == "min" else value > self.best
62
+ ):
63
+ self.best = value
64
+ return True
65
+ return False
66
+
67
+
68
+ class EarlyStopping(MonitorCallback):
69
+ """Sets ``trainer.should_stop`` after ``patience`` evaluations without
70
+ improvement of ``monitor``."""
71
+
72
+ def __init__(
73
+ self,
74
+ *,
75
+ monitor: str,
76
+ mode: Literal["min", "max"] = "min",
77
+ patience: int = 10,
78
+ ) -> None:
79
+ super().__init__(monitor=monitor, mode=mode)
80
+ self.patience = patience
81
+ self.counter = 0
82
+
83
+ def on_eval_end(self, trainer: Trainer) -> None:
84
+ if self.improved(trainer):
85
+ self.counter = 0
86
+ return
87
+ self.counter += 1
88
+ log.info(
89
+ "No improvement in '{}' for {}/{} evaluations.",
90
+ self.monitor,
91
+ self.counter,
92
+ self.patience,
93
+ )
94
+ if self.counter >= self.patience:
95
+ log.info("Early stopping triggered.")
96
+ trainer.should_stop = True
97
+
98
+
99
+ class CheckpointCallback(Callback):
100
+ """Checkpoint *schedule*: restore on train start, save every
101
+ ``save_freq`` iterations and once more at the end if needed.
102
+
103
+ The *mechanics* (file formats, best/latest policies, what CONTINUE vs
104
+ FORK restores) belong to the injected
105
+ :class:`~eztrain.checkpoint.Checkpointer`.
106
+ """
107
+
108
+ def __init__(self, *, checkpointer: Checkpointer, save_freq: int = 1) -> None:
109
+ self.checkpointer = checkpointer
110
+ self.save_freq = save_freq
111
+ self._last_saved: int | None = None
112
+
113
+ def on_train_start(self, trainer: Trainer) -> None:
114
+ self.checkpointer.setup(trainer)
115
+
116
+ def on_iteration_end(self, trainer: Trainer, iteration: int) -> None:
117
+ if iteration % self.save_freq == 0:
118
+ self.checkpointer.save(trainer, iteration, trainer.history)
119
+ self._last_saved = iteration
120
+
121
+ def on_train_end(self, trainer: Trainer) -> None:
122
+ if trainer.iteration > 0 and self._last_saved != trainer.iteration:
123
+ self.checkpointer.save(trainer, trainer.iteration, trainer.history)
124
+ self._last_saved = trainer.iteration
125
+ self.checkpointer.close()
eztrain/checkpoint.py ADDED
@@ -0,0 +1,32 @@
1
+ """Checkpointer protocol.
2
+
3
+ The schedule (when to save/restore) lives in
4
+ :class:`~eztrain.callbacks.CheckpointCallback`; the mechanics (what a
5
+ checkpoint physically is) live in framework-specific implementations
6
+ (torch/orbax extras). A trainer advertises what to persist through its
7
+ ``checkpointables`` mapping.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from collections.abc import Mapping
13
+ from typing import TYPE_CHECKING, Any, Protocol, runtime_checkable
14
+
15
+ if TYPE_CHECKING:
16
+ from eztrain.trainer import Trainer
17
+
18
+
19
+ @runtime_checkable
20
+ class Checkpointer(Protocol):
21
+ def setup(self, trainer: Trainer) -> None:
22
+ """Prepare storage for ``trainer.run.run_id`` and, if
23
+ ``trainer.run.restore_dir`` is set, restore according to
24
+ ``trainer.run.run_type`` (CONTINUE: full state, set
25
+ ``trainer.start_iteration``; FORK: weights only)."""
26
+ ...
27
+
28
+ def save(
29
+ self, trainer: Trainer, iteration: int, metrics: Mapping[str, Any]
30
+ ) -> None: ...
31
+
32
+ def close(self) -> None: ...
eztrain/loggers.py ADDED
@@ -0,0 +1,118 @@
1
+ """Logger protocol and built-in implementations.
2
+
3
+ The trainer owns *when* to log and with which ``step``; loggers only own
4
+ *where* the metrics go. ``WandbLogger`` imports wandb lazily so the core
5
+ package works without it installed.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import sys
11
+ from collections.abc import Mapping
12
+ from typing import Any, Protocol, runtime_checkable
13
+
14
+ from eztrain.media import Image, Video
15
+ from eztrain.run import RunInfo
16
+
17
+
18
+ @runtime_checkable
19
+ class Logger(Protocol):
20
+ def start(self, run: RunInfo, config: Mapping[str, Any] | None = None) -> None: ...
21
+
22
+ def log(self, metrics: Mapping[str, Any], step: int | None = None) -> None: ...
23
+
24
+ def finish(self) -> None: ...
25
+
26
+
27
+ class NullLogger:
28
+ """Discards everything. The default when no logger is given."""
29
+
30
+ def start(self, run: RunInfo, config: Mapping[str, Any] | None = None) -> None:
31
+ pass
32
+
33
+ def log(self, metrics: Mapping[str, Any], step: int | None = None) -> None:
34
+ pass
35
+
36
+ def finish(self) -> None:
37
+ pass
38
+
39
+
40
+ class RecordingLogger:
41
+ """Keeps every call in memory. Useful in tests and quick scripts."""
42
+
43
+ def __init__(self) -> None:
44
+ self.run: RunInfo | None = None
45
+ self.config: Mapping[str, Any] | None = None
46
+ self.records: list[tuple[dict[str, Any], int | None]] = []
47
+ self.finished: bool = False
48
+
49
+ def start(self, run: RunInfo, config: Mapping[str, Any] | None = None) -> None:
50
+ self.run = run
51
+ self.config = config
52
+
53
+ def log(self, metrics: Mapping[str, Any], step: int | None = None) -> None:
54
+ self.records.append((dict(metrics), step))
55
+
56
+ def finish(self) -> None:
57
+ self.finished = True
58
+
59
+
60
+ class WandbLogger:
61
+ """Weights & Biases logger. Requires the ``eztrain[wandb]`` extra.
62
+
63
+ Maps :class:`~eztrain.run.RunInfo` onto ``wandb.init``: the run id is the
64
+ wandb id (so CONTINUE runs resume the same wandb run, ``resume="must"``)
65
+ and the base name is the display name.
66
+ """
67
+
68
+ def __init__(
69
+ self,
70
+ *,
71
+ project: str,
72
+ entity: str | None = None,
73
+ group: str | None = None,
74
+ job_type: str | None = None,
75
+ ) -> None:
76
+ self.project = project
77
+ self.entity = entity
78
+ self.group = group
79
+ self.job_type = job_type
80
+
81
+ def start(self, run: RunInfo, config: Mapping[str, Any] | None = None) -> None:
82
+ import wandb
83
+
84
+ wandb.init(
85
+ entity=self.entity,
86
+ project=self.project,
87
+ group=self.group,
88
+ job_type=self.job_type,
89
+ name=run.name,
90
+ id=run.run_id,
91
+ resume=run.resume,
92
+ config=dict(config) if config is not None else None,
93
+ )
94
+
95
+ def log(self, metrics: Mapping[str, Any], step: int | None = None) -> None:
96
+ import wandb
97
+
98
+ wandb.log({k: self._convert(v) for k, v in metrics.items()}, step=step)
99
+
100
+ def finish(self) -> None:
101
+ import wandb
102
+
103
+ wandb.finish()
104
+
105
+ @staticmethod
106
+ def _convert(value: Any) -> Any:
107
+ import wandb
108
+
109
+ if isinstance(value, Image):
110
+ converted = wandb.Image(value.data)
111
+ # close matplotlib figures without importing matplotlib ourselves
112
+ plt = sys.modules.get("matplotlib.pyplot")
113
+ if plt is not None and hasattr(value.data, "savefig"):
114
+ plt.close(value.data)
115
+ return converted
116
+ if isinstance(value, Video):
117
+ return wandb.Video(value.frames, fps=value.fps, format="mp4")
118
+ return value
eztrain/media.py ADDED
@@ -0,0 +1,25 @@
1
+ """Framework-neutral wrappers for rich log payloads.
2
+
3
+ Metrics and callbacks produce these instead of tracker-specific objects
4
+ (``wandb.Image``/``wandb.Video``); each ``Logger`` implementation converts
5
+ them to whatever its backend understands, or drops them.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import Any
11
+
12
+
13
+ class Image:
14
+ """A static image: a matplotlib Figure or an HxWxC array."""
15
+
16
+ def __init__(self, data: Any) -> None:
17
+ self.data = data
18
+
19
+
20
+ class Video:
21
+ """A video as an array of frames (e.g. TxCxHxW), plus playback fps."""
22
+
23
+ def __init__(self, frames: Any, fps: int = 1) -> None:
24
+ self.frames = frames
25
+ self.fps = fps
eztrain/metrics.py ADDED
@@ -0,0 +1,61 @@
1
+ """Metric protocol and composite collection.
2
+
3
+ A metric accumulates state across ``update`` calls and produces named scalars
4
+ in ``compute``. The trainer only ever calls ``reset``/``compute``/``plot``;
5
+ ``update`` is called by *your* ``eval_step`` with whatever signature your
6
+ metrics need, so the protocol does not fix one.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ from collections.abc import Iterator, Mapping, Sequence
12
+ from typing import Any, Protocol, runtime_checkable
13
+
14
+ from eztrain.media import Image, Video
15
+
16
+
17
+ @runtime_checkable
18
+ class Metric(Protocol):
19
+ def reset(self) -> None: ...
20
+
21
+ def update(self, *args: Any, **kwargs: Any) -> None: ...
22
+
23
+ def compute(self) -> Mapping[str, float]: ...
24
+
25
+ # Optionally, a metric may also define:
26
+ # def plot(self) -> Mapping[str, Image | Video]: ...
27
+
28
+
29
+ class MetricCollection:
30
+ """Fans out to a list of metrics and merges their results."""
31
+
32
+ def __init__(self, metrics: Sequence[Metric] | None = None) -> None:
33
+ self.metrics: list[Metric] = list(metrics or [])
34
+
35
+ def __iter__(self) -> Iterator[Metric]:
36
+ return iter(self.metrics)
37
+
38
+ def __len__(self) -> int:
39
+ return len(self.metrics)
40
+
41
+ def reset(self) -> None:
42
+ for metric in self.metrics:
43
+ metric.reset()
44
+
45
+ def update(self, *args: Any, **kwargs: Any) -> None:
46
+ for metric in self.metrics:
47
+ metric.update(*args, **kwargs)
48
+
49
+ def compute(self) -> dict[str, float]:
50
+ results: dict[str, float] = {}
51
+ for metric in self.metrics:
52
+ results.update(metric.compute())
53
+ return results
54
+
55
+ def plot(self) -> dict[str, Image | Video]:
56
+ visuals: dict[str, Image | Video] = {}
57
+ for metric in self.metrics:
58
+ plot = getattr(metric, "plot", None)
59
+ if plot is not None:
60
+ visuals.update(plot())
61
+ return visuals
eztrain/py.typed ADDED
File without changes
eztrain/run.py ADDED
@@ -0,0 +1,58 @@
1
+ """Run lifecycle: identity, resume and fork semantics for a training run.
2
+
3
+ A run id is ``"<base_name>_<timestamp>"`` (e.g. ``"coinrun-exp-1_20260530_051406"``)
4
+ and is meant to be shared between the checkpoint directory and the experiment
5
+ tracker id (e.g. ``wandb.init(id=...)``), so that resuming a run resumes both.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import re
11
+ from dataclasses import dataclass
12
+ from datetime import datetime
13
+ from enum import Enum, auto
14
+ from typing import Literal, TypeAlias
15
+
16
+ ResumePolicy: TypeAlias = Literal["never", "must"]
17
+
18
+
19
+ class RunType(Enum):
20
+ FRESH = auto() # no checkpoint, train from scratch
21
+ CONTINUE = auto() # same run: same id/folder, resume tracker + optimizer state
22
+ FORK = auto() # new run: new id/folder, load weights only, fresh tracker run
23
+
24
+
25
+ _TIMESTAMP_FMT = "%Y%m%d_%H%M%S"
26
+ _TIMESTAMP_RE = re.compile(r"_\d{8}_\d{6}$")
27
+
28
+
29
+ def generate_run_id(base_name: str) -> str:
30
+ timestamp = datetime.now().strftime(_TIMESTAMP_FMT)
31
+ return f"{base_name}_{timestamp}"
32
+
33
+
34
+ def run_id_base(run_id: str) -> str:
35
+ return _TIMESTAMP_RE.sub("", run_id)
36
+
37
+
38
+ @dataclass(frozen=True)
39
+ class RunInfo:
40
+ name: str
41
+ run_id: str
42
+ run_type: RunType
43
+ restore_dir: str | None # checkpoint folder to load weights from (None = scratch)
44
+ resume: ResumePolicy # tracker resume semantics ("must" only when continuing)
45
+
46
+
47
+ def resolve_run(resume_from: str | None, name: str) -> RunInfo:
48
+ """3 modes from (resume_from, name):
49
+
50
+ - resume_from unset -> FRESH: new id, train from scratch.
51
+ - resume_from + same name -> CONTINUE: same id/folder, resume tracker.
52
+ - resume_from + new name -> FORK: new id/folder, load weights, fresh tracker.
53
+ """
54
+ if resume_from is None:
55
+ return RunInfo(name, generate_run_id(name), RunType.FRESH, None, "never")
56
+ if name == run_id_base(resume_from):
57
+ return RunInfo(name, resume_from, RunType.CONTINUE, resume_from, "must")
58
+ return RunInfo(name, generate_run_id(name), RunType.FORK, resume_from, "never")
eztrain/trainer.py ADDED
@@ -0,0 +1,227 @@
1
+ """The training loop.
2
+
3
+ :class:`Trainer` is a small template-method class: it owns the *skeleton*
4
+ every training run shares (run identity, hooks, periodic evaluation, history,
5
+ logging, graceful interruption) and delegates the *content* of one iteration
6
+ to a subclass. An iteration can be anything that returns a mapping of
7
+ metrics: an epoch over a dataloader (see :class:`EpochTrainer`), an RL
8
+ collect->update cycle, a world-model phase — the core never touches tensors.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from collections import defaultdict
14
+ from collections.abc import Iterable, Mapping
15
+ from typing import Any
16
+
17
+ from loguru import logger as log
18
+ from tqdm import tqdm
19
+
20
+ from eztrain.callbacks import Callback
21
+ from eztrain.loggers import Logger, NullLogger
22
+ from eztrain.metrics import MetricCollection
23
+ from eztrain.run import RunInfo, resolve_run
24
+
25
+
26
+ class Trainer:
27
+ """Base trainer: subclass and implement :meth:`train_iteration`.
28
+
29
+ Stable surface exposed to callbacks (safe to rely on):
30
+
31
+ - ``run``: :class:`~eztrain.run.RunInfo` (run_id, run_type, restore_dir)
32
+ - ``iteration``: current iteration (0 before the loop starts)
33
+ - ``start_iteration``: writable; a restoring callback sets it in
34
+ ``on_train_start`` and the loop resumes from ``start_iteration + 1``
35
+ - ``history``: dict with the latest value of every logged metric
36
+ - ``should_stop``: writable; checked after every iteration
37
+ - ``checkpointables``: what a checkpointer should persist (override it)
38
+ - ``logger`` and ``call_hook``
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ *,
44
+ max_iterations: int,
45
+ run_name: str = "run",
46
+ resume_from: str | None = None,
47
+ eval_freq: int = 0,
48
+ callbacks: list[Callback] | None = None,
49
+ logger: Logger | None = None,
50
+ config: Mapping[str, Any] | None = None,
51
+ unit: str = "epoch",
52
+ ) -> None:
53
+ self.run: RunInfo = resolve_run(resume_from, run_name)
54
+ self.max_iterations = max_iterations
55
+ self.eval_freq = eval_freq
56
+ self.callbacks: list[Callback] = callbacks or []
57
+ self.logger: Logger = logger or NullLogger()
58
+ self.config = config
59
+ self.unit = unit
60
+
61
+ self.start_iteration: int = 0
62
+ self.iteration: int = 0
63
+ self.history: dict[str, Any] = {}
64
+ self.should_stop: bool = False
65
+
66
+ if self.run.restore_dir is not None:
67
+ log.info(
68
+ "{} run '{}' (restoring from {})",
69
+ self.run.run_type.name,
70
+ self.run.run_id,
71
+ self.run.restore_dir,
72
+ )
73
+
74
+ # --- override points ---------------------------------------------------
75
+
76
+ def train_iteration(self, iteration: int) -> Mapping[str, Any]:
77
+ """One unit of training work (an epoch, an RL update, ...).
78
+
79
+ Returns the metrics to log for this iteration.
80
+ """
81
+ raise NotImplementedError
82
+
83
+ def evaluate(self) -> Mapping[str, Any]:
84
+ """Periodic evaluation, called every ``eval_freq`` iterations."""
85
+ return {}
86
+
87
+ def log_step(self) -> int:
88
+ """The ``step`` passed to the logger. Override e.g. with a global
89
+ environment-step counter in RL trainers."""
90
+ return self.iteration
91
+
92
+ @property
93
+ def checkpointables(self) -> Mapping[str, Any]:
94
+ """What a :class:`~eztrain.checkpoint.Checkpointer` should persist
95
+ (e.g. model/optimizer objects, keyed by name). Override it."""
96
+ return {}
97
+
98
+ # --- template loop -----------------------------------------------------
99
+
100
+ def fit(self) -> None:
101
+ self.logger.start(self.run, self.config)
102
+ try:
103
+ self.call_hook("on_train_start")
104
+ for iteration in tqdm(
105
+ range(self.start_iteration + 1, self.max_iterations + 1),
106
+ desc=self.unit,
107
+ colour="green",
108
+ ):
109
+ self.iteration = iteration
110
+
111
+ logs = dict(self.train_iteration(iteration))
112
+ self.history.update(logs)
113
+
114
+ if self.eval_freq and iteration % self.eval_freq == 0:
115
+ eval_logs = dict(self.evaluate())
116
+ logs.update(eval_logs)
117
+ self.history.update(eval_logs)
118
+ self.call_hook("on_eval_end")
119
+
120
+ self.logger.log(logs, step=self.log_step())
121
+ self.call_hook("on_iteration_end", iteration=iteration)
122
+
123
+ if self.should_stop:
124
+ log.info("Stop requested, ending training early.")
125
+ break
126
+ except KeyboardInterrupt:
127
+ log.warning("Training interrupted by user.")
128
+ finally:
129
+ self.call_hook("on_train_end")
130
+ self.logger.finish()
131
+
132
+ def call_hook(self, hook: str, **kwargs: Any) -> None:
133
+ """Invoke ``hook`` on every callback that defines it.
134
+
135
+ Dispatch is by name, so trainers may introduce custom hooks (e.g.
136
+ ``on_rollout_end``) that only some callbacks implement.
137
+ """
138
+ for callback in self.callbacks:
139
+ fn = getattr(callback, hook, None)
140
+ if fn is not None:
141
+ fn(self, **kwargs)
142
+
143
+
144
+ class EpochTrainer(Trainer):
145
+ """Supervised-style trainer: one iteration = one pass over ``train_loader``.
146
+
147
+ Still framework-agnostic: batches are opaque and only handled by your
148
+ ``train_step``/``eval_step``. Per-step metrics are averaged over the
149
+ epoch and prefixed ``train/`` and ``val/``; anything your metrics
150
+ ``compute``/``plot`` is added on evaluation. Your ``eval_step`` is
151
+ responsible for calling ``self.metrics.update(...)`` with whatever
152
+ arguments your metrics expect. For per-batch logging, call
153
+ ``self.logger.log({...})`` inside ``train_step`` (no ``step``).
154
+ """
155
+
156
+ def __init__(
157
+ self,
158
+ *,
159
+ train_loader: Iterable[Any],
160
+ val_loader: Iterable[Any] | None = None,
161
+ test_loader: Iterable[Any] | None = None,
162
+ metrics: MetricCollection | None = None,
163
+ **kwargs: Any,
164
+ ) -> None:
165
+ super().__init__(**kwargs)
166
+ self.train_loader = train_loader
167
+ self.val_loader = val_loader
168
+ self.test_loader = test_loader
169
+ self.metrics = metrics or MetricCollection()
170
+
171
+ # --- override points ---------------------------------------------------
172
+
173
+ def train_step(self, batch: Any) -> Mapping[str, float]:
174
+ """One optimization step. Returns per-step metrics (e.g. losses)."""
175
+ raise NotImplementedError
176
+
177
+ def eval_step(self, batch: Any) -> Mapping[str, float]:
178
+ """One evaluation step. Update ``self.metrics`` here and return
179
+ per-step metrics (e.g. losses) to be averaged."""
180
+ raise NotImplementedError
181
+
182
+ # --- Trainer implementation --------------------------------------------
183
+
184
+ def train_iteration(self, iteration: int) -> Mapping[str, Any]:
185
+ step_logs: defaultdict[str, float] = defaultdict(float)
186
+ num_steps = 0
187
+ for batch in tqdm(
188
+ self.train_loader,
189
+ desc=f"{self.unit} {iteration}",
190
+ leave=False,
191
+ colour="blue",
192
+ ):
193
+ for key, value in self.train_step(batch).items():
194
+ step_logs[key] += float(value)
195
+ num_steps += 1
196
+
197
+ if num_steps == 0:
198
+ return {}
199
+ return {f"train/{k}": v / num_steps for k, v in step_logs.items()}
200
+
201
+ def evaluate(self) -> Mapping[str, Any]:
202
+ if self.val_loader is None:
203
+ return {}
204
+ return self.evaluate_split(self.val_loader, "val")
205
+
206
+ def evaluate_split(self, loader: Iterable[Any], prefix: str) -> dict[str, Any]:
207
+ """Run ``eval_step`` over ``loader``; average step metrics and merge
208
+ ``metrics.compute()``/``metrics.plot()``, all under ``prefix/``."""
209
+ self.metrics.reset()
210
+
211
+ step_logs: defaultdict[str, float] = defaultdict(float)
212
+ num_steps = 0
213
+ for batch in tqdm(
214
+ loader, desc=f"evaluating {prefix}", leave=False, colour="red"
215
+ ):
216
+ for key, value in self.eval_step(batch).items():
217
+ step_logs[key] += float(value)
218
+ num_steps += 1
219
+
220
+ results: dict[str, Any] = {}
221
+ if num_steps > 0:
222
+ results.update(
223
+ {f"{prefix}/{k}": v / num_steps for k, v in step_logs.items()}
224
+ )
225
+ results.update({f"{prefix}/{k}": v for k, v in self.metrics.compute().items()})
226
+ results.update({f"{prefix}/{k}": v for k, v in self.metrics.plot().items()})
227
+ return results
@@ -0,0 +1,187 @@
1
+ Metadata-Version: 2.5
2
+ Name: eztrain
3
+ Version: 0.1.0
4
+ Summary: Small, framework-agnostic training-loop library. The trainer companion to EzConfy.
5
+ Project-URL: Homepage, https://github.com/alessioarcara/EzTrain
6
+ Project-URL: Repository, https://github.com/alessioarcara/EzTrain
7
+ Project-URL: Issues, https://github.com/alessioarcara/EzTrain/issues
8
+ Author: Alessio Arcara
9
+ License-Expression: MIT
10
+ License-File: LICENSE
11
+ Keywords: deep-learning,ezconfy,machine-learning,trainer,training
12
+ Classifier: Development Status :: 3 - Alpha
13
+ Classifier: Intended Audience :: Science/Research
14
+ Classifier: Operating System :: OS Independent
15
+ Classifier: Programming Language :: Python :: 3
16
+ Classifier: Programming Language :: Python :: 3.10
17
+ Classifier: Programming Language :: Python :: 3.11
18
+ Classifier: Programming Language :: Python :: 3.12
19
+ Classifier: Programming Language :: Python :: 3.13
20
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
21
+ Classifier: Typing :: Typed
22
+ Requires-Python: >=3.10
23
+ Requires-Dist: loguru>=0.7
24
+ Requires-Dist: tqdm>=4.66
25
+ Provides-Extra: wandb
26
+ Requires-Dist: wandb>=0.17; extra == 'wandb'
27
+ Description-Content-Type: text/markdown
28
+
29
+ # EzTrain
30
+
31
+ A small, framework-agnostic training-loop library. The trainer companion to
32
+ [EzConfy](https://github.com/alessioarcara/EzConfy).
33
+
34
+ Every ML project rewrites the same trainer: a loop over epochs or updates,
35
+ periodic evaluation, callbacks, early stopping, checkpoint scheduling, metric
36
+ logging, graceful `Ctrl+C`. **EzTrain** extracts exactly that skeleton and
37
+ nothing else:
38
+
39
+ - **Core never imports torch or jax.** It abstracts the *iteration*, not the
40
+ tensor: an iteration is anything that returns a `Mapping[str, Any]` of
41
+ metrics — an epoch over a dataloader, an RL collect→update cycle, a
42
+ world-model phase.
43
+ - **Run lifecycle built in.** A run id (`<name>_<timestamp>`) is shared
44
+ between checkpoints and the experiment tracker, with three resume modes:
45
+ **FRESH** (from scratch), **CONTINUE** (same run: weights + optimizer +
46
+ tracker), **FORK** (new run seeded with old weights).
47
+ - **Composition over inheritance** where it matters: `Logger`, `Metric` and
48
+ `Checkpointer` are Protocols you inject; `Callback` hooks are dispatched
49
+ dynamically by name so trainers can invent their own hooks.
50
+ - **No hidden defaults.** No optimizer/scheduler construction, no concrete
51
+ metrics, no config system (that's EzConfy's job), no distributed magic.
52
+
53
+ Core dependencies: `tqdm`, `loguru`. That's it.
54
+
55
+ ## Install
56
+
57
+ ```bash
58
+ uv add eztrain # core
59
+ uv add "eztrain[wandb]" # + Weights & Biases logger
60
+ ```
61
+
62
+ ## Supervised (epoch-based)
63
+
64
+ ```python
65
+ from eztrain import EpochTrainer, EarlyStopping, MetricCollection, WandbLogger
66
+
67
+ class MyTrainer(EpochTrainer):
68
+ def __init__(self, *, model, optimizer, **kwargs):
69
+ super().__init__(**kwargs)
70
+ self.model, self.optimizer = model, optimizer
71
+
72
+ def train_step(self, batch):
73
+ loss = ... # your forward/backward/step
74
+ return {"loss": loss.item()} # averaged over the epoch -> "train/loss"
75
+
76
+ def eval_step(self, batch):
77
+ preds, loss = ...
78
+ self.metrics.update(preds, batch.y) # your metrics, your signature
79
+ return {"loss": loss.item()} # -> "val/loss"
80
+
81
+ trainer = MyTrainer(
82
+ model=model,
83
+ optimizer=optimizer,
84
+ train_loader=train_loader,
85
+ val_loader=val_loader,
86
+ metrics=MetricCollection([MyAccuracy()]),
87
+ max_iterations=100, # epochs
88
+ eval_freq=1,
89
+ callbacks=[EarlyStopping(monitor="val/loss", patience=10)],
90
+ logger=WandbLogger(project="my-project"),
91
+ run_name="baseline",
92
+ )
93
+ trainer.fit()
94
+ ```
95
+
96
+ ## RL (update-based)
97
+
98
+ Subclass `Trainer` directly — one iteration is one update:
99
+
100
+ ```python
101
+ from eztrain import Trainer
102
+
103
+ class PPOTrainer(Trainer):
104
+ def __init__(self, *, env, agent, num_steps, **kwargs):
105
+ super().__init__(unit="update", **kwargs)
106
+ self.env, self.agent, self.num_steps = env, agent, num_steps
107
+ self.obs = env.reset()
108
+
109
+ def train_iteration(self, update):
110
+ segment, self.obs = collect_rollouts(self.env, self.agent, self.num_steps, self.obs)
111
+ advantages, returns = compute_gae(segment)
112
+ return self.agent.learn_from(segment, advantages, returns)
113
+
114
+ def evaluate(self):
115
+ return {"eval/reward": evaluate(self.agent)}
116
+
117
+ def log_step(self): # log by env steps, not updates
118
+ return self.iteration * self.num_steps * self.env.num_envs
119
+ ```
120
+
121
+ ## Resume and fork
122
+
123
+ ```python
124
+ Trainer(run_name="exp-1") # FRESH
125
+ Trainer(run_name="exp-1", resume_from="exp-1_20260530_051406") # CONTINUE
126
+ Trainer(run_name="exp-2", resume_from="exp-1_20260530_051406") # FORK
127
+ ```
128
+
129
+ `trainer.run` carries `run_id`, `run_type` and `restore_dir`. A
130
+ `CheckpointCallback` restores in `on_train_start` (setting
131
+ `trainer.start_iteration`) and saves on schedule; the checkpoint *mechanics*
132
+ live in a `Checkpointer` implementation you inject (torch/orbax
133
+ implementations ship as extras — coming next). `WandbLogger` reuses the run
134
+ id, so a CONTINUE run resumes the same wandb run.
135
+
136
+ ## Callbacks
137
+
138
+ ```python
139
+ from eztrain import Callback
140
+
141
+ class MyCallback(Callback):
142
+ def on_train_start(self, trainer): ...
143
+ def on_iteration_end(self, trainer, iteration): ...
144
+ def on_eval_end(self, trainer): ... # trainer.history has fresh metrics
145
+ def on_train_end(self, trainer): ... # always runs, even on Ctrl+C
146
+ ```
147
+
148
+ The stable surface callbacks can rely on: `trainer.run`, `trainer.iteration`,
149
+ `trainer.start_iteration` (writable), `trainer.history`,
150
+ `trainer.should_stop` (writable), `trainer.checkpointables`,
151
+ `trainer.logger`, `trainer.call_hook`. Custom hooks compose freely:
152
+ `self.call_hook("on_rollout_end", num_steps=n)` inside your trainer reaches
153
+ any callback that defines it.
154
+
155
+ ## With EzConfy
156
+
157
+ Every public class takes plain keyword arguments, so it can be instantiated
158
+ straight from YAML:
159
+
160
+ ```yaml
161
+ # schema.yaml
162
+ types:
163
+ Callback: eztrain.callbacks:Callback
164
+ Logger: eztrain.loggers:Logger
165
+ schema:
166
+ trainer:
167
+ callbacks: list[Callback]
168
+ logger: Logger
169
+ ```
170
+
171
+ ```yaml
172
+ # config.yaml
173
+ trainer:
174
+ logger:
175
+ _target_type_: eztrain.loggers:WandbLogger
176
+ _init_args_: { project: my-project, entity: me }
177
+ callbacks:
178
+ - _target_type_: eztrain.callbacks:EarlyStopping
179
+ _init_args_: { monitor: val/loss, mode: min, patience: 10 }
180
+ ```
181
+
182
+ ## What's deliberately *not* here
183
+
184
+ Concrete train steps, models, losses, optimizers/schedulers (inject your
185
+ own — a library default becomes a cage), metrics implementations, config
186
+ loading, distributed training. Keep those in your project; eztrain only
187
+ owns the loop around them.
@@ -0,0 +1,13 @@
1
+ eztrain/__init__.py,sha256=Oq84rWsb1UBWWEclh3Xr9qXotodYXpSDYcxkbEkUDcQ,1255
2
+ eztrain/callbacks.py,sha256=aCc4_jxy804LgZke15FK_79QYLaW-1rZJhb1F_UQLmU,4002
3
+ eztrain/checkpoint.py,sha256=M6QvOk-zWRDLvyBBaJyINuJUARhg_PSHLU-2D36WjTk,1015
4
+ eztrain/loggers.py,sha256=xizU3MdZuo7yfJz7JibMfgNEaDNTPTPBAUeFvEHgMcw,3523
5
+ eztrain/media.py,sha256=LZCUxmTBdR5QKGRequtu0KmakggJ1iM6yvu3-8ETIYo,673
6
+ eztrain/metrics.py,sha256=ssnUAOj5M1JZpYcHJiEfCthE01Ge92UUS1cesXG0Iy4,1865
7
+ eztrain/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
8
+ eztrain/run.py,sha256=IxDW1QMK137G0wvOWIldEshBZVhibI36bmwXbQ6JhW4,2016
9
+ eztrain/trainer.py,sha256=h6QFZ5igK27zxytbGVLQxcdWkj18Cd49fxi5mz6VeR8,8530
10
+ eztrain-0.1.0.dist-info/METADATA,sha256=kpMQr-gDdrvgSgjiieaaBODFxL08Pmk96gAIMmeBQzk,6715
11
+ eztrain-0.1.0.dist-info/WHEEL,sha256=zOwg4jB6zX2kU910N-cMawjivD6tO8NEWvE12je1bVk,87
12
+ eztrain-0.1.0.dist-info/licenses/LICENSE,sha256=lIIz0RQXSrdzWcGFtNs-wewg9LJ5D-O-CvElp2uaeY0,1071
13
+ eztrain-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.32.0
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Alessio Arcara
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.