royaleimitate 0.2.6__tar.gz

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.
Files changed (43) hide show
  1. royaleimitate-0.2.6/LICENSE +21 -0
  2. royaleimitate-0.2.6/PKG-INFO +22 -0
  3. royaleimitate-0.2.6/README.md +41 -0
  4. royaleimitate-0.2.6/pyproject.toml +74 -0
  5. royaleimitate-0.2.6/royaleimitate/__init__.py +55 -0
  6. royaleimitate-0.2.6/royaleimitate/__main__.py +6 -0
  7. royaleimitate-0.2.6/royaleimitate/alarms.py +36 -0
  8. royaleimitate-0.2.6/royaleimitate/artifacts.py +310 -0
  9. royaleimitate-0.2.6/royaleimitate/cli.py +84 -0
  10. royaleimitate-0.2.6/royaleimitate/cloning.py +226 -0
  11. royaleimitate-0.2.6/royaleimitate/config.py +240 -0
  12. royaleimitate-0.2.6/royaleimitate/extension.py +105 -0
  13. royaleimitate-0.2.6/royaleimitate/field_model.py +98 -0
  14. royaleimitate-0.2.6/royaleimitate/fit.py +170 -0
  15. royaleimitate-0.2.6/royaleimitate/init.py +111 -0
  16. royaleimitate-0.2.6/royaleimitate/public_log.py +417 -0
  17. royaleimitate-0.2.6/royaleimitate/references.py +184 -0
  18. royaleimitate-0.2.6/royaleimitate/regularisers.py +435 -0
  19. royaleimitate-0.2.6/royaleimitate/replay_cards.json +121 -0
  20. royaleimitate-0.2.6/royaleimitate/replays.py +324 -0
  21. royaleimitate-0.2.6/royaleimitate/schema.py +93 -0
  22. royaleimitate-0.2.6/royaleimitate/shards.py +591 -0
  23. royaleimitate-0.2.6/royaleimitate/split.py +36 -0
  24. royaleimitate-0.2.6/royaleimitate/warm_start.py +141 -0
  25. royaleimitate-0.2.6/royaleimitate.egg-info/PKG-INFO +22 -0
  26. royaleimitate-0.2.6/royaleimitate.egg-info/SOURCES.txt +41 -0
  27. royaleimitate-0.2.6/royaleimitate.egg-info/dependency_links.txt +1 -0
  28. royaleimitate-0.2.6/royaleimitate.egg-info/entry_points.txt +6 -0
  29. royaleimitate-0.2.6/royaleimitate.egg-info/requires.txt +17 -0
  30. royaleimitate-0.2.6/royaleimitate.egg-info/top_level.txt +1 -0
  31. royaleimitate-0.2.6/setup.cfg +4 -0
  32. royaleimitate-0.2.6/tests/test_clone.py +126 -0
  33. royaleimitate-0.2.6/tests/test_extension_contract.py +68 -0
  34. royaleimitate-0.2.6/tests/test_field_reference.py +176 -0
  35. royaleimitate-0.2.6/tests/test_imitation_alarms.py +140 -0
  36. royaleimitate-0.2.6/tests/test_imitation_anchor.py +711 -0
  37. royaleimitate-0.2.6/tests/test_imitation_init.py +560 -0
  38. royaleimitate-0.2.6/tests/test_public_log.py +767 -0
  39. royaleimitate-0.2.6/tests/test_replays.py +167 -0
  40. royaleimitate-0.2.6/tests/test_save_actor.py +91 -0
  41. royaleimitate-0.2.6/tests/test_shards.py +254 -0
  42. royaleimitate-0.2.6/tests/test_supported_surface.py +60 -0
  43. royaleimitate-0.2.6/tests/test_wheel_install.py +93 -0
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 RoyaleGym contributors
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.
@@ -0,0 +1,22 @@
1
+ Metadata-Version: 2.4
2
+ Name: royaleimitate
3
+ Version: 0.2.6
4
+ Summary: RoyaleImitate: imitation learning for RoyaleLearn, as an optional add-on
5
+ License: MIT
6
+ Requires-Python: >=3.12
7
+ License-File: LICENSE
8
+ Requires-Dist: royalelearn>=0.5.4
9
+ Requires-Dist: royalegym>=0.1.0
10
+ Requires-Dist: numpy>=2.0
11
+ Requires-Dist: msgspec>=0.18
12
+ Provides-Extra: torch
13
+ Requires-Dist: torch>=2.4; extra == "torch"
14
+ Requires-Dist: safetensors>=0.4; extra == "torch"
15
+ Provides-Extra: replays
16
+ Requires-Dist: huggingface_hub>=0.24; extra == "replays"
17
+ Requires-Dist: pyarrow>=15; extra == "replays"
18
+ Provides-Extra: dev
19
+ Requires-Dist: pytest>=8.0; extra == "dev"
20
+ Requires-Dist: ruff>=0.16; extra == "dev"
21
+ Requires-Dist: pyarrow>=15; extra == "dev"
22
+ Dynamic: license-file
@@ -0,0 +1,41 @@
1
+ <p align="center"><img src="docs/media/logo.png" width="128" alt="The RoyaleImitate logo: a purple crown shield with two white cards on it, outlined in gold"></p><h1 align="center">RoyaleImitate</h1>
2
+
3
+ <p align="center"><a href="https://github.com/RoyaleGym/RoyaleImitate/actions/workflows/suite.yml"><img alt="CI" src="https://github.com/RoyaleGym/RoyaleImitate/actions/workflows/suite.yml/badge.svg"></a> <img alt="License" src="https://img.shields.io/github/license/RoyaleGym/RoyaleImitate?style=flat-square&color=555"> <img alt="Python" src="https://img.shields.io/badge/python-3.12%20%7C%203.13%20%7C%203.14-3776AB?style=flat-square&logo=python&logoColor=white"> <a href="https://royalegym.github.io/RoyaleGym/"><img alt="Docs" src="https://img.shields.io/badge/docs-royalegym.github.io-8957e5?style=flat-square&logo=readthedocs&logoColor=white"></a> <a href="https://discord.gg/4D2BS5JBHP"><img alt="Discord" src="https://img.shields.io/discord/1551699576304705647?style=flat-square&logo=discord&logoColor=white&label=discord&color=5865F2"></a> <img alt="Last commit" src="https://img.shields.io/github/last-commit/RoyaleGym/RoyaleImitate?style=flat-square&color=555"></p>
4
+
5
+ Start a Clash Royale bot from one you already have, and keep it close to that bot while it learns.
6
+ It is an optional add-on to RoyaleLearn, the trainer: leave it out and nothing changes.
7
+
8
+ ## Install
9
+
10
+ pip install "royalegym[all]"
11
+
12
+ This is the `[imitate]` part. Until it is on PyPI, see [Install](https://royalegym.github.io/RoyaleGym/install/) for the exact line.
13
+
14
+ ## Try it
15
+
16
+ ```python
17
+ from royalegym import TowerHPReward, make_env
18
+ from royalelearn import Learner
19
+ from royaleimitate import save_actor
20
+
21
+ def build_env():
22
+ return make_env(reward=TowerHPReward())
23
+
24
+ teacher = Learner(build_env, save_dir="runs/teacher")
25
+ teacher.learn(total_steps=20_000)
26
+ digest = save_actor(teacher, "runs/teacher-actor")
27
+
28
+ start = {"path": "runs/teacher-actor", "sha256": digest}
29
+ student = Learner(build_env, save_dir="runs/student", extensions={"warm_start": {"init": start}})
30
+ student.learn(total_steps=20_000)
31
+ ```
32
+
33
+ The student starts from the teacher's weights; [examples/minimal.py](examples/minimal.py) also keeps it close.
34
+
35
+ ## Next
36
+
37
+ - The docs: [royalegym.github.io/RoyaleGym](https://royalegym.github.io/RoyaleGym/), and the [RoyaleImitate page](https://royalegym.github.io/RoyaleGym/resources/royaleimitate/)
38
+ - The guide: [guide](https://royalegym.github.io/RoyaleGym/repos/royaleimitate/guide/). The full contract (Advanced): [spec](https://royalegym.github.io/RoyaleGym/repos/royaleimitate/spec/)
39
+ - Questions: [Discord](https://discord.gg/4D2BS5JBHP)
40
+
41
+ MIT licensed. See [LICENSE](LICENSE).
@@ -0,0 +1,74 @@
1
+ [build-system]
2
+ requires = ["setuptools>=68"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "royaleimitate"
7
+ version = "0.2.6"
8
+ description = "RoyaleImitate: imitation learning for RoyaleLearn, as an optional add-on"
9
+ requires-python = ">=3.12"
10
+ license = { text = "MIT" }
11
+ # Neither is on PyPI: both are sibling checkouts, installed editable before this package
12
+ # (RoyaleSim -> RoyaleGym -> RoyaleViser -> RoyaleLearn -> RoyaleImitate). pip resolves them from
13
+ # the venv; it never fetches them.
14
+ dependencies = [
15
+ "royalelearn>=0.5.4",
16
+ "royalegym>=0.1.0",
17
+ "numpy>=2.0",
18
+ "msgspec>=0.18",
19
+ ]
20
+
21
+ [project.optional-dependencies]
22
+ torch = [
23
+ "torch>=2.4",
24
+ "safetensors>=0.4",
25
+ ]
26
+ # Human games from the IL_Replay dataset (``royaleimitate.from_replays``).
27
+ replays = [
28
+ "huggingface_hub>=0.24",
29
+ "pyarrow>=15",
30
+ ]
31
+ dev = [
32
+ "pytest>=8.0",
33
+ "ruff>=0.16",
34
+ "pyarrow>=15",
35
+ ]
36
+
37
+ # How RoyaleLearn finds the two config sections this package provides. The entry point's name is
38
+ # the section's key.
39
+ [project.entry-points."royalelearn.extensions"]
40
+ imitation = "royaleimitate.extension:EXTENSION"
41
+ warm_start = "royaleimitate.warm_start:EXTENSION"
42
+
43
+ [project.scripts]
44
+ royaleimitate = "royaleimitate.cli:main"
45
+
46
+ # Naming the package explicitly keeps anything else in this folder out of a wheel built here.
47
+ [tool.setuptools.packages.find]
48
+ include = ["royaleimitate", "royaleimitate.*"]
49
+
50
+ [tool.setuptools.package-data]
51
+ royaleimitate = ["replay_cards.json"]
52
+
53
+ [tool.pytest.ini_options]
54
+ testpaths = ["tests"]
55
+ python_files = ["test_*.py"]
56
+ pythonpath = ["."]
57
+ addopts = "-ra --strict-markers -m \"not slow and not engine\""
58
+ markers = [
59
+ "slow: takes more than about five seconds",
60
+ "engine: needs a royalesim build that matches the data on disk",
61
+ ]
62
+
63
+ [tool.ruff]
64
+ line-length = 100
65
+ target-version = "py312"
66
+ src = ["royaleimitate", "tests"]
67
+ extend-exclude = ["datasets"]
68
+
69
+ [tool.ruff.lint]
70
+ select = ["E", "F", "W", "I", "B", "UP", "SIM", "RUF"]
71
+ ignore = ["RUF001", "RUF002", "RUF003", "SIM108"]
72
+
73
+ [tool.ruff.lint.isort]
74
+ known-first-party = ["royaleimitate", "royalelearn", "royalegym", "royalesim"]
@@ -0,0 +1,55 @@
1
+ """Imitation learning for RoyaleLearn, as an optional add-on.
2
+
3
+ Two config sections, found by RoyaleLearn through this package's entry points:
4
+
5
+ - ``warm_start``: start the actor from saved weights, check them against their own probe rows,
6
+ and hold the actor still at first while the critic catches up;
7
+ - ``imitation``: anchor the policy to reference policies with the reference-KL regulariser and
8
+ its adaptive coefficient.
9
+
10
+ And the tools that make what they read: saved-policy folders (``artifacts``), field models
11
+ (``field_model``, ``fit``), and demonstration shards (``shards``, ``split``). And ``public_log``:
12
+ the observation's fair fields from a timed log of card plays, with no engine running.
13
+ And ``from_replays``: human games from the public IL_Replay dataset, as demonstration shards.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ __version__ = "0.2.6"
19
+
20
+ __all__ = ["clone", "from_replays", "record", "save_actor"]
21
+
22
+
23
+ def save_actor(learner: object, folder: object, *, seed: int = 0) -> str:
24
+ """Write a trained ``royalelearn.Learner``'s actor as an actor artifact; returns its digest.
25
+
26
+ See ``royaleimitate.artifacts.save_actor``. Imported on first use, so importing this package
27
+ does not import torch.
28
+ """
29
+ from .artifacts import save_actor as save
30
+
31
+ return save(learner, folder, seed=seed) # type: ignore[arg-type]
32
+
33
+
34
+ def record(learner: object, teacher: object, out: object, **kwargs: object) -> object:
35
+ """A teacher's battles in ``learner``'s environment, as shard rows; see
36
+ ``royaleimitate.cloning.record``."""
37
+ from .cloning import record as run
38
+
39
+ return run(learner, teacher, out, **kwargs) # type: ignore[arg-type]
40
+
41
+
42
+ def clone(learner: object, demonstrations: object, out: object, **kwargs: object) -> str:
43
+ """``learner``'s network trained to copy ``demonstrations``, written to ``out``; see
44
+ ``royaleimitate.cloning.clone``."""
45
+ from .cloning import clone as run
46
+
47
+ return run(learner, demonstrations, out, **kwargs) # type: ignore[arg-type]
48
+
49
+
50
+ def from_replays(learner: object, out: object, **kwargs: object) -> object:
51
+ """Human games from the IL_Replay dataset, replayed in ``learner``'s environment, as shard
52
+ rows; see ``royaleimitate.replays.from_replays``."""
53
+ from .replays import from_replays as run
54
+
55
+ return run(learner, out, **kwargs) # type: ignore[arg-type]
@@ -0,0 +1,6 @@
1
+ """``python -m royaleimitate``: the same command as the ``royaleimitate`` script, for a shell where
2
+ the virtual environment's scripts folder is not on PATH."""
3
+
4
+ from .cli import main
5
+
6
+ raise SystemExit(main())
@@ -0,0 +1,36 @@
1
+ """The regularisers' alarms (section 19.9). Both warn: a halted treatment run is censored out of
2
+ its comparison."""
3
+
4
+ from __future__ import annotations
5
+
6
+ from typing import TYPE_CHECKING
7
+
8
+ from royalelearn.extensions import FamilyAlarm
9
+
10
+ from .schema import IMITATION_ALARM_METRICS
11
+
12
+ if TYPE_CHECKING: # pragma: no cover - annotations only
13
+ from royalelearn.extensions import Alarm
14
+
15
+ __all__ = ["imitation_alarms"]
16
+
17
+
18
+ def imitation_alarms(*, ref_kl_warn: float, lambda_saturated_patience: int) -> list[Alarm]:
19
+ return [
20
+ FamilyAlarm(
21
+ "imitation_ref_kl_high",
22
+ lambda _member, kl: kl > ref_kl_warn,
23
+ keys=IMITATION_ALARM_METRICS["imitation_ref_kl_high"],
24
+ meaning="the policy has moved far from a reference it is anchored to",
25
+ ),
26
+ FamilyAlarm(
27
+ "imitation_lambda_saturated",
28
+ lambda _member, at_max: at_max >= 1.0,
29
+ keys=IMITATION_ALARM_METRICS["imitation_lambda_saturated"],
30
+ patience=lambda_saturated_patience,
31
+ meaning=(
32
+ "lambda has sat at coef.max: the reward is pulling harder than the anchor can "
33
+ "hold, which is a statement about the reward"
34
+ ),
35
+ ),
36
+ ]
@@ -0,0 +1,310 @@
1
+ """Actor artifacts: the folder a saved actor lives in, and how a config names it.
2
+
3
+ Section 19.2 of ``docs/harness-spec.md``. An artifact is a folder, and a config names it by a
4
+ path and a digest. The digest is over every file in the folder, not over the weights alone,
5
+ because the other files change numbers too: a model's normalisation can live in its
6
+ ``spec.json``, and an actor's self-test is in its probe rows.
7
+
8
+ An **actor artifact** is the snapshot layout of section 11.6 -- ``actor.safetensors`` beside a
9
+ ``SnapshotSpec`` -- with the weights in float32 rather than half precision, and a probe set when
10
+ it is meant to be loaded as an init. A pool snapshot is played against and never trained from;
11
+ an init is trained from, and a file that rounded the seeded weights could not reproduce a run
12
+ started from them.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import hashlib
18
+ from collections.abc import Mapping
19
+ from pathlib import Path
20
+ from typing import TYPE_CHECKING, Any, NamedTuple
21
+
22
+ import msgspec
23
+ import numpy as np
24
+
25
+ from royalelearn.extensions import SPEC_NAME, WEIGHTS_NAME, PreflightError, SnapshotSpec, digest_of
26
+
27
+ if TYPE_CHECKING: # pragma: no cover - annotations only
28
+ from torch import Tensor
29
+
30
+ from royalelearn.extensions import ObsBatch, RowCodec
31
+
32
+ __all__ = [
33
+ "PROBE_CHUNK",
34
+ "PROBE_NAME",
35
+ "ActorArtifact",
36
+ "ProbeSet",
37
+ "artifact_digest",
38
+ "check_actor_artifact",
39
+ "load_actor_state",
40
+ "probe_log_probs",
41
+ "read_actor_artifact",
42
+ "save_actor",
43
+ "self_test",
44
+ "verify_artifact",
45
+ "write_actor_artifact",
46
+ ]
47
+
48
+ PROBE_NAME = "probe.safetensors"
49
+ #: The forward the probe log-probabilities are taken in, on both sides. A batch's shape can
50
+ #: choose the kernel a convolution runs with, and a self-test that compared a 1,024-row forward
51
+ #: against a 256-row one would be measuring that choice.
52
+ PROBE_CHUNK = 256
53
+
54
+
55
+ def artifact_digest(folder: str | Path) -> str:
56
+ """sha256 of the canonical JSON of ``{relative path: sha256 of the file}`` over the folder.
57
+
58
+ Every file, at any depth, named by its path relative to the folder with forward slashes, so
59
+ the digest is the same on every platform and does not depend on where the folder sits.
60
+ """
61
+ root = Path(folder)
62
+ if not root.is_dir():
63
+ raise PreflightError(f"no artifact folder at {root}")
64
+ listing = {
65
+ path.relative_to(root).as_posix(): hashlib.sha256(path.read_bytes()).hexdigest()
66
+ for path in sorted(root.rglob("*"))
67
+ if path.is_file()
68
+ }
69
+ if not listing:
70
+ raise PreflightError(f"the artifact folder {root} is empty")
71
+ return digest_of(listing)
72
+
73
+
74
+ def verify_artifact(path: str | Path, stated: str, *, what: str) -> Path:
75
+ """The folder, if its digest is the one the config states; a refusal naming both if not.
76
+
77
+ Checked at every start, fresh or resumed. The identity carries the stated digest, so without
78
+ this the identity is only as true as the config's claim about the file.
79
+ """
80
+ folder = Path(path)
81
+ found = artifact_digest(folder)
82
+ if found != stated:
83
+ raise PreflightError(
84
+ f"{what}: the artifact at {folder} has digest {found}, and the config states "
85
+ f"{stated}. Either the folder changed since the config was written or the config "
86
+ "names the wrong one; `royaleimitate artifact-digest <folder>` prints a folder's digest"
87
+ )
88
+ return folder
89
+
90
+
91
+ # --------------------------------------------------------------------------
92
+ # Actor artifacts
93
+ # --------------------------------------------------------------------------
94
+
95
+
96
+ class ProbeSet(NamedTuple):
97
+ """Rows packed by the run's codec, and the masked log-probabilities recorded on them."""
98
+
99
+ rows: np.ndarray
100
+ log_probs: np.ndarray
101
+
102
+
103
+ class ActorArtifact(NamedTuple):
104
+ folder: Path
105
+ spec: SnapshotSpec
106
+ state: dict[str, Tensor]
107
+ probe: ProbeSet | None
108
+
109
+
110
+ def probe_log_probs(
111
+ actor: Any, codec: RowCodec, rows: np.ndarray, *, chunk: int = PROBE_CHUNK
112
+ ) -> tuple[np.ndarray, np.ndarray]:
113
+ """``(N, n_actions)`` float32 masked log-probabilities and the ``(N, n_actions)`` mask.
114
+
115
+ The one function both sides call: the tool that writes an artifact records what it returns,
116
+ and the self-test compares against what it returns now. Two implementations of "the
117
+ log-probabilities on these rows" would be two answers.
118
+ """
119
+ import torch
120
+
121
+ out: list[np.ndarray] = []
122
+ masks: list[np.ndarray] = []
123
+ with torch.no_grad():
124
+ for start in range(0, int(rows.shape[0]), chunk):
125
+ obs: ObsBatch = codec.decode(rows[start : start + chunk])
126
+ distribution = actor.distribution(obs)
127
+ out.append(distribution.log_probs.float().cpu().numpy())
128
+ masks.append(obs.mask.cpu().numpy())
129
+ if not out:
130
+ width = codec.spec.n_actions
131
+ return np.zeros((0, width), np.float32), np.zeros((0, width), bool)
132
+ return np.concatenate(out), np.concatenate(masks)
133
+
134
+
135
+ def self_test(actor: Any, codec: RowCodec, probe: ProbeSet, *, atol: float, what: str) -> float:
136
+ """The probe-logit self-test of section 19.4. Returns the largest difference it found.
137
+
138
+ A compatible spec says the tensors have the right shapes; only a forward on known rows says
139
+ the network computes the function that was saved. Compared on the legal actions only: an
140
+ illegal action's log-probability is the fill value's residue and carries nothing.
141
+ """
142
+ if probe.rows.shape[0] == 0:
143
+ raise PreflightError(f"{what}: the artifact's probe set is empty")
144
+ measured, mask = probe_log_probs(actor, codec, probe.rows)
145
+ if measured.shape != probe.log_probs.shape:
146
+ raise PreflightError(
147
+ f"{what}: the probe records log-probabilities of shape {probe.log_probs.shape} and "
148
+ f"this run's actor produces {measured.shape}"
149
+ )
150
+ difference = np.where(mask, np.abs(measured - probe.log_probs), 0.0)
151
+ worst = float(difference.max())
152
+ if not worst <= atol:
153
+ row = int(np.unravel_index(int(difference.argmax()), difference.shape)[0])
154
+ raise PreflightError(
155
+ f"{what}: the probe-logit self-test failed. On probe row {row} a legal action's "
156
+ f"log-probability differs by {worst:.3g} from the one recorded when the artifact "
157
+ f"was written, against a tolerance of {atol:g}. The weights loaded, so the network "
158
+ "they were loaded into does not compute the function that was saved: a different "
159
+ "architecture behind the same digest, a codec that decodes the rows differently, or "
160
+ "a different precision or device"
161
+ )
162
+ return worst
163
+
164
+
165
+ def write_actor_artifact(
166
+ folder: str | Path,
167
+ state: Mapping[str, Tensor],
168
+ spec: SnapshotSpec,
169
+ *,
170
+ probe: ProbeSet | None = None,
171
+ ) -> str:
172
+ """Write an actor artifact into a folder that does not exist yet. Returns its digest.
173
+
174
+ The folder must be new: an artifact half overwritten by a second writer is a folder whose
175
+ digest matches nothing anybody wrote down.
176
+ """
177
+ import torch
178
+ from safetensors.torch import save
179
+
180
+ root = Path(folder)
181
+ if root.exists() and any(root.iterdir()):
182
+ raise FileExistsError(f"{root} already holds files; an artifact is written once")
183
+ root.mkdir(parents=True, exist_ok=True)
184
+ tensors = {
185
+ name: tensor.detach().to(dtype=torch.float32, device="cpu").contiguous()
186
+ for name, tensor in state.items()
187
+ }
188
+ (root / WEIGHTS_NAME).write_bytes(save(tensors))
189
+ meta = dict(spec.meta)
190
+ meta["dtype"] = "float32"
191
+ (root / SPEC_NAME).write_bytes(msgspec.json.encode(msgspec.structs.replace(spec, meta=meta)))
192
+ if probe is not None:
193
+ (root / PROBE_NAME).write_bytes(
194
+ save(
195
+ {
196
+ "rows": torch.from_numpy(np.ascontiguousarray(probe.rows, dtype=np.uint8)),
197
+ "log_probs": torch.from_numpy(
198
+ np.ascontiguousarray(probe.log_probs, dtype=np.float32)
199
+ ),
200
+ }
201
+ )
202
+ )
203
+ return artifact_digest(root)
204
+
205
+
206
+ def read_actor_artifact(folder: str | Path) -> ActorArtifact:
207
+ """An actor artifact's spec, its tensors on the CPU, and its probe set if it has one."""
208
+ from safetensors.torch import load
209
+
210
+ root = Path(folder)
211
+ for name in (WEIGHTS_NAME, SPEC_NAME):
212
+ if not (root / name).is_file():
213
+ raise PreflightError(f"{root} is not an actor artifact: it has no {name}")
214
+ spec = msgspec.json.decode((root / SPEC_NAME).read_bytes(), type=SnapshotSpec)
215
+ state = load((root / WEIGHTS_NAME).read_bytes())
216
+ probe = None
217
+ if (root / PROBE_NAME).is_file():
218
+ blob = load((root / PROBE_NAME).read_bytes())
219
+ probe = ProbeSet(rows=blob["rows"].numpy(), log_probs=blob["log_probs"].numpy())
220
+ return ActorArtifact(folder=root, spec=spec, state=dict(state), probe=probe)
221
+
222
+
223
+ def check_actor_artifact(artifact: ActorArtifact, current: SnapshotSpec, *, what: str) -> None:
224
+ """``check_compatible`` plus the action layout, which an init or a reference also needs.
225
+
226
+ A pool snapshot shares the run's action space by construction; an artifact brought in from
227
+ outside the run is exactly the thing that might not.
228
+ """
229
+ from royalelearn.extensions import check_compatible
230
+
231
+ check_compatible(artifact.spec, current)
232
+ if artifact.spec.action_digest and artifact.spec.action_digest != current.action_digest:
233
+ raise PreflightError(
234
+ f"{what}: the artifact was written for action space {artifact.spec.action_digest} and "
235
+ f"this run's is {current.action_digest}"
236
+ )
237
+
238
+
239
+ def load_actor_state(actor: Any, artifact: ActorArtifact, *, what: str) -> None:
240
+ """``load_state_dict(strict=True)``, with a refusal that names the keys that did not match."""
241
+ reference = actor.state_dict()
242
+ missing = sorted(set(reference) - set(artifact.state))
243
+ unexpected = sorted(set(artifact.state) - set(reference))
244
+ if missing or unexpected:
245
+ raise PreflightError(
246
+ f"{what}: the artifact's tensors do not match this run's actor. Missing: "
247
+ f"{missing[:8]}; unexpected: {unexpected[:8]}"
248
+ )
249
+ shapes = [
250
+ name
251
+ for name, tensor in artifact.state.items()
252
+ if tuple(tensor.shape) != tuple(reference[name].shape)
253
+ ]
254
+ if shapes:
255
+ raise PreflightError(f"{what}: tensors of the wrong shape: {shapes[:8]}")
256
+ actor.load_state_dict(
257
+ {name: tensor.to(reference[name].device) for name, tensor in artifact.state.items()},
258
+ strict=True,
259
+ )
260
+
261
+
262
+ #: How many decisions ``save_actor`` plays to make the probe rows, two seats each.
263
+ SAVE_PROBE_STEPS = 32
264
+
265
+
266
+ def save_actor(learner: Any, folder: str | Path, *, seed: int = 0) -> str:
267
+ """Write a ``royalelearn.Learner``'s trained actor as an actor artifact. Returns its digest.
268
+
269
+ The folder is what ``warm_start.init`` and a ``snapshot`` reference read, and the digest is
270
+ the ``sha256`` a config names it by. The probe rows are both seats' observations from a short
271
+ battle of the learner's own ``build_env``, with random legal moves; the log-probabilities the
272
+ actor gives them are recorded beside the weights, so a run that loads the folder can check it
273
+ computes the same policy. Call ``learner.learn`` first.
274
+ """
275
+ run = getattr(learner, "run", None)
276
+ if run is None:
277
+ raise PreflightError("there is no trained actor to save yet: call learner.learn() first")
278
+ codec = run.row_codec()
279
+ env = learner.build_env()
280
+ rng = np.random.default_rng(seed)
281
+ observations: list[Mapping[str, Any]] = []
282
+ try:
283
+ obs, _info = env.reset(seed=seed)
284
+ for _ in range(SAVE_PROBE_STEPS):
285
+ actions = {}
286
+ for agent, seat in obs.items():
287
+ observations.append(seat)
288
+ legal = np.flatnonzero(np.asarray(seat["action_mask"]))
289
+ actions[agent] = int(rng.choice(legal))
290
+ obs, _rewards, terminated, truncated, _info = env.step(actions)
291
+ if any(terminated.values()) or any(truncated.values()):
292
+ obs, _info = env.reset(seed=seed + 1)
293
+ finally:
294
+ close = getattr(env, "close", None)
295
+ if close is not None:
296
+ close()
297
+ rows = codec.pack(observations)
298
+ actor = run.model.actor
299
+ log_probs, _mask = probe_log_probs(actor, codec, rows)
300
+ write_actor_artifact(
301
+ folder,
302
+ actor.state_dict(),
303
+ run.snapshot_template,
304
+ probe=ProbeSet(rows=rows, log_probs=log_probs),
305
+ )
306
+ # So that ``Learner.load_policy(folder)`` plays it, as well as a warm start loading it.
307
+ from royalelearn.extensions import write_policy_record
308
+
309
+ write_policy_record(folder, run.spec, learner.config.net)
310
+ return artifact_digest(folder)
@@ -0,0 +1,84 @@
1
+ """``royaleimitate``: the commands that make what the imitation sections read.
2
+
3
+ - ``artifact-digest <folder>`` prints the digest a config names an artifact folder by.
4
+ - ``fit-field-reference`` fits the play/wait model a ``field_mlp`` reference loads.
5
+
6
+ A training run itself is still ``royalelearn train``: RoyaleLearn finds this package through the
7
+ ``warm_start`` and ``imitation`` sections of the run's config.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import argparse
13
+ import sys
14
+ from collections.abc import Sequence
15
+ from pathlib import Path
16
+
17
+ __all__ = ["build_parser", "main"]
18
+
19
+
20
+ def build_parser() -> argparse.ArgumentParser:
21
+ parser = argparse.ArgumentParser(
22
+ prog="royaleimitate",
23
+ description="Make the artifacts RoyaleLearn's imitation sections read.",
24
+ )
25
+ commands = parser.add_subparsers(dest="command", required=True)
26
+
27
+ digest = commands.add_parser(
28
+ "artifact-digest", help="print the digest a config names an artifact folder by"
29
+ )
30
+ digest.add_argument("folder", type=Path)
31
+ digest.set_defaults(handler=_artifact_digest)
32
+
33
+ fit = commands.add_parser(
34
+ "fit-field-reference",
35
+ help="fit the play/wait model a field_mlp reference loads",
36
+ )
37
+ fit.add_argument("--rows", type=Path, required=True, help="an .npz of field columns")
38
+ fit.add_argument(
39
+ "--fields", required=True, help="comma-separated observation vector field names"
40
+ )
41
+ fit.add_argument("--out", type=Path, required=True, help="a new folder for the artifact")
42
+ fit.add_argument("--hidden", default="32,32", help="hidden layer widths, comma-separated")
43
+ fit.add_argument("--epochs", type=int, default=20)
44
+ fit.add_argument("--seed", type=int, default=20260924)
45
+ fit.set_defaults(handler=_fit_field_reference)
46
+ return parser
47
+
48
+
49
+ def _fit_field_reference(args: argparse.Namespace) -> int:
50
+ """Fit the play/wait model and write it as a field-model artifact."""
51
+ from .fit import FitConfig, fit_field_reference
52
+
53
+ fields = [name.strip() for name in args.fields.split(",") if name.strip()]
54
+ hidden = [int(width) for width in args.hidden.split(",") if width.strip()]
55
+ fit_field_reference(
56
+ args.rows,
57
+ fields,
58
+ args.out,
59
+ config=FitConfig(hidden=hidden, epochs=args.epochs, seed=args.seed),
60
+ )
61
+ return 0
62
+
63
+
64
+ def _artifact_digest(args: argparse.Namespace) -> int:
65
+ """The folder digest: over every file in the folder, by relative path."""
66
+ from .artifacts import artifact_digest
67
+
68
+ print(artifact_digest(args.folder))
69
+ return 0
70
+
71
+
72
+ def main(argv: Sequence[str] | None = None) -> int:
73
+ from royalelearn.extensions import PreflightError
74
+
75
+ args = build_parser().parse_args(argv)
76
+ try:
77
+ return int(args.handler(args))
78
+ except PreflightError as exc:
79
+ print(f"royaleimitate: {exc}", file=sys.stderr)
80
+ return 2
81
+
82
+
83
+ if __name__ == "__main__": # pragma: no cover
84
+ sys.exit(main())