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.
- royaleimitate-0.2.6/LICENSE +21 -0
- royaleimitate-0.2.6/PKG-INFO +22 -0
- royaleimitate-0.2.6/README.md +41 -0
- royaleimitate-0.2.6/pyproject.toml +74 -0
- royaleimitate-0.2.6/royaleimitate/__init__.py +55 -0
- royaleimitate-0.2.6/royaleimitate/__main__.py +6 -0
- royaleimitate-0.2.6/royaleimitate/alarms.py +36 -0
- royaleimitate-0.2.6/royaleimitate/artifacts.py +310 -0
- royaleimitate-0.2.6/royaleimitate/cli.py +84 -0
- royaleimitate-0.2.6/royaleimitate/cloning.py +226 -0
- royaleimitate-0.2.6/royaleimitate/config.py +240 -0
- royaleimitate-0.2.6/royaleimitate/extension.py +105 -0
- royaleimitate-0.2.6/royaleimitate/field_model.py +98 -0
- royaleimitate-0.2.6/royaleimitate/fit.py +170 -0
- royaleimitate-0.2.6/royaleimitate/init.py +111 -0
- royaleimitate-0.2.6/royaleimitate/public_log.py +417 -0
- royaleimitate-0.2.6/royaleimitate/references.py +184 -0
- royaleimitate-0.2.6/royaleimitate/regularisers.py +435 -0
- royaleimitate-0.2.6/royaleimitate/replay_cards.json +121 -0
- royaleimitate-0.2.6/royaleimitate/replays.py +324 -0
- royaleimitate-0.2.6/royaleimitate/schema.py +93 -0
- royaleimitate-0.2.6/royaleimitate/shards.py +591 -0
- royaleimitate-0.2.6/royaleimitate/split.py +36 -0
- royaleimitate-0.2.6/royaleimitate/warm_start.py +141 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/PKG-INFO +22 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/SOURCES.txt +41 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/dependency_links.txt +1 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/entry_points.txt +6 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/requires.txt +17 -0
- royaleimitate-0.2.6/royaleimitate.egg-info/top_level.txt +1 -0
- royaleimitate-0.2.6/setup.cfg +4 -0
- royaleimitate-0.2.6/tests/test_clone.py +126 -0
- royaleimitate-0.2.6/tests/test_extension_contract.py +68 -0
- royaleimitate-0.2.6/tests/test_field_reference.py +176 -0
- royaleimitate-0.2.6/tests/test_imitation_alarms.py +140 -0
- royaleimitate-0.2.6/tests/test_imitation_anchor.py +711 -0
- royaleimitate-0.2.6/tests/test_imitation_init.py +560 -0
- royaleimitate-0.2.6/tests/test_public_log.py +767 -0
- royaleimitate-0.2.6/tests/test_replays.py +167 -0
- royaleimitate-0.2.6/tests/test_save_actor.py +91 -0
- royaleimitate-0.2.6/tests/test_shards.py +254 -0
- royaleimitate-0.2.6/tests/test_supported_surface.py +60 -0
- 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,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())
|