cotterbot 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- cotter/__init__.py +125 -0
- cotter/backends.py +125 -0
- cotter/cli.py +221 -0
- cotter/compliance/__init__.py +33 -0
- cotter/compliance/base.py +74 -0
- cotter/config.py +201 -0
- cotter/envs/__init__.py +0 -0
- cotter/envs/factory.py +31 -0
- cotter/envs/registry.py +37 -0
- cotter/envs/wrapper.py +73 -0
- cotter/pipeline.py +214 -0
- cotter/policy.py +170 -0
- cotter/report.py +168 -0
- cotter/runner.py +282 -0
- cotter/stats.py +35 -0
- cotter/success.py +92 -0
- cotter/tests/__init__.py +0 -0
- cotter/tests/adversarial.py +398 -0
- cotter/tests/regression.py +166 -0
- cotter/tests/safety.py +155 -0
- cotter/tests/sprt.py +177 -0
- cotter/zoo/__init__.py +13 -0
- cotter/zoo/registry.py +173 -0
- cotterbot-0.1.0.dist-info/METADATA +409 -0
- cotterbot-0.1.0.dist-info/RECORD +28 -0
- cotterbot-0.1.0.dist-info/WHEEL +4 -0
- cotterbot-0.1.0.dist-info/entry_points.txt +3 -0
- cotterbot-0.1.0.dist-info/licenses/LICENSE +21 -0
cotter/config.py
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
"""YAML config schema for CLI-driven test runs.
|
|
2
|
+
|
|
3
|
+
A config file declares which test categories to run and their
|
|
4
|
+
parameters; categories are optional and only the ones present execute.
|
|
5
|
+
Example::
|
|
6
|
+
|
|
7
|
+
env: InvertedPendulum-v5
|
|
8
|
+
algo: PPO
|
|
9
|
+
base_seed: 0
|
|
10
|
+
success:
|
|
11
|
+
type: min_length
|
|
12
|
+
value: 1000
|
|
13
|
+
performance:
|
|
14
|
+
p0: 0.80
|
|
15
|
+
p1: 0.95
|
|
16
|
+
alpha: 0.05
|
|
17
|
+
beta: 0.05
|
|
18
|
+
n_max: 50
|
|
19
|
+
safety:
|
|
20
|
+
n_episodes: 20
|
|
21
|
+
limits:
|
|
22
|
+
cotter/joint_velocities: 5.0
|
|
23
|
+
cotter/actuator_forces: 2.5
|
|
24
|
+
regression:
|
|
25
|
+
baseline: artifacts/victim_ppo_inverted_pendulum.zip
|
|
26
|
+
n_pairs: 30
|
|
27
|
+
alpha: 0.05
|
|
28
|
+
adversarial:
|
|
29
|
+
epsilon: 0.07
|
|
30
|
+
n_episodes: 20
|
|
31
|
+
min_success_rate: 0.5
|
|
32
|
+
train: true
|
|
33
|
+
timesteps: 150000
|
|
34
|
+
report: report.json
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
from __future__ import annotations
|
|
38
|
+
|
|
39
|
+
from dataclasses import dataclass, field
|
|
40
|
+
from pathlib import Path
|
|
41
|
+
|
|
42
|
+
import yaml
|
|
43
|
+
|
|
44
|
+
from cotter.success import make_success_fn
|
|
45
|
+
from cotter.tests.safety import SafetyLimit
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
class ConfigError(ValueError):
|
|
49
|
+
"""The config file is invalid; the message says where and why."""
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
@dataclass
|
|
53
|
+
class PerformanceConfig:
|
|
54
|
+
p0: float = 0.80
|
|
55
|
+
p1: float = 0.95
|
|
56
|
+
alpha: float = 0.05
|
|
57
|
+
beta: float = 0.05
|
|
58
|
+
n_max: int = 50
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
@dataclass
|
|
62
|
+
class SafetyConfig:
|
|
63
|
+
limits: list[SafetyLimit] = field(default_factory=list)
|
|
64
|
+
n_episodes: int = 20
|
|
65
|
+
n_workers: int = 1 # parallel rollout workers; 1 = serial
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
@dataclass
|
|
69
|
+
class RegressionConfig:
|
|
70
|
+
baseline: Path = Path()
|
|
71
|
+
n_pairs: int = 30
|
|
72
|
+
alpha: float = 0.05
|
|
73
|
+
n_workers: int = 1 # parallel rollout workers; 1 = serial
|
|
74
|
+
metric: str = "return" # continuous metric for Wilcoxon: "return" or an info key
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
@dataclass
|
|
78
|
+
class AdversarialConfig:
|
|
79
|
+
epsilon: float = 0.05
|
|
80
|
+
n_episodes: int = 20
|
|
81
|
+
min_success_rate: float = 0.5
|
|
82
|
+
train: bool = True
|
|
83
|
+
timesteps: int = 150_000
|
|
84
|
+
use_zoo: bool = False # reuse/cache the trained adversary in the zoo
|
|
85
|
+
zoo_root: str | None = None # override the default ~/.cotter/zoo
|
|
86
|
+
n_workers: int = 1 # parallel rollout workers for the fixed-N eval; 1 = serial
|
|
87
|
+
max_seconds: float | None = None # wall-clock time-box on adversary training
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
@dataclass
|
|
91
|
+
class RunConfig:
|
|
92
|
+
env: str
|
|
93
|
+
success: dict
|
|
94
|
+
algo: str = "PPO"
|
|
95
|
+
base_seed: int = 0
|
|
96
|
+
backend: str = "gymnasium"
|
|
97
|
+
performance: PerformanceConfig | None = None
|
|
98
|
+
safety: SafetyConfig | None = None
|
|
99
|
+
regression: RegressionConfig | None = None
|
|
100
|
+
adversarial: AdversarialConfig | None = None
|
|
101
|
+
report: Path | None = None
|
|
102
|
+
|
|
103
|
+
def success_fn(self):
|
|
104
|
+
return make_success_fn(self.success)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
_KNOWN_TOP_KEYS = {
|
|
108
|
+
"env", "algo", "base_seed", "backend", "success",
|
|
109
|
+
"performance", "safety", "regression", "adversarial", "report",
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def _section(data: dict, name: str, cls, **transforms):
|
|
114
|
+
raw = data.get(name)
|
|
115
|
+
if raw is None:
|
|
116
|
+
return None
|
|
117
|
+
if not isinstance(raw, dict):
|
|
118
|
+
raise ConfigError(f"'{name}' must be a mapping; got {type(raw).__name__}")
|
|
119
|
+
valid = {f.name for f in cls.__dataclass_fields__.values()}
|
|
120
|
+
unknown = set(raw) - valid
|
|
121
|
+
if unknown:
|
|
122
|
+
raise ConfigError(f"unknown keys in '{name}': {sorted(unknown)} (expected {sorted(valid)})")
|
|
123
|
+
kwargs = dict(raw)
|
|
124
|
+
for key, fn in transforms.items():
|
|
125
|
+
if key in kwargs:
|
|
126
|
+
kwargs[key] = fn(kwargs[key])
|
|
127
|
+
try:
|
|
128
|
+
return cls(**kwargs)
|
|
129
|
+
except (TypeError, ValueError) as exc:
|
|
130
|
+
raise ConfigError(f"invalid '{name}' section: {exc}") from exc
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def _parse_limits(raw) -> list[SafetyLimit]:
|
|
134
|
+
if not isinstance(raw, dict) or not raw:
|
|
135
|
+
raise ConfigError("'safety.limits' must be a non-empty mapping of quantity -> max_abs")
|
|
136
|
+
limits = []
|
|
137
|
+
for quantity, max_abs in raw.items():
|
|
138
|
+
if not isinstance(max_abs, (int, float)) or isinstance(max_abs, bool):
|
|
139
|
+
raise ConfigError(f"safety limit '{quantity}' must be numeric; got {max_abs!r}")
|
|
140
|
+
try:
|
|
141
|
+
limits.append(SafetyLimit(str(quantity), float(max_abs)))
|
|
142
|
+
except ValueError as exc:
|
|
143
|
+
raise ConfigError(f"safety limit '{quantity}': {exc}") from exc
|
|
144
|
+
return limits
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def parse_config(data: dict, config_dir: Path | None = None) -> RunConfig:
|
|
148
|
+
"""Validate a parsed YAML mapping into a RunConfig.
|
|
149
|
+
|
|
150
|
+
Relative paths (regression baseline, report) resolve against
|
|
151
|
+
``config_dir`` when given, so configs work from any cwd.
|
|
152
|
+
"""
|
|
153
|
+
if not isinstance(data, dict):
|
|
154
|
+
raise ConfigError(f"config root must be a mapping; got {type(data).__name__}")
|
|
155
|
+
unknown = set(data) - _KNOWN_TOP_KEYS
|
|
156
|
+
if unknown:
|
|
157
|
+
raise ConfigError(f"unknown top-level keys: {sorted(unknown)} (expected {sorted(_KNOWN_TOP_KEYS)})")
|
|
158
|
+
if "env" not in data:
|
|
159
|
+
raise ConfigError("config must set 'env' (a Gymnasium env id)")
|
|
160
|
+
if "success" not in data:
|
|
161
|
+
raise ConfigError("config must set 'success' (a success criterion mapping)")
|
|
162
|
+
|
|
163
|
+
def resolve(p) -> Path:
|
|
164
|
+
path = Path(p)
|
|
165
|
+
if config_dir is not None and not path.is_absolute():
|
|
166
|
+
path = config_dir / path
|
|
167
|
+
return path
|
|
168
|
+
|
|
169
|
+
cfg = RunConfig(
|
|
170
|
+
env=str(data["env"]),
|
|
171
|
+
success=dict(data["success"]),
|
|
172
|
+
algo=str(data.get("algo", "PPO")),
|
|
173
|
+
base_seed=int(data.get("base_seed", 0)),
|
|
174
|
+
backend=str(data.get("backend", "gymnasium")),
|
|
175
|
+
performance=_section(data, "performance", PerformanceConfig),
|
|
176
|
+
safety=_section(data, "safety", SafetyConfig, limits=_parse_limits),
|
|
177
|
+
regression=_section(data, "regression", RegressionConfig, baseline=resolve),
|
|
178
|
+
adversarial=_section(data, "adversarial", AdversarialConfig),
|
|
179
|
+
report=resolve(data["report"]) if "report" in data else None,
|
|
180
|
+
)
|
|
181
|
+
if cfg.safety is not None and not cfg.safety.limits:
|
|
182
|
+
raise ConfigError("'safety' section requires a 'limits' mapping")
|
|
183
|
+
if cfg.regression is not None and cfg.regression.baseline == Path():
|
|
184
|
+
raise ConfigError("'regression' section requires a 'baseline' policy path")
|
|
185
|
+
make_success_fn(cfg.success) # fail at load time, not mid-run
|
|
186
|
+
return cfg
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
def load_config(path: str | Path) -> RunConfig:
|
|
190
|
+
"""Load and validate a YAML config file."""
|
|
191
|
+
path = Path(path)
|
|
192
|
+
if not path.exists():
|
|
193
|
+
raise FileNotFoundError(f"config file not found: {path}")
|
|
194
|
+
try:
|
|
195
|
+
data = yaml.safe_load(path.read_text())
|
|
196
|
+
except yaml.YAMLError as exc:
|
|
197
|
+
raise ConfigError(f"{path} is not valid YAML: {exc}") from exc
|
|
198
|
+
try:
|
|
199
|
+
return parse_config(data, config_dir=path.resolve().parent)
|
|
200
|
+
except ConfigError as exc:
|
|
201
|
+
raise ConfigError(f"{path}: {exc}") from exc
|
cotter/envs/__init__.py
ADDED
|
File without changes
|
cotter/envs/factory.py
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
"""Picklable environment factory for parallel rollouts.
|
|
2
|
+
|
|
3
|
+
``AsyncVectorEnv`` constructs its sub-environments inside worker
|
|
4
|
+
processes, so it needs a factory that survives pickling under the
|
|
5
|
+
``spawn`` start method used on macOS. A module-level callable class does;
|
|
6
|
+
a local closure does not. :class:`WrappedEnvFactory` rebuilds the same
|
|
7
|
+
instrumented env that :func:`cotter.pipeline.make_env` produces.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from dataclasses import dataclass
|
|
13
|
+
|
|
14
|
+
import gymnasium as gym
|
|
15
|
+
|
|
16
|
+
from cotter.envs.registry import make_env_by_id
|
|
17
|
+
from cotter.envs.wrapper import CotterWrapper
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True)
|
|
21
|
+
class WrappedEnvFactory:
|
|
22
|
+
"""Zero-argument callable that builds one wrapped env by id."""
|
|
23
|
+
|
|
24
|
+
env_id: str
|
|
25
|
+
|
|
26
|
+
def __call__(self) -> gym.Env:
|
|
27
|
+
env = make_env_by_id(self.env_id)
|
|
28
|
+
try:
|
|
29
|
+
return CotterWrapper(env)
|
|
30
|
+
except TypeError:
|
|
31
|
+
return env
|
cotter/envs/registry.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
"""Environment creation with optional extension registration.
|
|
2
|
+
|
|
3
|
+
gymnasium-robotics environments (Fetch, Shadow Hand, ...) are not
|
|
4
|
+
registered until the package is imported. :func:`make_env_by_id` retries
|
|
5
|
+
a failed ``gym.make`` after registering installed extension packages, so
|
|
6
|
+
config files can name any installed env without boilerplate.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import gymnasium as gym
|
|
12
|
+
from gymnasium.error import NameNotFound
|
|
13
|
+
|
|
14
|
+
_registered = False
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def register_extension_envs() -> None:
|
|
18
|
+
"""Register env packages that require an explicit import (idempotent)."""
|
|
19
|
+
global _registered
|
|
20
|
+
if _registered:
|
|
21
|
+
return
|
|
22
|
+
try:
|
|
23
|
+
import gymnasium_robotics
|
|
24
|
+
|
|
25
|
+
gym.register_envs(gymnasium_robotics)
|
|
26
|
+
except ImportError:
|
|
27
|
+
pass
|
|
28
|
+
_registered = True
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def make_env_by_id(env_id: str, **kwargs) -> gym.Env:
|
|
32
|
+
"""``gym.make`` that knows about installed extension packages."""
|
|
33
|
+
try:
|
|
34
|
+
return gym.make(env_id, **kwargs)
|
|
35
|
+
except NameNotFound:
|
|
36
|
+
register_extension_envs()
|
|
37
|
+
return gym.make(env_id, **kwargs)
|
cotter/envs/wrapper.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
"""Gymnasium wrapper that surfaces safety-relevant MuJoCo quantities.
|
|
2
|
+
|
|
3
|
+
:class:`CotterWrapper` guarantees that every ``step()`` info dict contains
|
|
4
|
+
the physical quantities the safety tests read, copied out of the MuJoCo
|
|
5
|
+
``data`` structure of the wrapped environment:
|
|
6
|
+
|
|
7
|
+
============================== =========================================
|
|
8
|
+
info key contents
|
|
9
|
+
============================== =========================================
|
|
10
|
+
``cotter/joint_velocities`` ``data.qvel`` — generalized joint
|
|
11
|
+
velocities (rad/s or m/s per DoF)
|
|
12
|
+
``cotter/actuator_forces`` ``data.actuator_force`` — scalar force/
|
|
13
|
+
torque produced by each actuator (N or Nm)
|
|
14
|
+
``cotter/contact_count`` ``data.ncon`` — number of active contact
|
|
15
|
+
points this step
|
|
16
|
+
``cotter/contact_forces`` ``data.cfrc_ext`` L2 norms — magnitude of
|
|
17
|
+
the external contact wrench on each body
|
|
18
|
+
(one value per body; N-scale)
|
|
19
|
+
============================== =========================================
|
|
20
|
+
|
|
21
|
+
The wrapper works with any Gymnasium MuJoCo environment (anything whose
|
|
22
|
+
unwrapped env exposes a ``data`` attribute with these fields) and fails
|
|
23
|
+
loudly at construction time otherwise.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
from __future__ import annotations
|
|
27
|
+
|
|
28
|
+
import gymnasium as gym
|
|
29
|
+
import numpy as np
|
|
30
|
+
|
|
31
|
+
JOINT_VELOCITIES = "cotter/joint_velocities"
|
|
32
|
+
ACTUATOR_FORCES = "cotter/actuator_forces"
|
|
33
|
+
CONTACT_COUNT = "cotter/contact_count"
|
|
34
|
+
CONTACT_FORCES = "cotter/contact_forces"
|
|
35
|
+
|
|
36
|
+
INSTRUMENTED_KEYS = (JOINT_VELOCITIES, ACTUATOR_FORCES, CONTACT_COUNT, CONTACT_FORCES)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class CotterWrapper(gym.Wrapper):
|
|
40
|
+
"""Instrument a MuJoCo env's info dict with safety-relevant quantities."""
|
|
41
|
+
|
|
42
|
+
def __init__(self, env: gym.Env) -> None:
|
|
43
|
+
super().__init__(env)
|
|
44
|
+
data = getattr(env.unwrapped, "data", None)
|
|
45
|
+
if data is None or not all(
|
|
46
|
+
hasattr(data, attr)
|
|
47
|
+
for attr in ("qvel", "actuator_force", "ncon", "cfrc_ext")
|
|
48
|
+
):
|
|
49
|
+
raise TypeError(
|
|
50
|
+
f"CotterWrapper requires a MuJoCo-backed environment exposing "
|
|
51
|
+
f"`unwrapped.data` with qvel/actuator_force/ncon; "
|
|
52
|
+
f"got {type(env.unwrapped).__name__}. Wrap a Gymnasium MuJoCo "
|
|
53
|
+
"env (e.g. gym.make('InvertedPendulum-v5'))."
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
def _instrument(self, info: dict) -> dict:
|
|
57
|
+
data = self.env.unwrapped.data
|
|
58
|
+
info[JOINT_VELOCITIES] = np.array(data.qvel, dtype=float, copy=True)
|
|
59
|
+
info[ACTUATOR_FORCES] = np.array(data.actuator_force, dtype=float, copy=True)
|
|
60
|
+
info[CONTACT_COUNT] = int(data.ncon)
|
|
61
|
+
# one L2 wrench magnitude per body; body 0 is the world and always 0
|
|
62
|
+
info[CONTACT_FORCES] = np.linalg.norm(
|
|
63
|
+
np.asarray(data.cfrc_ext, dtype=float), axis=1
|
|
64
|
+
)
|
|
65
|
+
return info
|
|
66
|
+
|
|
67
|
+
def reset(self, **kwargs):
|
|
68
|
+
obs, info = self.env.reset(**kwargs)
|
|
69
|
+
return obs, self._instrument(info)
|
|
70
|
+
|
|
71
|
+
def step(self, action):
|
|
72
|
+
obs, reward, terminated, truncated, info = self.env.step(action)
|
|
73
|
+
return obs, reward, terminated, truncated, self._instrument(info)
|
cotter/pipeline.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
1
|
+
"""Execute a RunConfig against a policy: the engine behind `cotter run`.
|
|
2
|
+
|
|
3
|
+
Runs whichever test categories the config declares and returns the
|
|
4
|
+
aggregated :class:`~cotter.report.TestReport`. Categories run in a fixed
|
|
5
|
+
order (performance, safety, regression, adversarial) with seeds derived
|
|
6
|
+
from ``base_seed`` per category, so adding or removing one category does
|
|
7
|
+
not change another's rollouts.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import importlib
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Callable
|
|
15
|
+
|
|
16
|
+
import gymnasium as gym
|
|
17
|
+
|
|
18
|
+
from cotter.backends import BackendFactory
|
|
19
|
+
from cotter.config import RunConfig
|
|
20
|
+
from cotter.envs.factory import WrappedEnvFactory
|
|
21
|
+
from cotter.envs.wrapper import CotterWrapper
|
|
22
|
+
from cotter.policy import Policy, load_policy
|
|
23
|
+
from cotter.report import TestReport
|
|
24
|
+
from cotter.runner import (
|
|
25
|
+
make_seed_sequence,
|
|
26
|
+
rollout_one,
|
|
27
|
+
run_rollouts,
|
|
28
|
+
run_rollouts_parallel,
|
|
29
|
+
)
|
|
30
|
+
from cotter.tests.adversarial import get_adversary, run_adversarial_test
|
|
31
|
+
from cotter.tests.regression import mcnemar_exact, wilcoxon_regression
|
|
32
|
+
from cotter.tests.safety import evaluate_safety
|
|
33
|
+
from cotter.tests.sprt import run_sprt
|
|
34
|
+
|
|
35
|
+
# per-category seed offsets, kept stable so results are comparable across
|
|
36
|
+
# configs that enable different category subsets
|
|
37
|
+
_PERF_SEED, _SAFETY_SEED, _REGRESSION_SEED, _ADV_SEED = 100, 200, 300, 400
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def resolve_algo(name: str):
|
|
41
|
+
"""Look up an SB3 algorithm class by name (PPO, SAC, TD3, ...)."""
|
|
42
|
+
sb3 = importlib.import_module("stable_baselines3")
|
|
43
|
+
algo = getattr(sb3, name, None)
|
|
44
|
+
if algo is None:
|
|
45
|
+
available = [a for a in ("A2C", "DDPG", "DQN", "PPO", "SAC", "TD3") if hasattr(sb3, a)]
|
|
46
|
+
raise ValueError(f"unknown SB3 algorithm '{name}' (available: {available})")
|
|
47
|
+
return algo
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def make_env(env_id: str, backend: str = "gymnasium", log: Callable[[str], None] = print) -> gym.Env:
|
|
51
|
+
"""Create the test environment via the named backend.
|
|
52
|
+
|
|
53
|
+
The env is instrumented with :class:`CotterWrapper` when the backend
|
|
54
|
+
produces a MuJoCo-backed env; otherwise safety quantities are
|
|
55
|
+
unavailable and a note is logged.
|
|
56
|
+
"""
|
|
57
|
+
env = BackendFactory.from_name(backend).make_env(env_id)
|
|
58
|
+
if not isinstance(env, CotterWrapper):
|
|
59
|
+
log(f"[cotter] {env_id} is not MuJoCo-backed; safety quantities unavailable")
|
|
60
|
+
return env
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _obtain_adversary(policy_path, policy, env, cfg, adv_cfg, log):
|
|
64
|
+
"""Load a cached adversary from the zoo, or train one and cache it.
|
|
65
|
+
|
|
66
|
+
Falls back to :func:`get_adversary` (train-with-random-fallback) when
|
|
67
|
+
the zoo is disabled or the victim cannot be hashed.
|
|
68
|
+
"""
|
|
69
|
+
if not adv_cfg.use_zoo:
|
|
70
|
+
return get_adversary(
|
|
71
|
+
policy, env, adv_cfg.epsilon, timesteps=adv_cfg.timesteps,
|
|
72
|
+
seed=cfg.base_seed, log=log, max_seconds=adv_cfg.max_seconds,
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
from cotter.tests.adversarial import train_adversary
|
|
76
|
+
from cotter.zoo import AdversaryZoo
|
|
77
|
+
|
|
78
|
+
zoo = AdversaryZoo(adv_cfg.zoo_root) if adv_cfg.zoo_root else AdversaryZoo()
|
|
79
|
+
# the artifact path hashes identically however it is later referenced
|
|
80
|
+
victim_ref = Path(policy_path)
|
|
81
|
+
|
|
82
|
+
cached = zoo.load(cfg.env, victim_ref, adv_cfg.epsilon)
|
|
83
|
+
if cached is not None:
|
|
84
|
+
log(f"[cotter] reusing cached adversary from zoo ({zoo.root})")
|
|
85
|
+
return cached, "reused cached zoo adversary"
|
|
86
|
+
|
|
87
|
+
log(f"[cotter] no cached adversary; training and storing in zoo ({zoo.root})")
|
|
88
|
+
adversary = train_adversary(
|
|
89
|
+
policy, env, adv_cfg.epsilon, timesteps=adv_cfg.timesteps,
|
|
90
|
+
seed=cfg.base_seed, max_seconds=adv_cfg.max_seconds,
|
|
91
|
+
)
|
|
92
|
+
entry = zoo.save(adversary, cfg.env, victim_ref, notes=f"trained via cotter run, {cfg.env}")
|
|
93
|
+
log(f"[cotter] stored adversary at {entry.path}")
|
|
94
|
+
return adversary, "trained PPO adversary (cached in zoo)"
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def run_from_config(
|
|
98
|
+
policy_path: str | Path,
|
|
99
|
+
cfg: RunConfig,
|
|
100
|
+
log: Callable[[str], None] = print,
|
|
101
|
+
) -> TestReport:
|
|
102
|
+
env = make_env(cfg.env, backend=cfg.backend, log=log)
|
|
103
|
+
if cfg.safety is not None and not isinstance(env, CotterWrapper):
|
|
104
|
+
raise ValueError(
|
|
105
|
+
f"config declares safety limits but {cfg.env} is not MuJoCo-backed, "
|
|
106
|
+
"so the instrumented quantities do not exist"
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
algo = resolve_algo(cfg.algo)
|
|
110
|
+
policy = load_policy(Path(policy_path), env, algo=algo)
|
|
111
|
+
success_fn = cfg.success_fn()
|
|
112
|
+
log(f"[cotter] loaded policy '{policy.name}' onto {cfg.env}")
|
|
113
|
+
|
|
114
|
+
report = TestReport(
|
|
115
|
+
policy_name=policy.name,
|
|
116
|
+
env_id=cfg.env,
|
|
117
|
+
metadata={"base_seed": cfg.base_seed, "success": cfg.success, "algo": cfg.algo},
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
if cfg.performance is not None:
|
|
121
|
+
p = cfg.performance
|
|
122
|
+
log(f"[cotter] performance: SPRT p0={p.p0} p1={p.p1} n_max={p.n_max}")
|
|
123
|
+
seeds = make_seed_sequence(p.n_max, cfg.base_seed + _PERF_SEED)
|
|
124
|
+
|
|
125
|
+
def trial(i: int) -> bool:
|
|
126
|
+
rec = rollout_one(policy, env, seeds[i], success_fn, record_infos=False)
|
|
127
|
+
log(f"[cotter] trial {i + 1}: length={rec.length} "
|
|
128
|
+
f"return={rec.total_reward:.1f} success={rec.success}")
|
|
129
|
+
return rec.success
|
|
130
|
+
|
|
131
|
+
result = run_sprt(trial, p0=p.p0, p1=p.p1, alpha=p.alpha, beta=p.beta, n_max=p.n_max)
|
|
132
|
+
report.add_sprt(result)
|
|
133
|
+
log(f"[cotter] => {result.decision.value} after {result.n_trials} trials")
|
|
134
|
+
|
|
135
|
+
factory = WrappedEnvFactory(cfg.env)
|
|
136
|
+
|
|
137
|
+
def dispatch(pol, n, *, seeds=None, base=0, record_infos, n_workers):
|
|
138
|
+
if n_workers > 1:
|
|
139
|
+
return run_rollouts_parallel(
|
|
140
|
+
pol, factory, n, success_fn, seeds=seeds, base_seed=base,
|
|
141
|
+
record_infos=record_infos, n_workers=n_workers,
|
|
142
|
+
)
|
|
143
|
+
return run_rollouts(
|
|
144
|
+
pol, env, n, success_fn, seeds=seeds, base_seed=base,
|
|
145
|
+
record_infos=record_infos,
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
if cfg.safety is not None:
|
|
149
|
+
s = cfg.safety
|
|
150
|
+
log(f"[cotter] safety: {len(s.limits)} limit(s) over {s.n_episodes} episodes"
|
|
151
|
+
+ (f" ({s.n_workers} workers)" if s.n_workers > 1 else ""))
|
|
152
|
+
rollouts = dispatch(
|
|
153
|
+
policy, s.n_episodes, base=cfg.base_seed + _SAFETY_SEED,
|
|
154
|
+
record_infos=True, n_workers=s.n_workers,
|
|
155
|
+
)
|
|
156
|
+
result = evaluate_safety(rollouts.episode_infos, s.limits)
|
|
157
|
+
report.add_safety(result)
|
|
158
|
+
for quantity, worst in result.worst_observed.items():
|
|
159
|
+
limit = next(l.max_abs for l in s.limits if l.quantity == quantity)
|
|
160
|
+
log(f"[cotter] worst |{quantity}| = {worst:.4f} (limit {limit})")
|
|
161
|
+
log(f"[cotter] => {result.decision.value}")
|
|
162
|
+
|
|
163
|
+
if cfg.regression is not None:
|
|
164
|
+
r = cfg.regression
|
|
165
|
+
log(f"[cotter] regression: vs baseline {r.baseline} on {r.n_pairs} paired seeds")
|
|
166
|
+
baseline = load_policy(r.baseline, env, algo=algo, name=f"baseline:{r.baseline.stem}")
|
|
167
|
+
seeds = make_seed_sequence(r.n_pairs, cfg.base_seed + _REGRESSION_SEED)
|
|
168
|
+
# info-based metrics need per-episode final info recorded
|
|
169
|
+
record = r.metric != "return"
|
|
170
|
+
base_rs = dispatch(baseline, r.n_pairs, seeds=seeds, record_infos=record, n_workers=r.n_workers)
|
|
171
|
+
cand_rs = dispatch(policy, r.n_pairs, seeds=seeds, record_infos=record, n_workers=r.n_workers)
|
|
172
|
+
# baseline vs candidate: the CLI policy is the candidate
|
|
173
|
+
mcnemar = mcnemar_exact(base_rs.successes, cand_rs.successes, alpha=r.alpha)
|
|
174
|
+
report.add_regression(mcnemar, name="success_mcnemar")
|
|
175
|
+
wilcoxon = wilcoxon_regression(
|
|
176
|
+
base_rs.metric_values(r.metric), cand_rs.metric_values(r.metric), alpha=r.alpha
|
|
177
|
+
)
|
|
178
|
+
report.add_regression(wilcoxon, name=f"{r.metric}_wilcoxon")
|
|
179
|
+
log(f"[cotter] baseline {base_rs.success_rate:.0%} vs candidate "
|
|
180
|
+
f"{cand_rs.success_rate:.0%} (metric: {r.metric})")
|
|
181
|
+
log(f"[cotter] => McNemar {mcnemar.decision.value} (p={mcnemar.p_value:.3g}), "
|
|
182
|
+
f"Wilcoxon {wilcoxon.decision.value} (p={wilcoxon.p_value:.3g})")
|
|
183
|
+
|
|
184
|
+
if cfg.adversarial is not None:
|
|
185
|
+
a = cfg.adversarial
|
|
186
|
+
log(f"[cotter] adversarial: eps={a.epsilon} over {a.n_episodes} episodes")
|
|
187
|
+
random_result = run_adversarial_test(
|
|
188
|
+
policy, env, success_fn, epsilon=a.epsilon, n_episodes=a.n_episodes,
|
|
189
|
+
min_success_rate=a.min_success_rate, base_seed=cfg.base_seed + _ADV_SEED,
|
|
190
|
+
notes="uniform random baseline",
|
|
191
|
+
n_workers=a.n_workers, env_factory=factory,
|
|
192
|
+
)
|
|
193
|
+
report.add_adversarial(random_result, name="random_baseline")
|
|
194
|
+
log(f"[cotter] random baseline: {random_result.clean_success_rate:.0%} clean -> "
|
|
195
|
+
f"{random_result.adversarial_success_rate:.0%} perturbed")
|
|
196
|
+
if a.train:
|
|
197
|
+
adversary, notes = _obtain_adversary(policy_path, policy, env, cfg, a, log)
|
|
198
|
+
learned_result = run_adversarial_test(
|
|
199
|
+
policy, env, success_fn, epsilon=a.epsilon, n_episodes=a.n_episodes,
|
|
200
|
+
adversary=adversary, min_success_rate=a.min_success_rate,
|
|
201
|
+
base_seed=cfg.base_seed + _ADV_SEED,
|
|
202
|
+
notes=notes or "trained PPO adversary",
|
|
203
|
+
n_workers=a.n_workers, env_factory=factory,
|
|
204
|
+
)
|
|
205
|
+
report.add_adversarial(learned_result, name=f"learned_{adversary.name}")
|
|
206
|
+
log(f"[cotter] {adversary.name} adversary: "
|
|
207
|
+
f"{learned_result.clean_success_rate:.0%} clean -> "
|
|
208
|
+
f"{learned_result.adversarial_success_rate:.0%} perturbed")
|
|
209
|
+
|
|
210
|
+
if cfg.report is not None:
|
|
211
|
+
path = report.to_json(cfg.report)
|
|
212
|
+
log(f"[cotter] JSON report written to {path}")
|
|
213
|
+
|
|
214
|
+
return report
|