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 +48 -0
- eztrain/callbacks.py +125 -0
- eztrain/checkpoint.py +32 -0
- eztrain/loggers.py +118 -0
- eztrain/media.py +25 -0
- eztrain/metrics.py +61 -0
- eztrain/py.typed +0 -0
- eztrain/run.py +58 -0
- eztrain/trainer.py +227 -0
- eztrain-0.1.0.dist-info/METADATA +187 -0
- eztrain-0.1.0.dist-info/RECORD +13 -0
- eztrain-0.1.0.dist-info/WHEEL +4 -0
- eztrain-0.1.0.dist-info/licenses/LICENSE +21 -0
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,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.
|