neorx 0.2.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.
Files changed (123) hide show
  1. neorx/__init__.py +118 -0
  2. neorx/causalbiorl/__init__.py +23 -0
  3. neorx/causalbiorl/__main__.py +193 -0
  4. neorx/causalbiorl/agents/__init__.py +6 -0
  5. neorx/causalbiorl/agents/baseline_agent.py +161 -0
  6. neorx/causalbiorl/agents/causal_agent.py +434 -0
  7. neorx/causalbiorl/benchmark.py +276 -0
  8. neorx/causalbiorl/causal/__init__.py +40 -0
  9. neorx/causalbiorl/causal/discovery.py +237 -0
  10. neorx/causalbiorl/causal/graph_encoder.py +555 -0
  11. neorx/causalbiorl/causal/planner.py +333 -0
  12. neorx/causalbiorl/causal/reward_learner.py +384 -0
  13. neorx/causalbiorl/causal/scm.py +373 -0
  14. neorx/causalbiorl/causal/surrogate_docker.py +441 -0
  15. neorx/causalbiorl/envs/__init__.py +11 -0
  16. neorx/causalbiorl/envs/cell_growth.py +289 -0
  17. neorx/causalbiorl/envs/drug_discovery.py +884 -0
  18. neorx/causalbiorl/envs/metabolic_pathway.py +303 -0
  19. neorx/causalbiorl/envs/registration.py +41 -0
  20. neorx/causalbiorl/envs/toggle_switch.py +269 -0
  21. neorx/causalbiorl/models.py +94 -0
  22. neorx/causalbiorl/viz.py +207 -0
  23. neorx/cli/__init__.py +46 -0
  24. neorx/cli/__main__.py +6 -0
  25. neorx/core/__init__.py +116 -0
  26. neorx/core/__main__.py +275 -0
  27. neorx/core/api.py +214 -0
  28. neorx/core/bio/__init__.py +0 -0
  29. neorx/core/bio/classifier.py +462 -0
  30. neorx/core/bio/tissue_filter.py +572 -0
  31. neorx/core/cache.py +173 -0
  32. neorx/core/causal/__init__.py +0 -0
  33. neorx/core/causal/counterfactual.py +312 -0
  34. neorx/core/causal/identifier.py +1291 -0
  35. neorx/core/graph/__init__.py +0 -0
  36. neorx/core/graph/graph_builder.py +446 -0
  37. neorx/core/graph/models.py +359 -0
  38. neorx/core/graph/persistence.py +349 -0
  39. neorx/core/literature_validator.py +281 -0
  40. neorx/core/pipeline.py +920 -0
  41. neorx/core/py.typed +1 -0
  42. neorx/core/report.py +450 -0
  43. neorx/core/scoring/__init__.py +0 -0
  44. neorx/core/scoring/admet.py +273 -0
  45. neorx/core/scoring/scorer.py +260 -0
  46. neorx/core/sources/__init__.py +42 -0
  47. neorx/core/sources/chembl.py +643 -0
  48. neorx/core/sources/kegg.py +228 -0
  49. neorx/core/sources/monarch.py +355 -0
  50. neorx/core/sources/open_targets.py +301 -0
  51. neorx/core/sources/pdb.py +196 -0
  52. neorx/core/sources/reactome.py +180 -0
  53. neorx/core/sources/string_db.py +198 -0
  54. neorx/core/sources/uniprot.py +208 -0
  55. neorx/core/templates/report.html +440 -0
  56. neorx/core/validator.py +412 -0
  57. neorx/dockbot/__init__.py +79 -0
  58. neorx/dockbot/__main__.py +315 -0
  59. neorx/dockbot/api.py +292 -0
  60. neorx/dockbot/binding_site.py +350 -0
  61. neorx/dockbot/docker.py +335 -0
  62. neorx/dockbot/ligand_prep.py +306 -0
  63. neorx/dockbot/models.py +209 -0
  64. neorx/dockbot/parallel.py +248 -0
  65. neorx/dockbot/protein_prep.py +424 -0
  66. neorx/dockbot/report.py +361 -0
  67. neorx/dockbot/scorer.py +238 -0
  68. neorx/dockbot/viz.py +278 -0
  69. neorx/experiments/__init__.py +24 -0
  70. neorx/experiments/__main__.py +100 -0
  71. neorx/experiments/capture.py +85 -0
  72. neorx/experiments/chembl.py +74 -0
  73. neorx/experiments/gates.py +314 -0
  74. neorx/experiments/record.py +204 -0
  75. neorx/experiments/registry.py +105 -0
  76. neorx/experiments/replay.py +129 -0
  77. neorx/genmol/__init__.py +60 -0
  78. neorx/genmol/__main__.py +335 -0
  79. neorx/genmol/api.py +235 -0
  80. neorx/genmol/assets/__init__.py +0 -0
  81. neorx/genmol/assets/molvae_chembl36.pt +0 -0
  82. neorx/genmol/assets/tokenizer.json +39 -0
  83. neorx/genmol/configs/default.yaml +58 -0
  84. neorx/genmol/data/__init__.py +23 -0
  85. neorx/genmol/data/dataset.py +204 -0
  86. neorx/genmol/data/download.py +239 -0
  87. neorx/genmol/data/preprocess.py +262 -0
  88. neorx/genmol/data/tokenizer.py +318 -0
  89. neorx/genmol/evaluation/__init__.py +44 -0
  90. neorx/genmol/evaluation/distribution.py +188 -0
  91. neorx/genmol/evaluation/metrics.py +238 -0
  92. neorx/genmol/evaluation/visualise.py +392 -0
  93. neorx/genmol/generate.py +371 -0
  94. neorx/genmol/models/__init__.py +17 -0
  95. neorx/genmol/models/cvae.py +493 -0
  96. neorx/genmol/models/vae.py +538 -0
  97. neorx/genmol/pretrained.py +61 -0
  98. neorx/genmol/train.py +387 -0
  99. neorx/mirrorfold/__init__.py +121 -0
  100. neorx/mirrorfold/__main__.py +306 -0
  101. neorx/mirrorfold/analysis.py +405 -0
  102. neorx/mirrorfold/api.py +225 -0
  103. neorx/mirrorfold/compare.py +520 -0
  104. neorx/mirrorfold/mirror.py +370 -0
  105. neorx/mirrorfold/models.py +219 -0
  106. neorx/mirrorfold/predictor.py +503 -0
  107. neorx/mirrorfold/therapeutic.py +381 -0
  108. neorx/mirrorfold/viz.py +559 -0
  109. neorx/molscreen/__init__.py +64 -0
  110. neorx/molscreen/__main__.py +87 -0
  111. neorx/molscreen/accessibility.py +131 -0
  112. neorx/molscreen/data/.gitignore +2 -0
  113. neorx/molscreen/filters.py +350 -0
  114. neorx/molscreen/models.py +285 -0
  115. neorx/molscreen/parser.py +204 -0
  116. neorx/molscreen/properties.py +211 -0
  117. neorx/molscreen/similarity.py +384 -0
  118. neorx/py.typed +0 -0
  119. neorx-0.2.0.dist-info/METADATA +497 -0
  120. neorx-0.2.0.dist-info/RECORD +123 -0
  121. neorx-0.2.0.dist-info/WHEEL +4 -0
  122. neorx-0.2.0.dist-info/entry_points.txt +6 -0
  123. neorx-0.2.0.dist-info/licenses/LICENSE +31 -0
