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/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
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
@@ -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