royalelearn 0.5.4__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.
- royalelearn/__init__.py +60 -0
- royalelearn/__main__.py +15 -0
- royalelearn/api/__init__.py +121 -0
- royalelearn/api/advantage.py +61 -0
- royalelearn/api/buffer.py +179 -0
- royalelearn/api/checkpoint.py +132 -0
- royalelearn/api/ladder.py +179 -0
- royalelearn/api/metrics.py +101 -0
- royalelearn/api/policy.py +189 -0
- royalelearn/api/rollout.py +404 -0
- royalelearn/api/schedule.py +59 -0
- royalelearn/api/update.py +202 -0
- royalelearn/checkpoint.py +547 -0
- royalelearn/cli.py +753 -0
- royalelearn/config.py +1368 -0
- royalelearn/coordinator.py +2908 -0
- royalelearn/determinism.py +146 -0
- royalelearn/errors.py +106 -0
- royalelearn/extensions.py +591 -0
- royalelearn/identity.py +798 -0
- royalelearn/ladder/__init__.py +56 -0
- royalelearn/ladder/actors.py +189 -0
- royalelearn/ladder/evaluate.py +404 -0
- royalelearn/ladder/eviction.py +84 -0
- royalelearn/ladder/farm.py +409 -0
- royalelearn/ladder/gate.py +375 -0
- royalelearn/ladder/matchmaker.py +498 -0
- royalelearn/ladder/pool.py +416 -0
- royalelearn/ladder/rating.py +445 -0
- royalelearn/ladder/results.py +345 -0
- royalelearn/ladder/seat_decks.py +66 -0
- royalelearn/ladder/snapshots.py +267 -0
- royalelearn/learn/__init__.py +97 -0
- royalelearn/learn/actor_critic.py +328 -0
- royalelearn/learn/buffer.py +891 -0
- royalelearn/learn/decode.py +125 -0
- royalelearn/learn/distribution.py +277 -0
- royalelearn/learn/freeze.py +92 -0
- royalelearn/learn/gae.py +192 -0
- royalelearn/learn/inference.py +622 -0
- royalelearn/learn/nets.py +690 -0
- royalelearn/learn/ppo.py +1816 -0
- royalelearn/learn/returns.py +184 -0
- royalelearn/learn/rows.py +63 -0
- royalelearn/learn/schedules.py +299 -0
- royalelearn/learner.py +600 -0
- royalelearn/metrics/__init__.py +34 -0
- royalelearn/metrics/alarms.py +627 -0
- royalelearn/metrics/behaviour.py +231 -0
- royalelearn/metrics/bundle.py +201 -0
- royalelearn/metrics/records.py +647 -0
- royalelearn/metrics/schema.py +1166 -0
- royalelearn/metrics/sinks.py +373 -0
- royalelearn/metrics/viser_sink.py +432 -0
- royalelearn/metrics/wandb_sink.py +141 -0
- royalelearn/obs_layout.py +136 -0
- royalelearn/rewards.py +466 -0
- royalelearn/rollout/__init__.py +37 -0
- royalelearn/rollout/codec.py +661 -0
- royalelearn/rollout/envspec.py +527 -0
- royalelearn/rollout/farm.py +560 -0
- royalelearn/rollout/inline.py +1626 -0
- royalelearn/rollout/layout.py +630 -0
- royalelearn/rollout/plan.py +266 -0
- royalelearn/rollout/preflight.py +787 -0
- royalelearn/rollout/scripted.py +151 -0
- royalelearn/rollout/worker.py +334 -0
- royalelearn/seeding.py +152 -0
- royalelearn/testing.py +686 -0
- royalelearn/version.py +42 -0
- royalelearn-0.5.4.dist-info/METADATA +20 -0
- royalelearn-0.5.4.dist-info/RECORD +76 -0
- royalelearn-0.5.4.dist-info/WHEEL +5 -0
- royalelearn-0.5.4.dist-info/entry_points.txt +2 -0
- royalelearn-0.5.4.dist-info/licenses/LICENSE +21 -0
- royalelearn-0.5.4.dist-info/top_level.txt +1 -0
royalelearn/__init__.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
"""royalelearn -- the self-play training harness for RoyaleGym environments.
|
|
2
|
+
|
|
3
|
+
Layers (each user-facing behaviour is an ABC in ``api/`` with a shipped default):
|
|
4
|
+
|
|
5
|
+
RolloutSource api/rollout.py ProcessRolloutSource, InlineRolloutSource
|
|
6
|
+
ActorCritic api/policy.py SeparateActorCritic, SharedTrunkActorCritic
|
|
7
|
+
ObsCodec api/buffer.py SpatialObsCodec
|
|
8
|
+
ExperienceBuffer api/buffer.py RectBuffer
|
|
9
|
+
Matchmaker/Rater api/ladder.py MixMatchmaker, BradleyTerryDavidsonRater
|
|
10
|
+
MetricsSink api/metrics.py JsonlSink, ConsoleSink, WandbSink
|
|
11
|
+
CheckpointStore api/checkpoint.py DirCheckpointStore
|
|
12
|
+
|
|
13
|
+
``LearningCoordinator`` (coordinator.py) is the only place the phases of an iteration are
|
|
14
|
+
ordered. ``docs/harness-spec.md`` specifies all of it; ``docs/design.md`` says why.
|
|
15
|
+
|
|
16
|
+
Importing this package pulls in nothing that needs torch: the ABCs use numpy and msgspec,
|
|
17
|
+
and the concrete learner lives behind the lazy attributes below. That is what lets
|
|
18
|
+
``royalelearn config``, ``royalelearn identity`` and ``royalelearn --help`` work in an
|
|
19
|
+
environment where torch is not installed. Asking for a name that does need torch raises an
|
|
20
|
+
ImportError naming the package that is missing.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
from __future__ import annotations
|
|
24
|
+
|
|
25
|
+
import importlib
|
|
26
|
+
from typing import Any
|
|
27
|
+
|
|
28
|
+
from .version import __version__, git_describe
|
|
29
|
+
|
|
30
|
+
# Public name -> (module, attribute). A None attribute exports the module itself.
|
|
31
|
+
# Resolution is deferred so that the import cost and the torch dependency of a name are
|
|
32
|
+
# paid only by the caller that asks for it.
|
|
33
|
+
_EXPORTS: dict[str, tuple[str, str | None]] = {
|
|
34
|
+
"Learner": (".learner", "Learner"),
|
|
35
|
+
"LearningCoordinator": (".coordinator", "LearningCoordinator"),
|
|
36
|
+
"RunConfig": (".config", "RunConfig"),
|
|
37
|
+
"load_config": (".config", "load_config"),
|
|
38
|
+
"api": (".api", None),
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
__all__ = ["__version__", "git_describe", *sorted(_EXPORTS)]
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def __getattr__(name: str) -> Any:
|
|
45
|
+
"""Resolve a public name on first use.
|
|
46
|
+
|
|
47
|
+
An unknown name is an AttributeError. A known name whose module cannot be imported
|
|
48
|
+
raises that import's own ImportError unchanged, because its message already names what
|
|
49
|
+
is missing and rewording it would hide which package to install.
|
|
50
|
+
"""
|
|
51
|
+
try:
|
|
52
|
+
module_name, attribute = _EXPORTS[name]
|
|
53
|
+
except KeyError:
|
|
54
|
+
raise AttributeError(f"module 'royalelearn' has no attribute {name!r}") from None
|
|
55
|
+
module = importlib.import_module(module_name, __name__)
|
|
56
|
+
return module if attribute is None else getattr(module, attribute)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def __dir__() -> list[str]:
|
|
60
|
+
return sorted(__all__)
|
royalelearn/__main__.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
"""``python -m royalelearn``.
|
|
2
|
+
|
|
3
|
+
The module body is one call. Everything the entry point has to do before torch is imported --
|
|
4
|
+
the deterministic cuBLAS workspace and the BLAS thread counts -- ``cli`` does as it is imported,
|
|
5
|
+
which is why this file imports it and nothing else.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import sys
|
|
11
|
+
|
|
12
|
+
from .cli import main
|
|
13
|
+
|
|
14
|
+
if __name__ == "__main__":
|
|
15
|
+
sys.exit(main())
|
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
"""Every ABC and struct the harness is written against.
|
|
2
|
+
|
|
3
|
+
This subpackage imports numpy and msgspec and nothing else that matters: torch appears in type
|
|
4
|
+
annotations only, behind ``TYPE_CHECKING``. That is what lets a rollout worker, the CLI's
|
|
5
|
+
config and identity commands, and ``import royalelearn`` itself run in an environment with no
|
|
6
|
+
torch installed -- and what keeps a worker's resident memory three hundred megabytes smaller
|
|
7
|
+
than the parent's.
|
|
8
|
+
|
|
9
|
+
The concrete implementations live outside ``api/`` and may import whatever they need:
|
|
10
|
+
|
|
11
|
+
RolloutSource rollout/farm.py, rollout/inline.py
|
|
12
|
+
ActorCritic learn/actor_critic.py
|
|
13
|
+
ObsCodec rollout/codec.py
|
|
14
|
+
ExperienceBuffer learn/buffer.py
|
|
15
|
+
AdvantageEstimator learn/gae.py
|
|
16
|
+
Update learn/ppo.py
|
|
17
|
+
Schedule learn/schedules.py
|
|
18
|
+
Matchmaker, Rater ladder/matchmaker.py, ladder/rating.py
|
|
19
|
+
MetricsSink, Alarm metrics/sinks.py, metrics/alarms.py
|
|
20
|
+
CheckpointStore checkpoint.py
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
from __future__ import annotations
|
|
24
|
+
|
|
25
|
+
from .advantage import AdvantageEstimator, AdvantageStats
|
|
26
|
+
from .buffer import CodecTable, ExperienceBuffer, ObsCodec
|
|
27
|
+
from .checkpoint import Checkpointable, CheckpointStore, Manifest, RngState
|
|
28
|
+
from .ladder import (
|
|
29
|
+
ConditionResult,
|
|
30
|
+
EvictionPolicy,
|
|
31
|
+
GateDecision,
|
|
32
|
+
Matchmaker,
|
|
33
|
+
PromotionGate,
|
|
34
|
+
Rater,
|
|
35
|
+
RatingTable,
|
|
36
|
+
SnapshotStore,
|
|
37
|
+
)
|
|
38
|
+
from .metrics import Alarm, AlarmResult, MetricRow, MetricsSink, MetricValue
|
|
39
|
+
from .policy import (
|
|
40
|
+
ActionDistribution,
|
|
41
|
+
Actor,
|
|
42
|
+
ActorCritic,
|
|
43
|
+
ActResult,
|
|
44
|
+
BackpropResult,
|
|
45
|
+
Critic,
|
|
46
|
+
NetworkFactory,
|
|
47
|
+
ObsBatch,
|
|
48
|
+
)
|
|
49
|
+
from .rollout import (
|
|
50
|
+
Assignment,
|
|
51
|
+
Close,
|
|
52
|
+
Defer,
|
|
53
|
+
EnvSpec,
|
|
54
|
+
EpisodeRecord,
|
|
55
|
+
ObsKeySpec,
|
|
56
|
+
Plan,
|
|
57
|
+
RolloutRound,
|
|
58
|
+
RolloutSource,
|
|
59
|
+
SetState,
|
|
60
|
+
SlotPlan,
|
|
61
|
+
Spaces,
|
|
62
|
+
Step,
|
|
63
|
+
WorkerCommand,
|
|
64
|
+
WorkerFailure,
|
|
65
|
+
)
|
|
66
|
+
from .schedule import Schedule, ScheduleState
|
|
67
|
+
from .update import ActorLossTerm, ActorTermInputs, Update, UpdateResult
|
|
68
|
+
|
|
69
|
+
__all__ = [
|
|
70
|
+
"ActResult",
|
|
71
|
+
"ActionDistribution",
|
|
72
|
+
"Actor",
|
|
73
|
+
"ActorCritic",
|
|
74
|
+
"ActorLossTerm",
|
|
75
|
+
"ActorTermInputs",
|
|
76
|
+
"AdvantageEstimator",
|
|
77
|
+
"AdvantageStats",
|
|
78
|
+
"Alarm",
|
|
79
|
+
"AlarmResult",
|
|
80
|
+
"Assignment",
|
|
81
|
+
"BackpropResult",
|
|
82
|
+
"CheckpointStore",
|
|
83
|
+
"Checkpointable",
|
|
84
|
+
"Close",
|
|
85
|
+
"CodecTable",
|
|
86
|
+
"ConditionResult",
|
|
87
|
+
"Critic",
|
|
88
|
+
"Defer",
|
|
89
|
+
"EnvSpec",
|
|
90
|
+
"EpisodeRecord",
|
|
91
|
+
"EvictionPolicy",
|
|
92
|
+
"ExperienceBuffer",
|
|
93
|
+
"GateDecision",
|
|
94
|
+
"Manifest",
|
|
95
|
+
"Matchmaker",
|
|
96
|
+
"MetricRow",
|
|
97
|
+
"MetricValue",
|
|
98
|
+
"MetricsSink",
|
|
99
|
+
"NetworkFactory",
|
|
100
|
+
"ObsBatch",
|
|
101
|
+
"ObsCodec",
|
|
102
|
+
"ObsKeySpec",
|
|
103
|
+
"Plan",
|
|
104
|
+
"PromotionGate",
|
|
105
|
+
"Rater",
|
|
106
|
+
"RatingTable",
|
|
107
|
+
"RngState",
|
|
108
|
+
"RolloutRound",
|
|
109
|
+
"RolloutSource",
|
|
110
|
+
"Schedule",
|
|
111
|
+
"ScheduleState",
|
|
112
|
+
"SetState",
|
|
113
|
+
"SlotPlan",
|
|
114
|
+
"SnapshotStore",
|
|
115
|
+
"Spaces",
|
|
116
|
+
"Step",
|
|
117
|
+
"Update",
|
|
118
|
+
"UpdateResult",
|
|
119
|
+
"WorkerCommand",
|
|
120
|
+
"WorkerFailure",
|
|
121
|
+
]
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
"""How a return is turned into an advantage.
|
|
2
|
+
|
|
3
|
+
One ABC, because the estimator is the piece most likely to be replaced -- V-trace against a
|
|
4
|
+
frozen pool, a different lambda rule, a per-bucket normalisation -- and because the one thing it
|
|
5
|
+
must never do is carry its recursion across an episode boundary.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
from abc import ABC, abstractmethod
|
|
11
|
+
from typing import TYPE_CHECKING
|
|
12
|
+
|
|
13
|
+
import msgspec
|
|
14
|
+
|
|
15
|
+
from .checkpoint import Checkpointable
|
|
16
|
+
|
|
17
|
+
if TYPE_CHECKING: # pragma: no cover - annotations only
|
|
18
|
+
from torch import Tensor
|
|
19
|
+
|
|
20
|
+
__all__ = ["AdvantageEstimator", "AdvantageStats"]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class AdvantageStats(msgspec.Struct):
|
|
24
|
+
"""What the estimator saw, for the metric row.
|
|
25
|
+
|
|
26
|
+
``reward_scale`` is the divisor the return scaler applied and ``clipped_reward_frac`` is how
|
|
27
|
+
much of the batch the clip bound touched: both are how a scaled reward stays legible after
|
|
28
|
+
the scaling.
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
raw_return_mean: float
|
|
32
|
+
raw_return_std: float
|
|
33
|
+
reward_scale: float
|
|
34
|
+
clipped_reward_frac: float
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class AdvantageEstimator(ABC, Checkpointable):
|
|
38
|
+
"""Rewards and values in, advantages and returns out."""
|
|
39
|
+
|
|
40
|
+
@abstractmethod
|
|
41
|
+
def compute(
|
|
42
|
+
self,
|
|
43
|
+
*,
|
|
44
|
+
rewards: Tensor,
|
|
45
|
+
values: Tensor,
|
|
46
|
+
final_values: Tensor,
|
|
47
|
+
terminated: Tensor,
|
|
48
|
+
truncated: Tensor,
|
|
49
|
+
trainable: Tensor,
|
|
50
|
+
gamma: float,
|
|
51
|
+
lam: float,
|
|
52
|
+
) -> tuple[Tensor, Tensor, AdvantageStats]:
|
|
53
|
+
"""``rewards``/``terminated``/``truncated``/``trainable`` are ``(T, R)``; ``values`` is
|
|
54
|
+
``(T+1, R)``; ``final_values`` is ``(T, R)`` and is read only where truncated.
|
|
55
|
+
Returns ``(advantages (T, R), returns (T, R), stats)``.
|
|
56
|
+
|
|
57
|
+
Implementers MUST bootstrap a terminated cell from 0 and a truncated cell from
|
|
58
|
+
``final_values``, and MUST NOT carry the recursion across an episode boundary. A
|
|
59
|
+
truncation is an episode that was cut, not one that was decided, and treating the two
|
|
60
|
+
alike throws away the value of every position a step limit ended.
|
|
61
|
+
"""
|
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""How an observation is stored, and the rectangle it is stored in.
|
|
2
|
+
|
|
3
|
+
The two are one subject. The worker packs a row straight into its final resting place in the
|
|
4
|
+
experience buffer, so there is no intermediate copy to agree about; the learner unpacks it on
|
|
5
|
+
the device inside the same kernel that gathers the frame stack. What the codec may decide is
|
|
6
|
+
how each key is stored, and that is decided from the observation space at preflight rather than
|
|
7
|
+
from a list of plane indices written down here.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from abc import ABC, abstractmethod
|
|
13
|
+
from collections.abc import Callable, Iterator, Sequence
|
|
14
|
+
from typing import TYPE_CHECKING
|
|
15
|
+
|
|
16
|
+
import msgspec
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from .checkpoint import Checkpointable
|
|
20
|
+
|
|
21
|
+
if TYPE_CHECKING: # pragma: no cover - annotations only
|
|
22
|
+
from torch import Tensor
|
|
23
|
+
|
|
24
|
+
from ..learn.buffer import Batch
|
|
25
|
+
from ..rollout.layout import BufferHandle
|
|
26
|
+
from .policy import ObsBatch
|
|
27
|
+
from .rollout import EnvSpec, RolloutRound, SlotPlan
|
|
28
|
+
|
|
29
|
+
__all__ = ["MIN_TABLE_STATES", "CodecTable", "ExperienceBuffer", "ObsCodec"]
|
|
30
|
+
|
|
31
|
+
#: How many real observations a codec table may be decided from. Storage is decided per plane
|
|
32
|
+
#: from what the sample contains, so the sample has to be large enough, and played rather than
|
|
33
|
+
#: idle, to have reached the states in which a plane takes the values that decide it.
|
|
34
|
+
MIN_TABLE_STATES = 1000
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class CodecTable(msgspec.Struct, frozen=True, omit_defaults=True):
|
|
38
|
+
"""How each observation key is stored, decided from ``EnvSpec.obs_space`` at preflight
|
|
39
|
+
rather than from a list of plane indices. Logged, hashed and written into every snapshot.
|
|
40
|
+
|
|
41
|
+
``plane`` is one entry per spatial plane: its name, its storage in
|
|
42
|
+
{"uint8", "float16", "static", "derived"}, and the divisor that takes the stored integer
|
|
43
|
+
back to the value the environment produced. "derived" is reserved: nothing rebuilds such a
|
|
44
|
+
plane yet, so the codec refuses it at bind. Two runs whose tables differ are not comparable,
|
|
45
|
+
and ``digest()`` is what says so.
|
|
46
|
+
"""
|
|
47
|
+
|
|
48
|
+
plane: tuple[tuple[str, str, float], ...]
|
|
49
|
+
vector: str
|
|
50
|
+
mask: str
|
|
51
|
+
#: How ``card_ids`` is stored: "uint8", exact, or None when the observation has none. It is
|
|
52
|
+
#: omitted from the encoding when None (``omit_defaults``), so every table decided before
|
|
53
|
+
#: card identity existed hashes exactly as it did and no saved run's digest moves.
|
|
54
|
+
ids: str | None = None
|
|
55
|
+
|
|
56
|
+
def digest(self) -> str:
|
|
57
|
+
"""sha256 of the canonical JSON of this table."""
|
|
58
|
+
from ..rollout.envspec import digest_of
|
|
59
|
+
|
|
60
|
+
return digest_of(self)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class ObsCodec(ABC):
|
|
64
|
+
"""Quantisation of one observation row. The worker packs; the learner unpacks on the GPU.
|
|
65
|
+
|
|
66
|
+
Implementers MUST be exact round-trips for the integer-valued channels and MUST declare
|
|
67
|
+
``row_bytes`` as a constant given an ``EnvSpec`` and its ``CodecTable``.
|
|
68
|
+
"""
|
|
69
|
+
|
|
70
|
+
@abstractmethod
|
|
71
|
+
def table(
|
|
72
|
+
self,
|
|
73
|
+
spec: EnvSpec,
|
|
74
|
+
sample: Sequence[dict[str, np.ndarray]],
|
|
75
|
+
*,
|
|
76
|
+
min_states: int = MIN_TABLE_STATES,
|
|
77
|
+
) -> CodecTable:
|
|
78
|
+
"""Decide storage per key from the declared bounds and a sample of real observations.
|
|
79
|
+
|
|
80
|
+
Storage is decided from a sample; EXISTENCE is decided from the declaration. A plane the
|
|
81
|
+
layout does not declare static is stored even when it is constant across the sample,
|
|
82
|
+
because the tower planes are constant in any sample in which no tower falls.
|
|
83
|
+
|
|
84
|
+
Implementers MUST refuse a sample of fewer than ``min_states`` observations. The row
|
|
85
|
+
size of the whole run follows from this one decision, and a sample too small or too
|
|
86
|
+
idle to have reached the states a plane varies in decides it wrongly and in silence.
|
|
87
|
+
"""
|
|
88
|
+
|
|
89
|
+
@abstractmethod
|
|
90
|
+
def row_bytes(self, spec: EnvSpec) -> int: ...
|
|
91
|
+
|
|
92
|
+
@abstractmethod
|
|
93
|
+
def pack(self, obs: dict[str, np.ndarray], out: memoryview, row: int) -> None: ...
|
|
94
|
+
|
|
95
|
+
@abstractmethod
|
|
96
|
+
def static_planes(self, obs: dict[str, np.ndarray]) -> np.ndarray:
|
|
97
|
+
"""The planes ``EnvSpec.spatial_layout`` declares static; stored once per seat, never
|
|
98
|
+
per row."""
|
|
99
|
+
|
|
100
|
+
@abstractmethod
|
|
101
|
+
def unpack_to_device(self, raw: Tensor, statics: Tensor, out: ObsBatch) -> None:
|
|
102
|
+
"""Dequantise, scatter the static planes in, reshape the stored mask into the mask
|
|
103
|
+
planes, and gather the frame-stack history."""
|
|
104
|
+
|
|
105
|
+
@property
|
|
106
|
+
@abstractmethod
|
|
107
|
+
def codec_version(self) -> int:
|
|
108
|
+
"""The RULE's version. The table it produces is data and travels separately."""
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
class ExperienceBuffer(ABC, Checkpointable):
|
|
112
|
+
"""A rectangle of ``(T + frame_stack)`` cycles x ``R`` slots: ``T`` collected cycles, one
|
|
113
|
+
bootstrap row, and ``frame_stack - 1`` history rows carried over from the previous
|
|
114
|
+
iteration. Owns the shared-memory block the workers write into.
|
|
115
|
+
|
|
116
|
+
Implementers may assume each ``(cycle, slot)`` cell is written exactly once, by the worker
|
|
117
|
+
that owns that slot; the learner only reads.
|
|
118
|
+
"""
|
|
119
|
+
|
|
120
|
+
@abstractmethod
|
|
121
|
+
def shared_handle(self) -> BufferHandle:
|
|
122
|
+
"""Name, size and the numbers the offsets follow from; picklable, and sent to workers."""
|
|
123
|
+
|
|
124
|
+
@abstractmethod
|
|
125
|
+
def begin_iteration(self, plan: SlotPlan, cycles: int) -> None: ...
|
|
126
|
+
|
|
127
|
+
@abstractmethod
|
|
128
|
+
def record_round(
|
|
129
|
+
self, r: RolloutRound, actions: np.ndarray, log_probs: np.ndarray
|
|
130
|
+
) -> None:
|
|
131
|
+
"""Scalars only: observations are already in place. O(n), no observation copy."""
|
|
132
|
+
|
|
133
|
+
@abstractmethod
|
|
134
|
+
def set_values(self, values: Tensor) -> None:
|
|
135
|
+
"""``(T+1, R)`` float32, from the whole-iteration critic pass."""
|
|
136
|
+
|
|
137
|
+
@abstractmethod
|
|
138
|
+
def set_n_legal(self, counts: Tensor) -> None:
|
|
139
|
+
"""``(T+1, R)`` integer: how many actions each cell's mask left, from the critic's pass.
|
|
140
|
+
|
|
141
|
+
Implementers keep the collected cycles and may drop the bootstrap row, in which no
|
|
142
|
+
action was taken. The count is what tells the update which rows had a choice at all
|
|
143
|
+
without unpacking an observation to find out.
|
|
144
|
+
"""
|
|
145
|
+
|
|
146
|
+
@abstractmethod
|
|
147
|
+
def set_final_values(self, cells: np.ndarray, values: Tensor) -> None:
|
|
148
|
+
"""V(final_obs) for truncated cells; ``cells`` is int64[(k, 2)] of (cycle, slot)."""
|
|
149
|
+
|
|
150
|
+
@abstractmethod
|
|
151
|
+
def set_advantages(self, adv: Tensor, ret: Tensor) -> None:
|
|
152
|
+
"""``(T, R)`` float32 each."""
|
|
153
|
+
|
|
154
|
+
@abstractmethod
|
|
155
|
+
def trainable_mask(self) -> Tensor:
|
|
156
|
+
"""``(T, R)`` bool: which cells reach the update. What decides it is the seat's group
|
|
157
|
+
and the cell's validity, never where the row was written."""
|
|
158
|
+
|
|
159
|
+
@abstractmethod
|
|
160
|
+
def batches(
|
|
161
|
+
self,
|
|
162
|
+
batch_size: int,
|
|
163
|
+
minibatch_size: int,
|
|
164
|
+
epochs: int,
|
|
165
|
+
rng_for_epoch: Callable[[int], np.random.Generator],
|
|
166
|
+
*,
|
|
167
|
+
choice_first: bool = False,
|
|
168
|
+
) -> Iterator[Batch]:
|
|
169
|
+
"""Yield batches; each Batch knows its true sample count and iterates device-resident
|
|
170
|
+
minibatches. A batch never straddles an epoch boundary, and an epoch holds as many whole
|
|
171
|
+
batches as it can fill with the rows over spread one each across them, so every batch is
|
|
172
|
+
at least ``batch_size`` and an epoch is exactly ``n // batch_size`` optimizer steps. An
|
|
173
|
+
epoch with nothing trainable in it yields no batches at all. Each minibatch is weighted by
|
|
174
|
+
its share of its own batch. Gathers per MINIBATCH, never per batch.
|
|
175
|
+
|
|
176
|
+
``choice_first`` reorders each batch's cells so that the ones with more than one legal
|
|
177
|
+
action come first, keeping the permutation's order inside each class. Implementers MUST
|
|
178
|
+
move no cell between batches: it is a reordering, so that a caller skipping the forced
|
|
179
|
+
rows skips whole minibatches of them, and every batch-level denominator is unchanged."""
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
"""What a checkpoint is: components that save themselves, a store that writes them atomically,
|
|
2
|
+
and a manifest that says what was written and proves the bytes are still what was written.
|
|
3
|
+
|
|
4
|
+
Every component owns its folder and its own pair of methods. Adding a component to a checkpoint
|
|
5
|
+
is adding a folder name and a dict entry; there is no central serialiser to edit, and no place
|
|
6
|
+
where a component's state is described twice.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from abc import ABC, abstractmethod
|
|
12
|
+
from collections.abc import Mapping
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Any, Protocol, runtime_checkable
|
|
15
|
+
|
|
16
|
+
import msgspec
|
|
17
|
+
|
|
18
|
+
# Imported at run time, not behind TYPE_CHECKING: Manifest carries a RunIdentity, and msgspec
|
|
19
|
+
# resolves a Struct's annotations when a decoder for it is built, so a name that exists only
|
|
20
|
+
# for a type checker would make a manifest unreadable.
|
|
21
|
+
from ..identity import RunIdentity
|
|
22
|
+
|
|
23
|
+
__all__ = ["CheckpointStore", "Checkpointable", "Manifest", "RngState"]
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
@runtime_checkable
|
|
27
|
+
class Checkpointable(Protocol):
|
|
28
|
+
"""Anything that goes into a checkpoint folder of its own.
|
|
29
|
+
|
|
30
|
+
``FORMAT_VERSION`` is the component's own, not the checkpoint's: a component may change its
|
|
31
|
+
files without the checkpoint format moving, and the manifest records each one so that a
|
|
32
|
+
load can say which component it could not read.
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
FORMAT_VERSION: int
|
|
36
|
+
|
|
37
|
+
def save_checkpoint(self, folder: Path) -> None: ...
|
|
38
|
+
|
|
39
|
+
def load_checkpoint(self, folder: Path, *, strict: bool) -> None:
|
|
40
|
+
"""With ``strict=False``, a missing file prints the exact path it wanted and continues
|
|
41
|
+
with a default. With ``strict=True`` it raises. ``strict`` defaults to True on resume:
|
|
42
|
+
tolerate-everything is right for a research tool and wrong for a harness that promises
|
|
43
|
+
the curve continues."""
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class Manifest(msgspec.Struct):
|
|
47
|
+
"""What one checkpoint is, beside the folders that hold it.
|
|
48
|
+
|
|
49
|
+
``config`` is the resolved config verbatim, so a checkpoint is self-describing without the
|
|
50
|
+
run directory around it, and ``files`` is every file in the checkpoint with its sha256, so a
|
|
51
|
+
truncated write from a crash six weeks ago is an error on load rather than a silently wrong
|
|
52
|
+
resume.
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
format_version: int
|
|
56
|
+
run_id: str
|
|
57
|
+
run_name: str
|
|
58
|
+
identity: RunIdentity
|
|
59
|
+
config: dict[str, Any]
|
|
60
|
+
config_hash: str
|
|
61
|
+
iteration: int
|
|
62
|
+
cumulative_env_steps: int
|
|
63
|
+
cumulative_timesteps: int
|
|
64
|
+
cumulative_model_updates: int
|
|
65
|
+
wall_seconds: float
|
|
66
|
+
created_unix_ns: int
|
|
67
|
+
state_digest: str
|
|
68
|
+
component_versions: dict[str, int]
|
|
69
|
+
files: dict[str, str]
|
|
70
|
+
#: Seconds spent gating so far, and the last gate's decision: run state, like the wall clock.
|
|
71
|
+
#: None on a manifest written before either was recorded.
|
|
72
|
+
gate_seconds: float | None = None
|
|
73
|
+
#: Stored as plain data and converted by the coordinator, so this module need not import
|
|
74
|
+
#: the ladder's types: ``api.ladder`` already imports this one.
|
|
75
|
+
last_decision: dict[str, Any] | None = None
|
|
76
|
+
#: The env step each cadence last fired at: the periodic checkpoint, the gate's candidate
|
|
77
|
+
#: and the floor's admission. A resume puts them back, so each fires where it would have in
|
|
78
|
+
#: a run that never stopped. None on a manifest written before they were recorded; a resume
|
|
79
|
+
#: then starts all three at this checkpoint's env step, as every resume did before.
|
|
80
|
+
last_checkpoint_step: int | None = None
|
|
81
|
+
last_candidate_step: int | None = None
|
|
82
|
+
last_floor_step: int | None = None
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class RngState(msgspec.Struct):
|
|
86
|
+
"""Every random stream's position at the moment a checkpoint was written.
|
|
87
|
+
|
|
88
|
+
One reference learner saves none of this; the other re-seeds from the run's INITIAL seed on
|
|
89
|
+
load, so a run resumed at ten million steps draws the same action noise it drew at step
|
|
90
|
+
zero. Because every stream here is name-addressed, restoring the iteration counter and the
|
|
91
|
+
shard positions restores the stream, not merely the parameters.
|
|
92
|
+
"""
|
|
93
|
+
|
|
94
|
+
master_seed: int
|
|
95
|
+
torch_cpu: str
|
|
96
|
+
torch_cuda: list[str]
|
|
97
|
+
python_random: list[Any]
|
|
98
|
+
numpy_minibatch: dict[str, Any]
|
|
99
|
+
iteration: int
|
|
100
|
+
shard_streams: list[dict[str, Any]]
|
|
101
|
+
eval_seed_set_sha: str
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
class CheckpointStore(ABC):
|
|
105
|
+
"""Where checkpoints live, and the only thing that writes or reads them."""
|
|
106
|
+
|
|
107
|
+
@abstractmethod
|
|
108
|
+
def write(self, components: Mapping[str, Checkpointable], manifest: Manifest) -> Path:
|
|
109
|
+
"""Write atomically: into ``<name>.partial/``, fsync each file, then ``os.replace``.
|
|
110
|
+
|
|
111
|
+
The directory itself is not fsynced. ``os.fsync`` on a directory handle is a POSIX
|
|
112
|
+
guarantee and raises on Windows, where this harness's default profile runs, so the
|
|
113
|
+
durability step is per file and the atomic step is the rename -- which is atomic on both
|
|
114
|
+
platforms, and is the property recovery actually needs.
|
|
115
|
+
"""
|
|
116
|
+
|
|
117
|
+
@abstractmethod
|
|
118
|
+
def read(
|
|
119
|
+
self, path: Path, components: Mapping[str, Checkpointable], *, strict: bool
|
|
120
|
+
) -> Manifest:
|
|
121
|
+
"""Verify every hash in the manifest, then load each component. A mismatch raises
|
|
122
|
+
``CheckpointFormatError`` naming the first file that failed."""
|
|
123
|
+
|
|
124
|
+
@abstractmethod
|
|
125
|
+
def latest(self, run_dir: Path) -> Path | None:
|
|
126
|
+
"""The newest checkpoint, from the run index rather than from parsing directory names:
|
|
127
|
+
``int(x) for x in os.listdir(...)`` crashes on a stray file, and the save is called from
|
|
128
|
+
the crash handler, which is where the user most needs it not to."""
|
|
129
|
+
|
|
130
|
+
@abstractmethod
|
|
131
|
+
def prune(self, run_dir: Path, keep: int) -> list[Path]:
|
|
132
|
+
"""Remove all but the newest ``keep``, and return what was removed."""
|