marn 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.
- marn/__init__.py +154 -0
- marn/callbacks/__init__.py +7 -0
- marn/callbacks/base.py +51 -0
- marn/callbacks/early_stopping.py +60 -0
- marn/callbacks/logger.py +47 -0
- marn/checkpoint/__init__.py +12 -0
- marn/checkpoint/load.py +174 -0
- marn/checkpoint/save.py +122 -0
- marn/checkpoint/schema.py +54 -0
- marn/config/__init__.py +61 -0
- marn/config/loss.py +50 -0
- marn/config/model.py +67 -0
- marn/config/trainer.py +39 -0
- marn/distributed/__init__.py +18 -0
- marn/distributed/ddp.py +180 -0
- marn/distributed/memory.py +157 -0
- marn/generators/__init__.py +22 -0
- marn/generators/base.py +47 -0
- marn/generators/finetuning.py +129 -0
- marn/generators/grouped.py +138 -0
- marn/generators/layer_generator.py +31 -0
- marn/generators/layerwise.py +36 -0
- marn/generators/lazy.py +105 -0
- marn/generators/lrd.py +157 -0
- marn/generators/single_vector.py +91 -0
- marn/losses/__init__.py +23 -0
- marn/losses/alignment.py +87 -0
- marn/losses/base.py +23 -0
- marn/losses/context.py +40 -0
- marn/losses/mapping_loss.py +226 -0
- marn/losses/outputs.py +24 -0
- marn/losses/smoothness.py +147 -0
- marn/losses/stability.py +79 -0
- marn/losses/task.py +79 -0
- marn/mappers/__init__.py +6 -0
- marn/mappers/base.py +28 -0
- marn/mappers/mlp_mapper.py +123 -0
- marn/models/__init__.py +7 -0
- marn/models/forward_result.py +29 -0
- marn/models/mapping_model.py +142 -0
- marn/models/target_model.py +80 -0
- marn/modulation/__init__.py +8 -0
- marn/modulation/additive.py +25 -0
- marn/modulation/affine.py +29 -0
- marn/modulation/base.py +16 -0
- marn/modulation/low_rank.py +34 -0
- marn/py.typed +1 -0
- marn/registry.py +111 -0
- marn/runtime/__init__.py +13 -0
- marn/runtime/functional.py +41 -0
- marn/runtime/lazy.py +58 -0
- marn/runtime/parameter_spec.py +160 -0
- marn/runtime/parameter_tree.py +54 -0
- marn/strategies/__init__.py +17 -0
- marn/strategies/base.py +14 -0
- marn/strategies/finetuning.py +58 -0
- marn/strategies/grouped.py +18 -0
- marn/strategies/layerwise.py +15 -0
- marn/strategies/lrd.py +41 -0
- marn/strategies/slvt.py +24 -0
- marn/trainers/__init__.py +17 -0
- marn/trainers/batch_adapter.py +76 -0
- marn/trainers/lr_finder.py +374 -0
- marn/trainers/trainer.py +464 -0
- marn-0.1.0.dist-info/METADATA +125 -0
- marn-0.1.0.dist-info/RECORD +68 -0
- marn-0.1.0.dist-info/WHEEL +4 -0
- marn-0.1.0.dist-info/licenses/LICENSE +161 -0
marn/__init__.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
"""Train PyTorch models through low-dimensional parameter mappings."""
|
|
2
|
+
|
|
3
|
+
from marn.generators import (
|
|
4
|
+
GroupedGenerator,
|
|
5
|
+
LayerGenerator,
|
|
6
|
+
LazyLayerwiseGenerator,
|
|
7
|
+
LayerwiseGenerator,
|
|
8
|
+
ParameterGenerator,
|
|
9
|
+
SingleVectorGenerator,
|
|
10
|
+
FineTuningGenerator,
|
|
11
|
+
LRDGenerator,
|
|
12
|
+
)
|
|
13
|
+
from marn.distributed import (
|
|
14
|
+
benchmark_strategies,
|
|
15
|
+
cleanup_ddp,
|
|
16
|
+
is_ddp_available,
|
|
17
|
+
profile_peak_memory,
|
|
18
|
+
setup_ddp,
|
|
19
|
+
wrap_ddp,
|
|
20
|
+
)
|
|
21
|
+
from marn.models import (
|
|
22
|
+
ForwardResult,
|
|
23
|
+
MappingModel,
|
|
24
|
+
TargetModel,
|
|
25
|
+
UnsupportedTargetModelError,
|
|
26
|
+
)
|
|
27
|
+
from marn.losses import (
|
|
28
|
+
AlignmentLoss,
|
|
29
|
+
BaseLoss,
|
|
30
|
+
ClassificationLoss,
|
|
31
|
+
LossOutput,
|
|
32
|
+
MappingLoss,
|
|
33
|
+
RegressionLoss,
|
|
34
|
+
SmoothnessLoss,
|
|
35
|
+
StabilityLoss,
|
|
36
|
+
TaskLoss,
|
|
37
|
+
TrainingContext,
|
|
38
|
+
)
|
|
39
|
+
from marn.mappers import BaseMapper, MLPMapper, ResidualMLPMapper
|
|
40
|
+
from marn.modulation import (
|
|
41
|
+
AdditiveModulation,
|
|
42
|
+
AffineModulation,
|
|
43
|
+
BaseModulation,
|
|
44
|
+
LowRankModulation,
|
|
45
|
+
)
|
|
46
|
+
from marn.runtime import ParameterEntry, ParameterSpec, ParameterTree
|
|
47
|
+
from marn.strategies import (
|
|
48
|
+
GroupedStrategy,
|
|
49
|
+
LayerwiseStrategy,
|
|
50
|
+
SLVTStrategy,
|
|
51
|
+
FineTuningStrategy,
|
|
52
|
+
LRDStrategy,
|
|
53
|
+
)
|
|
54
|
+
from marn.config import (
|
|
55
|
+
GeneratorConfig,
|
|
56
|
+
LossConfig,
|
|
57
|
+
MapperConfig,
|
|
58
|
+
MappingConfig,
|
|
59
|
+
TaskLossConfig,
|
|
60
|
+
TrainerConfig,
|
|
61
|
+
load_config,
|
|
62
|
+
)
|
|
63
|
+
from marn.registry import (
|
|
64
|
+
GENERATOR_REGISTRY,
|
|
65
|
+
LOSS_REGISTRY,
|
|
66
|
+
MAPPER_REGISTRY,
|
|
67
|
+
MODULATION_REGISTRY,
|
|
68
|
+
Registry,
|
|
69
|
+
)
|
|
70
|
+
from marn.callbacks import Callback, EarlyStopping, MetricLogger
|
|
71
|
+
from marn.trainers import (
|
|
72
|
+
BatchAdapter,
|
|
73
|
+
MappingBatchAdapter,
|
|
74
|
+
MappingTrainer,
|
|
75
|
+
TupleBatchAdapter,
|
|
76
|
+
LRFinderResult,
|
|
77
|
+
)
|
|
78
|
+
from marn.checkpoint import (
|
|
79
|
+
CheckpointCompatibilityError,
|
|
80
|
+
CheckpointSchema,
|
|
81
|
+
load_checkpoint,
|
|
82
|
+
save_checkpoint,
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
__all__ = [
|
|
86
|
+
"AdditiveModulation",
|
|
87
|
+
"AffineModulation",
|
|
88
|
+
"AlignmentLoss",
|
|
89
|
+
"BaseLoss",
|
|
90
|
+
"BaseMapper",
|
|
91
|
+
"BaseModulation",
|
|
92
|
+
"BatchAdapter",
|
|
93
|
+
"Callback",
|
|
94
|
+
"CheckpointCompatibilityError",
|
|
95
|
+
"CheckpointSchema",
|
|
96
|
+
"ClassificationLoss",
|
|
97
|
+
"EarlyStopping",
|
|
98
|
+
"ForwardResult",
|
|
99
|
+
"FineTuningGenerator",
|
|
100
|
+
"GENERATOR_REGISTRY",
|
|
101
|
+
"GeneratorConfig",
|
|
102
|
+
"GroupedGenerator",
|
|
103
|
+
"GroupedStrategy",
|
|
104
|
+
"LOSS_REGISTRY",
|
|
105
|
+
"LayerGenerator",
|
|
106
|
+
"LazyLayerwiseGenerator",
|
|
107
|
+
"LayerwiseGenerator",
|
|
108
|
+
"LayerwiseStrategy",
|
|
109
|
+
"LRDGenerator",
|
|
110
|
+
"LossConfig",
|
|
111
|
+
"LossOutput",
|
|
112
|
+
"LowRankModulation",
|
|
113
|
+
"MAPPER_REGISTRY",
|
|
114
|
+
"MODULATION_REGISTRY",
|
|
115
|
+
"MapperConfig",
|
|
116
|
+
"MappingBatchAdapter",
|
|
117
|
+
"MappingConfig",
|
|
118
|
+
"MappingLoss",
|
|
119
|
+
"MappingModel",
|
|
120
|
+
"MappingTrainer",
|
|
121
|
+
"MLPMapper",
|
|
122
|
+
"MetricLogger",
|
|
123
|
+
"ParameterEntry",
|
|
124
|
+
"ParameterGenerator",
|
|
125
|
+
"ParameterSpec",
|
|
126
|
+
"ParameterTree",
|
|
127
|
+
"RegressionLoss",
|
|
128
|
+
"Registry",
|
|
129
|
+
"ResidualMLPMapper",
|
|
130
|
+
"SLVTStrategy",
|
|
131
|
+
"SingleVectorGenerator",
|
|
132
|
+
"FineTuningStrategy",
|
|
133
|
+
"LRDStrategy",
|
|
134
|
+
"LRFinderResult",
|
|
135
|
+
"SmoothnessLoss",
|
|
136
|
+
"StabilityLoss",
|
|
137
|
+
"TargetModel",
|
|
138
|
+
"TaskLoss",
|
|
139
|
+
"TaskLossConfig",
|
|
140
|
+
"TrainerConfig",
|
|
141
|
+
"TrainingContext",
|
|
142
|
+
"TupleBatchAdapter",
|
|
143
|
+
"UnsupportedTargetModelError",
|
|
144
|
+
"benchmark_strategies",
|
|
145
|
+
"cleanup_ddp",
|
|
146
|
+
"is_ddp_available",
|
|
147
|
+
"load_checkpoint",
|
|
148
|
+
"load_config",
|
|
149
|
+
"profile_peak_memory",
|
|
150
|
+
"save_checkpoint",
|
|
151
|
+
"setup_ddp",
|
|
152
|
+
"wrap_ddp",
|
|
153
|
+
]
|
|
154
|
+
__version__ = "0.1.0"
|
marn/callbacks/base.py
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
1
|
+
"""Observer callback base class for training lifecycle hooks."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from typing import TYPE_CHECKING, Any
|
|
6
|
+
|
|
7
|
+
if TYPE_CHECKING:
|
|
8
|
+
from marn.losses.outputs import LossOutput
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class Callback:
|
|
12
|
+
"""Base class for training lifecycle callbacks.
|
|
13
|
+
|
|
14
|
+
All hooks are no-ops by default — subclass and override the ones you need.
|
|
15
|
+
Callbacks receive the trainer instance so they can inspect or modify state
|
|
16
|
+
(e.g. set ``trainer.should_stop = True`` for early termination).
|
|
17
|
+
|
|
18
|
+
Hook execution order within each event:
|
|
19
|
+
|
|
20
|
+
1. ``on_fit_start``
|
|
21
|
+
2. For each epoch:
|
|
22
|
+
a. ``on_epoch_start``
|
|
23
|
+
b. For each batch: ``on_batch_start`` → ``on_batch_end``
|
|
24
|
+
c. ``on_validation_start`` → ``on_validation_end``
|
|
25
|
+
d. ``on_epoch_end``
|
|
26
|
+
3. ``on_fit_end``
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
def on_fit_start(self, trainer: Any) -> None:
|
|
30
|
+
"""Called once at the beginning of :meth:`fit`."""
|
|
31
|
+
|
|
32
|
+
def on_fit_end(self, trainer: Any) -> None:
|
|
33
|
+
"""Called once at the end of :meth:`fit`."""
|
|
34
|
+
|
|
35
|
+
def on_epoch_start(self, trainer: Any, epoch: int) -> None:
|
|
36
|
+
"""Called at the beginning of each epoch."""
|
|
37
|
+
|
|
38
|
+
def on_epoch_end(self, trainer: Any, epoch: int, metrics: dict[str, float]) -> None:
|
|
39
|
+
"""Called at the end of each epoch with aggregated metrics."""
|
|
40
|
+
|
|
41
|
+
def on_batch_start(self, trainer: Any, batch_idx: int) -> None:
|
|
42
|
+
"""Called before processing each batch."""
|
|
43
|
+
|
|
44
|
+
def on_batch_end(self, trainer: Any, batch_idx: int, loss_output: LossOutput) -> None:
|
|
45
|
+
"""Called after processing each batch with the loss output."""
|
|
46
|
+
|
|
47
|
+
def on_validation_start(self, trainer: Any) -> None:
|
|
48
|
+
"""Called before the validation loop."""
|
|
49
|
+
|
|
50
|
+
def on_validation_end(self, trainer: Any, metrics: dict[str, float]) -> None:
|
|
51
|
+
"""Called after the validation loop with validation metrics."""
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
"""Early stopping callback."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
from marn.callbacks.base import Callback
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class EarlyStopping(Callback):
|
|
12
|
+
"""Stop training when a monitored metric stops improving.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
monitor: Name of the metric to monitor (e.g. ``"val_loss"``).
|
|
16
|
+
patience: Number of epochs with no improvement before stopping.
|
|
17
|
+
min_delta: Minimum change to qualify as an improvement.
|
|
18
|
+
mode: ``"min"`` to minimize the metric, ``"max"`` to maximize it.
|
|
19
|
+
|
|
20
|
+
Example::
|
|
21
|
+
|
|
22
|
+
trainer = MappingTrainer(
|
|
23
|
+
...,
|
|
24
|
+
callbacks=[EarlyStopping(monitor="val_loss", patience=5)],
|
|
25
|
+
)
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
monitor: str = "val_loss",
|
|
31
|
+
patience: int = 10,
|
|
32
|
+
min_delta: float = 0.0,
|
|
33
|
+
mode: str = "min",
|
|
34
|
+
) -> None:
|
|
35
|
+
if mode not in ("min", "max"):
|
|
36
|
+
raise ValueError(f"mode must be 'min' or 'max', got {mode!r}")
|
|
37
|
+
self.monitor = monitor
|
|
38
|
+
self.patience = patience
|
|
39
|
+
self.min_delta = min_delta
|
|
40
|
+
self.mode = mode
|
|
41
|
+
self._best: float = math.inf if mode == "min" else -math.inf
|
|
42
|
+
self._wait: int = 0
|
|
43
|
+
|
|
44
|
+
def _is_improvement(self, current: float) -> bool:
|
|
45
|
+
if self.mode == "min":
|
|
46
|
+
return current < self._best - self.min_delta
|
|
47
|
+
return current > self._best + self.min_delta
|
|
48
|
+
|
|
49
|
+
def on_epoch_end(self, trainer: Any, epoch: int, metrics: dict[str, float]) -> None:
|
|
50
|
+
current = metrics.get(self.monitor)
|
|
51
|
+
if current is None:
|
|
52
|
+
return
|
|
53
|
+
|
|
54
|
+
if self._is_improvement(current):
|
|
55
|
+
self._best = current
|
|
56
|
+
self._wait = 0
|
|
57
|
+
else:
|
|
58
|
+
self._wait += 1
|
|
59
|
+
if self._wait >= self.patience:
|
|
60
|
+
trainer.should_stop = True
|
marn/callbacks/logger.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
"""Structured metric logging callback."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
from typing import TYPE_CHECKING, Any
|
|
7
|
+
|
|
8
|
+
from marn.callbacks.base import Callback
|
|
9
|
+
|
|
10
|
+
if TYPE_CHECKING:
|
|
11
|
+
from marn.losses.outputs import LossOutput
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger("marn.trainer")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class MetricLogger(Callback):
|
|
17
|
+
"""Log training and validation metrics using the standard ``logging`` module.
|
|
18
|
+
|
|
19
|
+
Args:
|
|
20
|
+
log_every_n_batches: Log batch metrics every N batches.
|
|
21
|
+
|
|
22
|
+
Example::
|
|
23
|
+
|
|
24
|
+
trainer = MappingTrainer(
|
|
25
|
+
...,
|
|
26
|
+
callbacks=[MetricLogger(log_every_n_batches=50)],
|
|
27
|
+
)
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
def __init__(self, log_every_n_batches: int = 10) -> None:
|
|
31
|
+
self.log_every_n_batches = log_every_n_batches
|
|
32
|
+
|
|
33
|
+
def on_epoch_start(self, trainer: Any, epoch: int) -> None:
|
|
34
|
+
logger.info("Epoch %d/%d started", epoch + 1, trainer.config.max_epochs)
|
|
35
|
+
|
|
36
|
+
def on_batch_end(self, trainer: Any, batch_idx: int, loss_output: LossOutput) -> None:
|
|
37
|
+
if (batch_idx + 1) % self.log_every_n_batches == 0:
|
|
38
|
+
metrics_str = ", ".join(f"{k}={v:.4f}" for k, v in loss_output.metrics.items())
|
|
39
|
+
logger.info(" batch %d: %s", batch_idx + 1, metrics_str)
|
|
40
|
+
|
|
41
|
+
def on_epoch_end(self, trainer: Any, epoch: int, metrics: dict[str, float]) -> None:
|
|
42
|
+
metrics_str = ", ".join(f"{k}={v:.4f}" for k, v in metrics.items())
|
|
43
|
+
logger.info("Epoch %d/%d: %s", epoch + 1, trainer.config.max_epochs, metrics_str)
|
|
44
|
+
|
|
45
|
+
def on_validation_end(self, trainer: Any, metrics: dict[str, float]) -> None:
|
|
46
|
+
metrics_str = ", ".join(f"{k}={v:.4f}" for k, v in metrics.items())
|
|
47
|
+
logger.info(" validation: %s", metrics_str)
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""Versioned checkpoint save and load for mapping network training state."""
|
|
2
|
+
|
|
3
|
+
from marn.checkpoint.save import save_checkpoint
|
|
4
|
+
from marn.checkpoint.load import load_checkpoint, CheckpointCompatibilityError
|
|
5
|
+
from marn.checkpoint.schema import CheckpointSchema
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"CheckpointCompatibilityError",
|
|
9
|
+
"CheckpointSchema",
|
|
10
|
+
"load_checkpoint",
|
|
11
|
+
"save_checkpoint",
|
|
12
|
+
]
|
marn/checkpoint/load.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
"""Load a mapping network checkpoint from disk and restore model/trainer state."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import TYPE_CHECKING, Any
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
|
|
11
|
+
from marn.checkpoint.schema import CURRENT_SCHEMA_VERSION, CheckpointSchema
|
|
12
|
+
|
|
13
|
+
if TYPE_CHECKING:
|
|
14
|
+
from marn.models.mapping_model import MappingModel
|
|
15
|
+
from marn.trainers.trainer import MappingTrainer
|
|
16
|
+
|
|
17
|
+
logger = logging.getLogger("marn.checkpoint")
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class CheckpointCompatibilityError(ValueError):
|
|
21
|
+
"""Raised when a checkpoint's architecture is incompatible with the model.
|
|
22
|
+
|
|
23
|
+
Possible causes:
|
|
24
|
+
|
|
25
|
+
* Schema version from a future release is not understood by this version.
|
|
26
|
+
* Target ``ParameterSpec`` names or shapes differ (different architecture).
|
|
27
|
+
* Latent or buffer tensors do not match the current generator's keys.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _load_payload(path: Path, map_location: Any) -> dict[str, Any]:
|
|
32
|
+
raw = torch.load(path, map_location=map_location, weights_only=False)
|
|
33
|
+
if not isinstance(raw, dict):
|
|
34
|
+
raise CheckpointCompatibilityError(
|
|
35
|
+
f"Checkpoint at {path} is not a valid marn checkpoint "
|
|
36
|
+
f"(expected a dict, got {type(raw).__name__})."
|
|
37
|
+
)
|
|
38
|
+
return raw
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _parse_schema(payload: dict[str, Any], path: Path) -> CheckpointSchema:
|
|
42
|
+
schema_version = payload.get("schema_version")
|
|
43
|
+
if schema_version is None:
|
|
44
|
+
raise CheckpointCompatibilityError(
|
|
45
|
+
f"Checkpoint at {path} is missing 'schema_version'. "
|
|
46
|
+
"It may have been saved by an incompatible tool."
|
|
47
|
+
)
|
|
48
|
+
if schema_version > CURRENT_SCHEMA_VERSION:
|
|
49
|
+
raise CheckpointCompatibilityError(
|
|
50
|
+
f"Checkpoint schema version {schema_version} is newer than the "
|
|
51
|
+
f"current package supports ({CURRENT_SCHEMA_VERSION}). "
|
|
52
|
+
"Upgrade marn to load this checkpoint."
|
|
53
|
+
)
|
|
54
|
+
return CheckpointSchema(
|
|
55
|
+
schema_version=schema_version,
|
|
56
|
+
package_version=payload.get("package_version", "unknown"),
|
|
57
|
+
latent_state=payload["latent_state"],
|
|
58
|
+
mapper_buffers=payload["mapper_buffers"],
|
|
59
|
+
parameter_spec_names=payload["parameter_spec_names"],
|
|
60
|
+
config=payload.get("config"),
|
|
61
|
+
trainer_state=payload.get("trainer_state"),
|
|
62
|
+
metadata=payload.get("metadata", {}),
|
|
63
|
+
)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _validate_spec(schema: CheckpointSchema, model: MappingModel, path: Path) -> None:
|
|
67
|
+
"""Raise ``CheckpointCompatibilityError`` if the target spec has changed."""
|
|
68
|
+
spec = model.target.parameter_spec
|
|
69
|
+
current: list[tuple[str, list[int]]] = [(entry.name, list(entry.shape)) for entry in spec]
|
|
70
|
+
saved = [(name, shape) for name, shape in schema.parameter_spec_names]
|
|
71
|
+
|
|
72
|
+
if len(current) != len(saved):
|
|
73
|
+
raise CheckpointCompatibilityError(
|
|
74
|
+
f"Checkpoint at {path} has {len(saved)} target parameters but "
|
|
75
|
+
f"the model has {len(current)}. The target architecture differs."
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
mismatches: list[str] = []
|
|
79
|
+
for (cur_name, cur_shape), (sav_name, sav_shape) in zip(current, saved):
|
|
80
|
+
if cur_name != sav_name:
|
|
81
|
+
mismatches.append(f" name mismatch: checkpoint={sav_name!r} model={cur_name!r}")
|
|
82
|
+
elif cur_shape != sav_shape:
|
|
83
|
+
mismatches.append(
|
|
84
|
+
f" shape mismatch for {cur_name!r}: checkpoint={sav_shape} model={cur_shape}"
|
|
85
|
+
)
|
|
86
|
+
|
|
87
|
+
if mismatches:
|
|
88
|
+
detail = "\n".join(mismatches)
|
|
89
|
+
raise CheckpointCompatibilityError(
|
|
90
|
+
f"Checkpoint at {path} is incompatible with this model:\n{detail}"
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def load_checkpoint(
|
|
95
|
+
path: str | Path,
|
|
96
|
+
model: MappingModel,
|
|
97
|
+
*,
|
|
98
|
+
trainer: MappingTrainer | None = None,
|
|
99
|
+
map_location: Any = None,
|
|
100
|
+
strict_latents: bool = True,
|
|
101
|
+
) -> CheckpointSchema:
|
|
102
|
+
"""Load a checkpoint and restore model (and optionally trainer) state.
|
|
103
|
+
|
|
104
|
+
The target's ``ParameterSpec`` is validated against the checkpoint before
|
|
105
|
+
any state is applied. If the architecture differs a
|
|
106
|
+
``CheckpointCompatibilityError`` is raised and the model is left unchanged.
|
|
107
|
+
|
|
108
|
+
Args:
|
|
109
|
+
path: Path to the ``.pt`` checkpoint file.
|
|
110
|
+
model: A ``MappingModel`` whose generator state will be restored.
|
|
111
|
+
trainer: Optional ``MappingTrainer`` whose optimizer, scheduler, epoch,
|
|
112
|
+
and GradScaler state will be restored when present in the checkpoint.
|
|
113
|
+
map_location: Passed to ``torch.load`` for device remapping.
|
|
114
|
+
strict_latents: When *True*, ``load_state_dict`` is called in strict
|
|
115
|
+
mode — every saved latent and buffer key must match the generator.
|
|
116
|
+
Set to *False* to allow partial loading (e.g. subsets of layers).
|
|
117
|
+
|
|
118
|
+
Returns:
|
|
119
|
+
The parsed ``CheckpointSchema`` so callers can inspect ``metadata``,
|
|
120
|
+
``config``, or ``trainer_state`` independently.
|
|
121
|
+
|
|
122
|
+
Raises:
|
|
123
|
+
FileNotFoundError: If the checkpoint file does not exist.
|
|
124
|
+
CheckpointCompatibilityError: If schema version or architecture is
|
|
125
|
+
incompatible.
|
|
126
|
+
|
|
127
|
+
Example::
|
|
128
|
+
|
|
129
|
+
schema = load_checkpoint("checkpoints/epoch10.pt", model=model, trainer=trainer)
|
|
130
|
+
print(f"Resumed from epoch {schema.trainer_state['epoch']}")
|
|
131
|
+
"""
|
|
132
|
+
path = Path(path)
|
|
133
|
+
if not path.exists():
|
|
134
|
+
raise FileNotFoundError(f"Checkpoint not found: {path}")
|
|
135
|
+
|
|
136
|
+
payload = _load_payload(path, map_location)
|
|
137
|
+
schema = _parse_schema(payload, path)
|
|
138
|
+
|
|
139
|
+
# Validate architecture before touching any model state
|
|
140
|
+
_validate_spec(schema, model, path)
|
|
141
|
+
|
|
142
|
+
# Restore generator state (latents + buffers together)
|
|
143
|
+
combined_state: dict[str, Any] = {**schema.latent_state, **schema.mapper_buffers}
|
|
144
|
+
model.generator.load_state_dict(combined_state, strict=strict_latents)
|
|
145
|
+
|
|
146
|
+
logger.info(
|
|
147
|
+
"Checkpoint loaded from %s (schema_version=%d package=%s)",
|
|
148
|
+
path,
|
|
149
|
+
schema.schema_version,
|
|
150
|
+
schema.package_version,
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
# Restore trainer state if both are available
|
|
154
|
+
if trainer is not None and schema.trainer_state is not None:
|
|
155
|
+
ts = schema.trainer_state
|
|
156
|
+
trainer.current_epoch = ts.get("epoch", 0)
|
|
157
|
+
trainer.should_stop = ts.get("should_stop", False)
|
|
158
|
+
|
|
159
|
+
if "optimizer" in ts:
|
|
160
|
+
trainer.optimizer.load_state_dict(ts["optimizer"])
|
|
161
|
+
|
|
162
|
+
if "scheduler" in ts and trainer.scheduler is not None:
|
|
163
|
+
trainer.scheduler.load_state_dict(ts["scheduler"])
|
|
164
|
+
|
|
165
|
+
if "scaler" in ts and trainer.scaler is not None:
|
|
166
|
+
trainer.scaler.load_state_dict(ts["scaler"])
|
|
167
|
+
|
|
168
|
+
logger.info(
|
|
169
|
+
"Trainer state restored: epoch=%d should_stop=%s",
|
|
170
|
+
trainer.current_epoch,
|
|
171
|
+
trainer.should_stop,
|
|
172
|
+
)
|
|
173
|
+
|
|
174
|
+
return schema
|
marn/checkpoint/save.py
ADDED
|
@@ -0,0 +1,122 @@
|
|
|
1
|
+
"""Save a mapping network checkpoint to disk."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import TYPE_CHECKING, Any
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
|
|
11
|
+
from marn.checkpoint.schema import CURRENT_SCHEMA_VERSION, CheckpointSchema
|
|
12
|
+
|
|
13
|
+
if TYPE_CHECKING:
|
|
14
|
+
from marn.config.model import MappingConfig
|
|
15
|
+
from marn.models.mapping_model import MappingModel
|
|
16
|
+
from marn.trainers.trainer import MappingTrainer
|
|
17
|
+
|
|
18
|
+
logger = logging.getLogger("marn.checkpoint")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def save_checkpoint(
|
|
22
|
+
path: str | Path,
|
|
23
|
+
model: MappingModel,
|
|
24
|
+
*,
|
|
25
|
+
config: MappingConfig | None = None,
|
|
26
|
+
trainer: MappingTrainer | None = None,
|
|
27
|
+
save_trainer_state: bool = True,
|
|
28
|
+
metadata: dict[str, Any] | None = None,
|
|
29
|
+
) -> None:
|
|
30
|
+
"""Save a compact, versioned checkpoint for the mapping network.
|
|
31
|
+
|
|
32
|
+
Only trainable latent vectors and required mapper buffers are written.
|
|
33
|
+
Generated target weights are excluded — they are ephemeral and can be
|
|
34
|
+
fully reconstructed at inference time.
|
|
35
|
+
|
|
36
|
+
Args:
|
|
37
|
+
path: Destination file path (e.g. ``"ckpt/epoch10.pt"``).
|
|
38
|
+
Parent directories are created automatically.
|
|
39
|
+
model: A ``MappingModel`` instance (may be on any device).
|
|
40
|
+
config: Optional ``MappingConfig`` saved for reference. Not used
|
|
41
|
+
during load, but useful for reproducibility auditing.
|
|
42
|
+
trainer: Optional ``MappingTrainer`` whose optimizer, scheduler, and
|
|
43
|
+
epoch state are serialised when ``save_trainer_state`` is ``True``.
|
|
44
|
+
save_trainer_state: When *True* and a *trainer* is supplied, include
|
|
45
|
+
optimizer state dict, scheduler state dict (if any), current epoch,
|
|
46
|
+
and ``should_stop`` flag so training can be resumed exactly.
|
|
47
|
+
metadata: Free-form dict of user annotations (dataset name, notes …).
|
|
48
|
+
|
|
49
|
+
Example::
|
|
50
|
+
|
|
51
|
+
save_checkpoint(
|
|
52
|
+
"checkpoints/epoch10.pt",
|
|
53
|
+
model=model,
|
|
54
|
+
config=mapping_config,
|
|
55
|
+
trainer=trainer,
|
|
56
|
+
)
|
|
57
|
+
"""
|
|
58
|
+
import marn # local import to avoid circular at module level
|
|
59
|
+
|
|
60
|
+
path = Path(path)
|
|
61
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
62
|
+
|
|
63
|
+
# ── Latent state ──────────────────────────────────────────────────────
|
|
64
|
+
# Only save trainable parameters (latent vectors); mapper buffers are
|
|
65
|
+
# captured separately so the two sets are clearly delineated.
|
|
66
|
+
latent_names = {name for name, _ in model.generator.named_latent_vectors()}
|
|
67
|
+
full_state = model.generator.state_dict()
|
|
68
|
+
|
|
69
|
+
latent_state: dict[str, Any] = {k: v for k, v in full_state.items() if k in latent_names}
|
|
70
|
+
mapper_buffers: dict[str, Any] = {k: v for k, v in full_state.items() if k not in latent_names}
|
|
71
|
+
|
|
72
|
+
# ── ParameterSpec fingerprint ─────────────────────────────────────────
|
|
73
|
+
spec = model.target.parameter_spec
|
|
74
|
+
spec_names: list[tuple[str, list[int]]] = [(entry.name, list(entry.shape)) for entry in spec]
|
|
75
|
+
|
|
76
|
+
# ── Config ───────────────────────────────────────────────────────────
|
|
77
|
+
config_dict: dict[str, Any] | None = None
|
|
78
|
+
if config is not None:
|
|
79
|
+
config_dict = config.model_dump()
|
|
80
|
+
|
|
81
|
+
# ── Trainer state ────────────────────────────────────────────────────
|
|
82
|
+
trainer_state: dict[str, Any] | None = None
|
|
83
|
+
if trainer is not None and save_trainer_state:
|
|
84
|
+
trainer_state = {
|
|
85
|
+
"epoch": trainer.current_epoch,
|
|
86
|
+
"should_stop": trainer.should_stop,
|
|
87
|
+
"optimizer": trainer.optimizer.state_dict(),
|
|
88
|
+
}
|
|
89
|
+
if trainer.scheduler is not None:
|
|
90
|
+
trainer_state["scheduler"] = trainer.scheduler.state_dict()
|
|
91
|
+
if trainer.scaler is not None:
|
|
92
|
+
trainer_state["scaler"] = trainer.scaler.state_dict()
|
|
93
|
+
|
|
94
|
+
schema = CheckpointSchema(
|
|
95
|
+
schema_version=CURRENT_SCHEMA_VERSION,
|
|
96
|
+
package_version=marn.__version__,
|
|
97
|
+
latent_state=latent_state,
|
|
98
|
+
mapper_buffers=mapper_buffers,
|
|
99
|
+
parameter_spec_names=spec_names,
|
|
100
|
+
config=config_dict,
|
|
101
|
+
trainer_state=trainer_state,
|
|
102
|
+
metadata=metadata or {},
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
payload = {
|
|
106
|
+
"schema_version": schema.schema_version,
|
|
107
|
+
"package_version": schema.package_version,
|
|
108
|
+
"latent_state": schema.latent_state,
|
|
109
|
+
"mapper_buffers": schema.mapper_buffers,
|
|
110
|
+
"parameter_spec_names": schema.parameter_spec_names,
|
|
111
|
+
"config": schema.config,
|
|
112
|
+
"trainer_state": schema.trainer_state,
|
|
113
|
+
"metadata": schema.metadata,
|
|
114
|
+
}
|
|
115
|
+
|
|
116
|
+
torch.save(payload, path)
|
|
117
|
+
logger.info(
|
|
118
|
+
"Checkpoint saved to %s (latents=%d buffers=%d)",
|
|
119
|
+
path,
|
|
120
|
+
len(latent_state),
|
|
121
|
+
len(mapper_buffers),
|
|
122
|
+
)
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""Versioned checkpoint schema for mapping network training state.
|
|
2
|
+
|
|
3
|
+
The checkpoint stores only what is needed to reconstruct and continue training:
|
|
4
|
+
|
|
5
|
+
- ``schema_version``: integer bumped on breaking format changes.
|
|
6
|
+
- ``package_version``: the marn package version string at save time.
|
|
7
|
+
- ``latent_state``: ``state_dict`` subset containing only the trainable latent
|
|
8
|
+
tensors (keys produced by ``model.generator.named_latent_vectors``).
|
|
9
|
+
- ``mapper_buffers``: buffer state required for deterministic reconstruction
|
|
10
|
+
(e.g. fixed orthogonal projection matrices).
|
|
11
|
+
- ``parameter_spec_names``: ordered parameter names and shapes from the target's
|
|
12
|
+
``ParameterSpec``. Used to validate that the same architecture is loaded.
|
|
13
|
+
- ``config``: optional serialised ``MappingConfig`` dict saved for reference.
|
|
14
|
+
- ``trainer_state``: optional dict containing optimizer/scheduler/epoch state
|
|
15
|
+
so training can be resumed at the same point.
|
|
16
|
+
- ``metadata``: free-form dict for user annotations (dataset, notes, …).
|
|
17
|
+
|
|
18
|
+
Generated target weights are **not** included. Checkpoints are compact because
|
|
19
|
+
only latent vectors and small fixed buffers are saved.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
from __future__ import annotations
|
|
23
|
+
|
|
24
|
+
from dataclasses import dataclass, field
|
|
25
|
+
from typing import Any
|
|
26
|
+
|
|
27
|
+
CURRENT_SCHEMA_VERSION: int = 1
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass
|
|
31
|
+
class CheckpointSchema:
|
|
32
|
+
"""All fields written to / read from a checkpoint file.
|
|
33
|
+
|
|
34
|
+
Attributes:
|
|
35
|
+
schema_version: Integer incremented on breaking format changes.
|
|
36
|
+
package_version: ``marn.__version__`` at save time.
|
|
37
|
+
latent_state: Ordered ``state_dict`` of trainable latent tensors only.
|
|
38
|
+
mapper_buffers: Buffer ``state_dict`` required for deterministic
|
|
39
|
+
reconstruction (projection matrices, running statistics …).
|
|
40
|
+
parameter_spec_names: Ordered list of ``(name, shape)`` tuples from
|
|
41
|
+
the target ``ParameterSpec``. Validated on load.
|
|
42
|
+
config: Optional serialised ``MappingConfig`` dict (may be ``None``).
|
|
43
|
+
trainer_state: Optional ``dict`` with optimizer/scheduler/epoch entries.
|
|
44
|
+
metadata: User-supplied annotations, e.g. dataset or description.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
schema_version: int
|
|
48
|
+
package_version: str
|
|
49
|
+
latent_state: dict[str, Any]
|
|
50
|
+
mapper_buffers: dict[str, Any]
|
|
51
|
+
parameter_spec_names: list[tuple[str, list[int]]]
|
|
52
|
+
config: dict[str, Any] | None = None
|
|
53
|
+
trainer_state: dict[str, Any] | None = None
|
|
54
|
+
metadata: dict[str, Any] = field(default_factory=dict)
|