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.
- neorx/__init__.py +118 -0
- neorx/causalbiorl/__init__.py +23 -0
- neorx/causalbiorl/__main__.py +193 -0
- neorx/causalbiorl/agents/__init__.py +6 -0
- neorx/causalbiorl/agents/baseline_agent.py +161 -0
- neorx/causalbiorl/agents/causal_agent.py +434 -0
- neorx/causalbiorl/benchmark.py +276 -0
- neorx/causalbiorl/causal/__init__.py +40 -0
- neorx/causalbiorl/causal/discovery.py +237 -0
- neorx/causalbiorl/causal/graph_encoder.py +555 -0
- neorx/causalbiorl/causal/planner.py +333 -0
- neorx/causalbiorl/causal/reward_learner.py +384 -0
- neorx/causalbiorl/causal/scm.py +373 -0
- neorx/causalbiorl/causal/surrogate_docker.py +441 -0
- neorx/causalbiorl/envs/__init__.py +11 -0
- neorx/causalbiorl/envs/cell_growth.py +289 -0
- neorx/causalbiorl/envs/drug_discovery.py +884 -0
- neorx/causalbiorl/envs/metabolic_pathway.py +303 -0
- neorx/causalbiorl/envs/registration.py +41 -0
- neorx/causalbiorl/envs/toggle_switch.py +269 -0
- neorx/causalbiorl/models.py +94 -0
- neorx/causalbiorl/viz.py +207 -0
- neorx/cli/__init__.py +46 -0
- neorx/cli/__main__.py +6 -0
- neorx/core/__init__.py +116 -0
- neorx/core/__main__.py +275 -0
- neorx/core/api.py +214 -0
- neorx/core/bio/__init__.py +0 -0
- neorx/core/bio/classifier.py +462 -0
- neorx/core/bio/tissue_filter.py +572 -0
- neorx/core/cache.py +173 -0
- neorx/core/causal/__init__.py +0 -0
- neorx/core/causal/counterfactual.py +312 -0
- neorx/core/causal/identifier.py +1291 -0
- neorx/core/graph/__init__.py +0 -0
- neorx/core/graph/graph_builder.py +446 -0
- neorx/core/graph/models.py +359 -0
- neorx/core/graph/persistence.py +349 -0
- neorx/core/literature_validator.py +281 -0
- neorx/core/pipeline.py +920 -0
- neorx/core/py.typed +1 -0
- neorx/core/report.py +450 -0
- neorx/core/scoring/__init__.py +0 -0
- neorx/core/scoring/admet.py +273 -0
- neorx/core/scoring/scorer.py +260 -0
- neorx/core/sources/__init__.py +42 -0
- neorx/core/sources/chembl.py +643 -0
- neorx/core/sources/kegg.py +228 -0
- neorx/core/sources/monarch.py +355 -0
- neorx/core/sources/open_targets.py +301 -0
- neorx/core/sources/pdb.py +196 -0
- neorx/core/sources/reactome.py +180 -0
- neorx/core/sources/string_db.py +198 -0
- neorx/core/sources/uniprot.py +208 -0
- neorx/core/templates/report.html +440 -0
- neorx/core/validator.py +412 -0
- neorx/dockbot/__init__.py +79 -0
- neorx/dockbot/__main__.py +315 -0
- neorx/dockbot/api.py +292 -0
- neorx/dockbot/binding_site.py +350 -0
- neorx/dockbot/docker.py +335 -0
- neorx/dockbot/ligand_prep.py +306 -0
- neorx/dockbot/models.py +209 -0
- neorx/dockbot/parallel.py +248 -0
- neorx/dockbot/protein_prep.py +424 -0
- neorx/dockbot/report.py +361 -0
- neorx/dockbot/scorer.py +238 -0
- neorx/dockbot/viz.py +278 -0
- neorx/experiments/__init__.py +24 -0
- neorx/experiments/__main__.py +100 -0
- neorx/experiments/capture.py +85 -0
- neorx/experiments/chembl.py +74 -0
- neorx/experiments/gates.py +314 -0
- neorx/experiments/record.py +204 -0
- neorx/experiments/registry.py +105 -0
- neorx/experiments/replay.py +129 -0
- neorx/genmol/__init__.py +60 -0
- neorx/genmol/__main__.py +335 -0
- neorx/genmol/api.py +235 -0
- neorx/genmol/assets/__init__.py +0 -0
- neorx/genmol/assets/molvae_chembl36.pt +0 -0
- neorx/genmol/assets/tokenizer.json +39 -0
- neorx/genmol/configs/default.yaml +58 -0
- neorx/genmol/data/__init__.py +23 -0
- neorx/genmol/data/dataset.py +204 -0
- neorx/genmol/data/download.py +239 -0
- neorx/genmol/data/preprocess.py +262 -0
- neorx/genmol/data/tokenizer.py +318 -0
- neorx/genmol/evaluation/__init__.py +44 -0
- neorx/genmol/evaluation/distribution.py +188 -0
- neorx/genmol/evaluation/metrics.py +238 -0
- neorx/genmol/evaluation/visualise.py +392 -0
- neorx/genmol/generate.py +371 -0
- neorx/genmol/models/__init__.py +17 -0
- neorx/genmol/models/cvae.py +493 -0
- neorx/genmol/models/vae.py +538 -0
- neorx/genmol/pretrained.py +61 -0
- neorx/genmol/train.py +387 -0
- neorx/mirrorfold/__init__.py +121 -0
- neorx/mirrorfold/__main__.py +306 -0
- neorx/mirrorfold/analysis.py +405 -0
- neorx/mirrorfold/api.py +225 -0
- neorx/mirrorfold/compare.py +520 -0
- neorx/mirrorfold/mirror.py +370 -0
- neorx/mirrorfold/models.py +219 -0
- neorx/mirrorfold/predictor.py +503 -0
- neorx/mirrorfold/therapeutic.py +381 -0
- neorx/mirrorfold/viz.py +559 -0
- neorx/molscreen/__init__.py +64 -0
- neorx/molscreen/__main__.py +87 -0
- neorx/molscreen/accessibility.py +131 -0
- neorx/molscreen/data/.gitignore +2 -0
- neorx/molscreen/filters.py +350 -0
- neorx/molscreen/models.py +285 -0
- neorx/molscreen/parser.py +204 -0
- neorx/molscreen/properties.py +211 -0
- neorx/molscreen/similarity.py +384 -0
- neorx/py.typed +0 -0
- neorx-0.2.0.dist-info/METADATA +497 -0
- neorx-0.2.0.dist-info/RECORD +123 -0
- neorx-0.2.0.dist-info/WHEEL +4 -0
- neorx-0.2.0.dist-info/entry_points.txt +6 -0
- 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)
|