neorx/__init__.py ADDED
@@ -0,0 +1,118 @@
1
+ """
2
+ NeoRx — Public API
3
+ ===========================
4
+
5
+ Thin re-export package so users can write clean imports::
6
+
7
+ from neorx import run_pipeline, build_disease_graph
8
+ from neorx import identify_causal_targets
9
+ from neorx import predict_admet, save_graph
10
+
11
+ Instead of the internal layout::
12
+
13
+ from neorx.core import run_pipeline # also works
14
+ """
15
+
16
+ from neorx.core import ( # noqa: F401
17
+ # Enums
18
+ NodeType,
19
+ EdgeType,
20
+ JobStatus,
21
+ TargetClassification,
22
+ # Graph models
23
+ GraphNode,
24
+ GraphEdge,
25
+ DiseaseGraph,
26
+ # Analysis models
27
+ NeoRxResult,
28
+ ScoredCandidate,
29
+ # Pipeline models
30
+ PipelineJob,
31
+ PipelineResult,
32
+ # API models
33
+ RunRequest,
34
+ GraphRequest,
35
+ IdentifyRequest,
36
+ ScreenTargetRequest,
37
+ StatusResponse,
38
+ # Core functions
39
+ build_disease_graph,
40
+ disease_graph_to_networkx,
41
+ identify_causal_targets,
42
+ score_candidate,
43
+ rank_candidates,
44
+ normalise_affinity,
45
+ normalise_sa,
46
+ run_pipeline,
47
+ generate_report,
48
+ # Cache
49
+ get_cache,
50
+ cached_api_call,
51
+ store_api_response,
52
+ # Persistence
53
+ save_graph,
54
+ load_graph,
55
+ list_saved_graphs,
56
+ # ADMET
57
+ predict_admet,
58
+ ADMETProfile,
59
+ )
60
+
61
+ # --- Subpackages -----------------------------------------------------
62
+ from neorx import causalbiorl, core, dockbot, genmol, mirrorfold, molscreen # noqa: F401,E402
63
+
64
+ __version__ = "0.2.0"
65
+
66
+ __all__ = [
67
+ # Enums
68
+ "NodeType",
69
+ "EdgeType",
70
+ "JobStatus",
71
+ "TargetClassification",
72
+ # Graph models
73
+ "GraphNode",
74
+ "GraphEdge",
75
+ "DiseaseGraph",
76
+ # Analysis models
77
+ "NeoRxResult",
78
+ "ScoredCandidate",
79
+ # Pipeline models
80
+ "PipelineJob",
81
+ "PipelineResult",
82
+ # API models
83
+ "RunRequest",
84
+ "GraphRequest",
85
+ "IdentifyRequest",
86
+ "ScreenTargetRequest",
87
+ "StatusResponse",
88
+ # Core functions
89
+ "build_disease_graph",
90
+ "disease_graph_to_networkx",
91
+ "identify_causal_targets",
92
+ "score_candidate",
93
+ "rank_candidates",
94
+ "normalise_affinity",
95
+ "normalise_sa",
96
+ "run_pipeline",
97
+ "generate_report",
98
+ # Cache
99
+ "get_cache",
100
+ "cached_api_call",
101
+ "store_api_response",
102
+ # Persistence
103
+ "save_graph",
104
+ "load_graph",
105
+ "list_saved_graphs",
106
+ # ADMET
107
+ "predict_admet",
108
+ "ADMETProfile",
109
+ ]
110
+
111
+ __all__ = __all__ + [
112
+ "core",
113
+ "genmol",
114
+ "causalbiorl",
115
+ "molscreen",
116
+ "dockbot",
117
+ "mirrorfold",
118
+ ]
@@ -0,0 +1,23 @@
1
+ """
2
+ CausalBioRL — Causal Reinforcement Learning Environments for Biological System Control.
3
+
4
+ Provides Gymnasium-compatible RL environments that simulate biological systems
5
+ (gene expression, metabolic flux, cell populations) and causal RL agents that
6
+ use structural causal models as world models.
7
+
8
+ Reference:
9
+ CausalBioRL: Causal Reinforcement Learning Environments for
10
+ Biological System Control (2026).
11
+ """
12
+
13
+ from neorx.causalbiorl.envs.registration import register_envs
14
+
15
+ from . import agents, causal, envs # noqa: F401,E402
16
+ from .agents.causal_agent import CausalAgent # noqa: F401,E402
17
+ from .envs.drug_discovery import DrugDiscoveryEnv # noqa: F401,E402
18
+
19
+ __version__ = "0.1.0"
20
+ __all__ = ["envs", "agents", "causal", "DrugDiscoveryEnv", "CausalAgent"]
21
+
22
+ # Register Gymnasium environments on import
23
+ register_envs()
@@ -0,0 +1,193 @@
1
+ """
2
+ CausalBioRL command-line interface.
3
+
4
+ Usage:
5
+ python -m causalbiorl train --env GeneticToggle-v0 --agent causal --episodes 1000
6
+ python -m causalbiorl benchmark --envs all --agents all --seeds 10
7
+ python -m causalbiorl play --env GeneticToggle-v0
8
+ python -m causalbiorl visualise --env GeneticToggle-v0 --agent causal
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from pathlib import Path
14
+ from typing import Annotated, Optional
15
+
16
+ import gymnasium as gym
17
+ import numpy as np
18
+ import typer
19
+
20
+ import neorx.causalbiorl # register envs # noqa: F401
21
+
22
+ app = typer.Typer(
23
+ name="causalbiorl",
24
+ help="CausalBioRL — Causal RL environments for biological system control.",
25
+ add_completion=False,
26
+ )
27
+
28
+
29
+ # ────────────────────────────────────────────────────────────────────────── #
30
+ # train #
31
+ # ────────────────────────────────────────────────────────────────────────── #
32
+
33
+
34
+ @app.command()
35
+ def train(
36
+ env: Annotated[str, typer.Option(help="Gymnasium environment ID")] = "GeneticToggle-v0",
37
+ agent: Annotated[str, typer.Option(help="Agent type: causal, ppo, sac, random")] = "causal",
38
+ episodes: Annotated[int, typer.Option(help="Number of training episodes")] = 500,
39
+ difficulty: Annotated[str, typer.Option(help="easy / medium / hard")] = "medium",
40
+ seed: Annotated[Optional[int], typer.Option(help="Random seed")] = None,
41
+ output_dir: Annotated[str, typer.Option(help="Output directory")] = "results",
42
+ ) -> None:
43
+ """Train an agent on a CausalBioRL environment."""
44
+ from neorx.causalbiorl.benchmark import _make_agent
45
+
46
+ typer.echo(f"🧬 Training {agent} on {env} ({difficulty}) for {episodes} episodes …")
47
+ environment = gym.make(env, difficulty=difficulty)
48
+ ag = _make_agent(agent, environment, seed=seed or 0)
49
+ metrics = ag.train(n_episodes=episodes, verbose=True)
50
+ environment.close()
51
+
52
+ out = Path(output_dir)
53
+ out.mkdir(parents=True, exist_ok=True)
54
+
55
+ # Save reward curve
56
+ import json
57
+ with open(out / f"{env}_{agent}_rewards.json", "w") as f:
58
+ json.dump(metrics["episode_rewards"], f)
59
+
60
+ final_mean = float(np.mean(metrics["episode_rewards"][-50:]))
61
+ typer.echo(f"✅ Done — final 50-episode mean reward: {final_mean:.3f}")
62
+
63
+
64
+ # ────────────────────────────────────────────────────────────────────────── #
65
+ # benchmark #
66
+ # ────────────────────────────────────────────────────────────────────────── #
67
+
68
+
69
+ @app.command()
70
+ def benchmark(
71
+ envs: Annotated[str, typer.Option(help="Comma-separated env IDs or 'all'")] = "all",
72
+ agents: Annotated[str, typer.Option(help="Comma-separated agent types or 'all'")] = "all",
73
+ seeds: Annotated[int, typer.Option(help="Number of random seeds")] = 10,
74
+ episodes: Annotated[int, typer.Option(help="Episodes per run")] = 500,
75
+ difficulty: Annotated[str, typer.Option(help="easy / medium / hard")] = "medium",
76
+ output_dir: Annotated[str, typer.Option(help="Output directory")] = "results",
77
+ ) -> None:
78
+ """Run the full benchmark suite."""
79
+ from neorx.causalbiorl.benchmark import ALL_AGENTS, ALL_ENVS, run_benchmark
80
+
81
+ env_list = ALL_ENVS if envs == "all" else [e.strip() for e in envs.split(",")]
82
+ agent_list = ALL_AGENTS if agents == "all" else [a.strip() for a in agents.split(",")]
83
+
84
+ typer.echo(f"🧬 Benchmark: {env_list} × {agent_list} × {seeds} seeds")
85
+ run_benchmark(
86
+ envs=env_list,
87
+ agents=agent_list,
88
+ n_seeds=seeds,
89
+ n_episodes=episodes,
90
+ difficulty=difficulty,
91
+ output_dir=output_dir,
92
+ )
93
+ typer.echo("✅ Benchmark complete — results saved to " + output_dir)
94
+
95
+
96
+ # ────────────────────────────────────────────────────────────────────────── #
97
+ # play #
98
+ # ────────────────────────────────────────────────────────────────────────── #
99
+
100
+
101
+ @app.command()
102
+ def play(
103
+ env: Annotated[str, typer.Option(help="Gymnasium environment ID")] = "GeneticToggle-v0",
104
+ difficulty: Annotated[str, typer.Option(help="easy / medium / hard")] = "easy",
105
+ steps: Annotated[int, typer.Option(help="Max steps")] = 100,
106
+ ) -> None:
107
+ """Interactive manual control of an environment.
108
+
109
+ Each step prompts you for action values (comma-separated floats in [0, 1]).
110
+ """
111
+ environment = gym.make(env, difficulty=difficulty)
112
+ state, info = environment.reset()
113
+ action_dim = int(np.prod(environment.action_space.shape)) # type: ignore[union-attr]
114
+
115
+ typer.echo(f"🎮 Playing {env} ({difficulty}) — action dim = {action_dim}")
116
+ typer.echo(f" State: {state}")
117
+
118
+ for t in range(steps):
119
+ raw = typer.prompt(f"Step {t} — enter {action_dim} action values (comma-sep, 0-1)")
120
+ try:
121
+ action = np.array([float(x.strip()) for x in raw.split(",")], dtype=np.float32)
122
+ except ValueError:
123
+ typer.echo(" ⚠ Invalid input — using zeros.")
124
+ action = np.zeros(action_dim, dtype=np.float32)
125
+
126
+ state, reward, terminated, truncated, info = environment.step(action)
127
+ typer.echo(f" State: {np.round(state, 3)} | Reward: {reward:.3f}")
128
+
129
+ if terminated or truncated:
130
+ typer.echo(" Episode finished.")
131
+ break
132
+
133
+ environment.close()
134
+
135
+
136
+ # ────────────────────────────────────────────────────────────────────────── #
137
+ # visualise #
138
+ # ────────────────────────────────────────────────────────────────────────── #
139
+
140
+
141
+ @app.command()
142
+ def visualise(
143
+ env: Annotated[str, typer.Option(help="Gymnasium environment ID")] = "GeneticToggle-v0",
144
+ agent: Annotated[str, typer.Option(help="Agent type")] = "causal",
145
+ episodes: Annotated[int, typer.Option(help="Training episodes before viz")] = 200,
146
+ difficulty: Annotated[str, typer.Option(help="Difficulty")] = "medium",
147
+ output_dir: Annotated[str, typer.Option(help="Output directory")] = "results",
148
+ ) -> None:
149
+ """Train an agent and visualise the learned causal graph + trajectories."""
150
+ from neorx.causalbiorl.benchmark import _make_agent
151
+ from neorx.causalbiorl.viz import plot_causal_graph, plot_trajectories
152
+
153
+ typer.echo(f"🔬 Training {agent} on {env} for visualisation …")
154
+ environment = gym.make(env, difficulty=difficulty)
155
+ ag = _make_agent(agent, environment, seed=42)
156
+ metrics = ag.train(n_episodes=episodes, verbose=True)
157
+ environment.close()
158
+
159
+ out = Path(output_dir)
160
+ out.mkdir(parents=True, exist_ok=True)
161
+
162
+ # Reward trajectory
163
+ plot_trajectories(
164
+ {agent: metrics["episode_rewards"]},
165
+ title=f"{env} — {agent}",
166
+ save_path=out / f"{env}_{agent}_trajectory.png",
167
+ )
168
+
169
+ # Causal graph (if causal agent)
170
+ if hasattr(ag, "get_learned_graph"):
171
+ graph = ag.get_learned_graph()
172
+ if graph is not None:
173
+ plot_causal_graph(
174
+ graph,
175
+ title=f"Learned Causal Graph — {env}",
176
+ save_path=out / f"{env}_{agent}_causal_graph.png",
177
+ )
178
+ typer.echo(f" Learned graph: {graph.number_of_nodes()} nodes, {graph.number_of_edges()} edges")
179
+
180
+ typer.echo(f"✅ Visualisations saved to {out}")
181
+
182
+
183
+ # ────────────────────────────────────────────────────────────────────────── #
184
+ # Entry point #
185
+ # ────────────────────────────────────────────────────────────────────────── #
186
+
187
+
188
+ def main() -> None:
189
+ app()
190
+
191
+
192
+ if __name__ == "__main__":
193
+ main()
@@ -0,0 +1,6 @@
1
+ """CausalBioRL agents — causal and baseline RL agents."""
2
+
3
+ from neorx.causalbiorl.agents.causal_agent import CausalAgent
4
+ from neorx.causalbiorl.agents.baseline_agent import PPOAgent, SACAgent, RandomAgent
5
+
6
+ __all__ = ["CausalAgent", "PPOAgent", "SACAgent", "RandomAgent"]
@@ -0,0 +1,161 @@
1
+ """
2
+ Baseline RL agents for comparison — PPO, SAC, and Random.
3
+
4
+ Wraps ``stable-baselines3`` implementations with a consistent interface
5
+ so that benchmarking code can treat all agents uniformly.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import Any
11
+
12
+ import gymnasium as gym
13
+ import numpy as np
14
+ from numpy.typing import NDArray
15
+ from tqdm import tqdm
16
+
17
+
18
+ # ────────────────────────────────────────────────────────────────────────── #
19
+ # Random agent #
20
+ # ────────────────────────────────────────────────────────────────────────── #
21
+
22
+
23
+ class RandomAgent:
24
+ """Agent that samples uniformly from the action space."""
25
+
26
+ def __init__(self, env: gym.Env, seed: int | None = None) -> None:
27
+ self.env = env
28
+ self.rng = np.random.default_rng(seed)
29
+ self.episode_rewards: list[float] = []
30
+ self.episode_lengths: list[int] = []
31
+
32
+ def train(
33
+ self, n_episodes: int = 500, verbose: bool = True, **_kwargs: object
34
+ ) -> dict[str, Any]:
35
+ pbar = tqdm(range(n_episodes), desc="RandomAgent", disable=not verbose)
36
+ total_steps = 0
37
+ for _ in pbar:
38
+ state, _ = self.env.reset(seed=int(self.rng.integers(0, 2**31)))
39
+ ep_reward = 0.0
40
+ ep_len = 0
41
+ terminated = truncated = False
42
+ while not (terminated or truncated):
43
+ action = self.env.action_space.sample()
44
+ state, reward, terminated, truncated, _ = self.env.step(action)
45
+ ep_reward += float(reward)
46
+ ep_len += 1
47
+ total_steps += 1
48
+ self.episode_rewards.append(ep_reward)
49
+ self.episode_lengths.append(ep_len)
50
+ return {
51
+ "episode_rewards": self.episode_rewards,
52
+ "episode_lengths": self.episode_lengths,
53
+ "total_steps": total_steps,
54
+ }
55
+
56
+ def act(self, state: NDArray[np.floating]) -> NDArray[np.floating]:
57
+ return self.env.action_space.sample()
58
+
59
+
60
+ # ────────────────────────────────────────────────────────────────────────── #
61
+ # Stable-Baselines3 wrappers #
62
+ # ────────────────────────────────────────────────────────────────────────── #
63
+
64
+
65
+ class _SB3Agent:
66
+ """Base wrapper for stable-baselines3 algorithms."""
67
+
68
+ _algo_cls: type | None = None
69
+
70
+ def __init__(
71
+ self,
72
+ env: gym.Env,
73
+ seed: int | None = None,
74
+ policy: str = "MlpPolicy",
75
+ **sb3_kwargs: object,
76
+ ) -> None:
77
+ self.env = env
78
+ self.seed = seed
79
+ self.episode_rewards: list[float] = []
80
+ self.episode_lengths: list[int] = []
81
+
82
+ try:
83
+ import stable_baselines3 # noqa: F401
84
+ except ImportError as exc:
85
+ raise ImportError(
86
+ "stable-baselines3 is required for PPO/SAC baselines: "
87
+ "pip install stable-baselines3"
88
+ ) from exc
89
+
90
+ assert self._algo_cls is not None
91
+ self.model = self._algo_cls(
92
+ policy,
93
+ env,
94
+ seed=seed,
95
+ verbose=0,
96
+ **sb3_kwargs,
97
+ )
98
+
99
+ def train(
100
+ self,
101
+ n_episodes: int = 500,
102
+ verbose: bool = True,
103
+ **_kwargs: object,
104
+ ) -> dict[str, Any]:
105
+ """Train for *n_episodes* via SB3's ``learn`` method.
106
+
107
+ SB3 operates in total timesteps, so we estimate the budget as
108
+ ``n_episodes × env.spec.max_episode_steps`` (or 500 fallback).
109
+ """
110
+ max_ep = getattr(self.env.spec, "max_episode_steps", None) or 500
111
+ total_timesteps = n_episodes * max_ep
112
+
113
+ # Use a callback to record per-episode statistics
114
+ self.model.learn(total_timesteps=total_timesteps, progress_bar=verbose)
115
+
116
+ # Evaluate to gather reward traces
117
+ self._evaluate(n_episodes, verbose)
118
+ return {
119
+ "episode_rewards": self.episode_rewards,
120
+ "episode_lengths": self.episode_lengths,
121
+ "total_steps": sum(self.episode_lengths),
122
+ }
123
+
124
+ def act(self, state: NDArray[np.floating]) -> NDArray[np.floating]:
125
+ action, _ = self.model.predict(state, deterministic=True)
126
+ return action
127
+
128
+ def _evaluate(self, n_episodes: int, verbose: bool) -> None:
129
+ pbar = tqdm(range(n_episodes), desc=self.__class__.__name__ + " eval", disable=not verbose)
130
+ for _ in pbar:
131
+ state, _ = self.env.reset()
132
+ ep_reward = 0.0
133
+ ep_len = 0
134
+ terminated = truncated = False
135
+ while not (terminated or truncated):
136
+ action = self.act(state)
137
+ state, reward, terminated, truncated, _ = self.env.step(action)
138
+ ep_reward += float(reward)
139
+ ep_len += 1
140
+ self.episode_rewards.append(ep_reward)
141
+ self.episode_lengths.append(ep_len)
142
+
143
+
144
+ class PPOAgent(_SB3Agent):
145
+ """Proximal Policy Optimisation (PPO) via stable-baselines3."""
146
+
147
+ def __init__(self, env: gym.Env, seed: int | None = None, **kwargs: object) -> None:
148
+ from stable_baselines3 import PPO
149
+
150
+ self._algo_cls = PPO # type: ignore[assignment]
151
+ super().__init__(env, seed=seed, **kwargs)
152
+
153
+
154
+ class SACAgent(_SB3Agent):
155
+ """Soft Actor-Critic (SAC) via stable-baselines3."""
156
+
157
+ def __init__(self, env: gym.Env, seed: int | None = None, **kwargs: object) -> None:
158
+ from stable_baselines3 import SAC
159
+
160
+ self._algo_cls = SAC # type: ignore[assignment]
161
+ super().__init__(env, seed=seed, **kwargs)