diff --git a/.gitignore b/.gitignore index 83cfc3f..dd76b50 100644 --- a/.gitignore +++ b/.gitignore @@ -5,6 +5,7 @@ wandb/ outputs/ multirun/ metrics/ +adjacency_debug.txt # Python-generated files __pycache__/ diff --git a/README.md b/README.md index 572fd73..e9cabc6 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,21 @@ To set up the UV module, you can run the following command: uv sync --frozen ``` +## Repository Structure + +```text +. +├── configs/ # Hydra configuration files (YAML) +├── docs/ # Comprehensive documentation and API guides +├── runs/ # Default output directory for Hydra and training artifacts +├── scripts/ # High-level entrypoints for training, simulation, and evaluation +├── src/ +│ └── brittle_star_project/ # Core library and environment logic +│ ├── evaluation/ # Checkpoint evaluation, rollout logic, and metrics persistence +│ └── trainers/ # Training loop implementations (e.g., PPO) +└── tests/ # Unit and integration tests +``` + ## Usage For detailed instructions on how to use the project, please refer to the **[API Documentation](docs/README.md)**. diff --git a/configs/centralized-final.yaml b/configs/centralized-final.yaml new file mode 100644 index 0000000..f40570d --- /dev/null +++ b/configs/centralized-final.yaml @@ -0,0 +1,71 @@ +# Custom Main Configuration +# +# Use with: +# uv run python scripts/train.py --config-name main_config_custom +# +# This keeps the project defaults intact while giving you a single custom +# training entrypoint you can edit freely. + +defaults: + - brittle_star_config + - experiment: base + - logging: default + - evaluation: default + - ppo: default + - architecture: centralized + - morphology: 5_arms_full + - arena: default + - environment: directed_locomotion + - simulation: default + - _self_ + +morphology: + morph_mode: CENTRALIZED + +experiment: + exp_name: "final-models/centralized/" + seed: 42 + torch_deterministic: true + cuda: true + +logging: + track: true + save_model: true + save_checkpoints: true + upload_final_model: true + upload_checkpoints: true + checkpoint_frequency: 20 + wandb_project_name: "final-models" + +evaluation: + evaluate_checkpoints: true + eval_max_steps: 2000 + eval_seed: 0 + +ppo: + learning_rate: 0.0001 + total_timesteps: 16384000 + num_envs: 128 + num_steps: 64 + anneal_lr: true + gamma: 0.99 + gae_lambda: 0.95 + num_minibatches: 32 + update_epochs: 4 + norm_adv: true + clip_coef: 0.2 + clip_vloss: true + ent_coef: 0.001 + vf_coef: 1.0 + max_grad_norm: 0.5 + target_kl: 0.02 + +environment: + simulation_time: 100000.0 + target_distance: 3.0 + +hydra: + job: + chdir: true + run: + dir: ${experiment.base_run_dir}/${experiment.exp_name}/${now:%Y-%m-%d}/${now:%H-%M-%S} diff --git a/configs/environment/directed_locomotion.yaml b/configs/environment/directed_locomotion.yaml index 465d80e..b664ca1 100644 --- a/configs/environment/directed_locomotion.yaml +++ b/configs/environment/directed_locomotion.yaml @@ -2,11 +2,11 @@ # Baseline task setting. task: DIRECTED_LOCOMOTION -simulation_time: 5000.0 +simulation_time: 100000.0 num_physics_steps_per_control_step: 10 time_scale: 2 camera_ids: [0, 1] render_size: [480, 640] joint_randomization_noise_scale: 0.0 -target_distance: 0.6 +target_distance: 3.0 light_perlin_noise_scale: 0 diff --git a/configs/environment/light_escape.yaml b/configs/environment/light_escape.yaml index c79cfb9..ade85ba 100644 --- a/configs/environment/light_escape.yaml +++ b/configs/environment/light_escape.yaml @@ -2,7 +2,7 @@ # Advanced task requiring movement away from light source. task: LIGHT_ESCAPE -simulation_time: 5.0 +simulation_time: 100000.0 num_physics_steps_per_control_step: 10 time_scale: 2 camera_ids: [0, 1] diff --git a/configs/evaluation/default.yaml b/configs/evaluation/default.yaml index ac49e46..a7c2a00 100644 --- a/configs/evaluation/default.yaml +++ b/configs/evaluation/default.yaml @@ -3,6 +3,6 @@ evaluate_checkpoints: false # Max number of control steps during evaluation rollout. -eval_max_steps: 5000 +eval_max_steps: 2000 # Seed for deterministic evaluation reset. eval_seed: 0 diff --git a/configs/evaluation/poster.yaml b/configs/evaluation/poster.yaml new file mode 100644 index 0000000..0777cd1 --- /dev/null +++ b/configs/evaluation/poster.yaml @@ -0,0 +1,23 @@ +# @package evaluation +# Configuration for the models used in the poster comparison. + +# Standard evaluation settings +evaluate_checkpoints: false +eval_max_steps: 5000 +eval_seed: 0 + +# Cross-model comparison settings +# We use 10 episodes to get a more robust average for the final poster results. +comparison_base_seed: 0 +comparison_num_episodes: 2 +comparison_output_csv: "runs/evaluation/comparison.csv" + +# Paths to the .cleanrl_model files to be compared (relative to workspace root). +comparison_models: + - "runs/input-space-2-arms/2026-05-02/08-14-58/final_model.flax" + +# Path to the morphologies to evaluate against. +comparison_morphologies: + - "configs/morphology/5_arms_full.yaml" + - "configs/morphology/3_arms.yaml" + - "configs/morphology/2_arms.yaml" diff --git a/configs/fully-connected-final.yaml b/configs/fully-connected-final.yaml new file mode 100644 index 0000000..8d60451 --- /dev/null +++ b/configs/fully-connected-final.yaml @@ -0,0 +1,74 @@ +# Custom Main Configuration +# +# Use with: +# uv run python scripts/train.py --config-name main_config_custom +# +# This keeps the project defaults intact while giving you a single custom +# training entrypoint you can edit freely. + +defaults: + - brittle_star_config + - experiment: base + - logging: default + - evaluation: default + - ppo: default + - architecture: decentralized + - morphology: 5_arms_full + - arena: default + - environment: directed_locomotion + - simulation: default + - _self_ + +architecture: + topology_type: "fully_connected" + +morphology: + morph_mode: FULLY_CONNECTED + +experiment: + exp_name: "final-models/fully-connected/" + seed: 42 + torch_deterministic: true + cuda: true + +logging: + track: true + save_model: true + save_checkpoints: true + upload_final_model: true + upload_checkpoints: true + checkpoint_frequency: 20 + wandb_project_name: "final-models" + +evaluation: + evaluate_checkpoints: true + eval_max_steps: 2000 + eval_seed: 0 + +ppo: + learning_rate: 0.0001 + total_timesteps: 16384000 + num_envs: 128 + num_steps: 64 + anneal_lr: true + gamma: 0.99 + gae_lambda: 0.95 + num_minibatches: 32 + update_epochs: 4 + norm_adv: true + clip_coef: 0.2 + clip_vloss: true + ent_coef: 0.001 + vf_coef: 1.0 + max_grad_norm: 0.5 + target_kl: 0.02 + +environment: + simulation_time: 100000.0 + target_distance: 3.0 + +hydra: + job: + chdir: true + run: + dir: ${experiment.base_run_dir}/${experiment.exp_name}/${now:%Y-%m-%d}/${now:%H-%M-%S} diff --git a/configs/ring-final.yaml b/configs/ring-final.yaml new file mode 100644 index 0000000..ffba64f --- /dev/null +++ b/configs/ring-final.yaml @@ -0,0 +1,74 @@ +# Custom Main Configuration +# +# Use with: +# uv run python scripts/train.py --config-name main_config_custom +# +# This keeps the project defaults intact while giving you a single custom +# training entrypoint you can edit freely. + +defaults: + - brittle_star_config + - experiment: base + - logging: default + - evaluation: default + - ppo: default + - architecture: decentralized + - morphology: 5_arms_full + - arena: default + - environment: directed_locomotion + - simulation: default + - _self_ + +architecture: + topology_type: "ring" + +morphology: + morph_mode: RING + +experiment: + exp_name: "final-models/ring/" + seed: 42 + torch_deterministic: true + cuda: true + +logging: + track: true + save_model: true + save_checkpoints: true + upload_final_model: true + upload_checkpoints: true + checkpoint_frequency: 20 + wandb_project_name: "final-models" + +evaluation: + evaluate_checkpoints: true + eval_max_steps: 2000 + eval_seed: 0 + +ppo: + learning_rate: 0.0001 + total_timesteps: 16384000 + num_envs: 128 + num_steps: 64 + anneal_lr: true + gamma: 0.99 + gae_lambda: 0.95 + num_minibatches: 32 + update_epochs: 4 + norm_adv: true + clip_coef: 0.2 + clip_vloss: true + ent_coef: 0.001 + vf_coef: 1.0 + max_grad_norm: 0.5 + target_kl: 0.02 + +environment: + simulation_time: 100000.0 + target_distance: 3.0 + +hydra: + job: + chdir: true + run: + dir: ${experiment.base_run_dir}/${experiment.exp_name}/${now:%Y-%m-%d}/${now:%H-%M-%S} diff --git a/docs/api/evaluation.md b/docs/api/evaluation.md new file mode 100644 index 0000000..e5f91d6 --- /dev/null +++ b/docs/api/evaluation.md @@ -0,0 +1,56 @@ +# Checkpoint & Model Evaluation + +This guide covers how to evaluate trained brittle star models, with a focus on measuring defect tolerance (amputations) across different controller architectures. + +## Checkpoint Evaluation (During Training) + +The `PPOTrainer` can automatically evaluate every saved checkpoint using the fast MJX backend. This is enabled via configuration. + +### Configuration + +In your experiment config or via CLI: +```bash +python scripts/train.py evaluation.evaluate_checkpoints=true evaluation.eval_max_steps=5000 +``` + +Results are saved to `runs//metrics/checkpoint_evaluation.csv` and synced to Weights & Biases if enabled. + +## Cross-Model & Defect Tolerance Analysis + +To measure how well different controllers handle damage (amputations), use `scripts/compare_models.py`. This script performs a grid search over models x morphologies. + +1. Create or update a YAML file in `configs/evaluation`. +2. Run the benchmark: + +```bash +python scripts/compare_models.py evaluation=poster +``` + +The script will evaluate every combination of model and morphology for the specified number of episodes. + +The results are saved to a CSV (default: `metrics/model_comparison.csv`). + +### CSV Schema + +| Column | Description | +|-----------------------|--------------------------------------------------------------| +| `model_path` | Path to the trained weights. | +| `architecture` | The `morph_mode` of the model (e.g., `CENTRALIZED`, `RING`). | +| `arm_0` ... `arm_4` | Number of segments in each arm slot (0 = amputated). | +| `num_active_arms` | Total number of arms with segments > 0. | +| `seed` | The episode seed. | +| `eval_return` | Accumulated shaped reward. | +| `approx_max_velocity` | Average velocity: `(initial_dist - final_dist) / steps`. | +| `reached_target` | Whether the robot finished within the success radius. | + +## Post-hoc Checkpoint Scanning + +If you need to re-evaluate every saved checkpoint in a run (e.g., to generate a learning curve with different metrics): + +```bash +python scripts/evaluate_checkpoints.py \ + simulation.model_path=runs//final_model.flax \ + evaluation.eval_max_steps=2000 +``` + +This script scans the `checkpoints/` directory of the specified run and evaluates every `.flax` file it finds using the model's training morphology. diff --git a/docs/api/simulation.md b/docs/api/simulation.md index 0a0e041..1e288b9 100644 --- a/docs/api/simulation.md +++ b/docs/api/simulation.md @@ -37,3 +37,5 @@ uv run scripts/simulate.py \ Videos and evaluation metadata are stored in timestamped folders alongside the model: `runs/your_run/final_model_evaluations/eval_/simulation.mp4` + +For batch evaluation and cross-model comparison, see the **[Evaluation Guide](./evaluation.md)**. diff --git a/docs/api/training.md b/docs/api/training.md index a9678bb..d3e6c51 100644 --- a/docs/api/training.md +++ b/docs/api/training.md @@ -38,12 +38,18 @@ To run with your custom experiment file: uv run python scripts/train.py experiment=my_experiment ``` -### Command-Line Overrides - -You can override any parameter directly from the command line using Hydra's dot notation. This is useful for quick tests: - ```bash uv run python scripts/train.py ppo.learning_rate=0.001 ppo.num_envs=32 logging.track=true ``` +## Evaluation During Training + +By default, the trainer saves checkpoints but does not evaluate them. To enable automatic headless evaluation of every saved checkpoint, set `evaluation.evaluate_checkpoints=true`: + +```bash +uv run python scripts/train.py evaluation.evaluate_checkpoints=true +``` + +For more details on evaluation metrics and comparison tools, see [Evaluation](./evaluation.md). + For more details on tracking your experiments, see [Tracking & Monitoring](./tracking.md). diff --git a/scripts/compare_models.py b/scripts/compare_models.py new file mode 100644 index 0000000..b3e8338 --- /dev/null +++ b/scripts/compare_models.py @@ -0,0 +1,182 @@ +"""Compare multiple trained policies across shared evaluation conditions. + +For each model listed in evaluation.comparison_models, this script runs +`comparison_num_episodes` headless rollouts (seeded sequentially from +`comparison_base_seed`) and writes a results CSV to `comparison_output_csv`. + +Results include two metrics per episode: +- `eval_return` — shaped reward (same function used during training) +- `max_velocity` — approximated as initial_xy_dist / steps taken + +Usage: + # With the default evaluation config + python scripts/compare_models.py evaluation=poster + + # Override the output path on the fly + python scripts/compare_models.py evaluation=poster \\ + evaluation.comparison_output_csv=metrics/quick_comparison.csv +""" + +from __future__ import annotations + +import csv +import logging +import time +from pathlib import Path + +import hydra +from omegaconf import DictConfig, OmegaConf + +from brittle_star_project.configs.main_config import BrittleStarConfig +from brittle_star_project.configs.register_configs import register_configs +from brittle_star_project.evaluation import build_eval_env +from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs +from brittle_star_project.evaluation.rollout import rollout_headless + +_FIELDNAMES = [ + "model_path", + "architecture", + "arm_0", + "arm_1", + "arm_2", + "arm_3", + "arm_4", + "num_active_arms", + "seed", + "reached_target", + "episode_length", + "eval_return", + "initial_target_distance", + "final_xy_dist", + "approx_max_velocity", +] + + +def _approx_max_velocity(result) -> float | None: + """Approximate max velocity as distance covered per step. + + This is a rough upper bound: (initial_dist - final_dist) / steps. + """ + if result.initial_target_distance is None or result.final_xy_dist is None or result.length <= 0: + return None + dist_covered = result.initial_target_distance - result.final_xy_dist + return dist_covered / result.length + + +@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") +def main(dict_cfg: DictConfig) -> None: + logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s") + logger = logging.getLogger(__name__) + + cfg: BrittleStarConfig = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) + ) + eval_cfg = cfg.evaluation + + model_paths = [str(p) for p in eval_cfg.comparison_models] + if not model_paths: + raise ValueError( + "evaluation.comparison_models is empty. " + "Add at least one model path in your evaluation config." + ) + + base_seed = int(eval_cfg.comparison_base_seed) + num_episodes = int(eval_cfg.comparison_num_episodes) + max_steps = int(eval_cfg.eval_max_steps) + + seeds = list(range(base_seed, base_seed + num_episodes)) + + output_path = Path(hydra.utils.to_absolute_path(eval_cfg.comparison_output_csv)) + output_path.parent.mkdir(parents=True, exist_ok=True) + + logger.info( + f"Comparing {len(model_paths)} models over {num_episodes} episodes " + f"(seeds {seeds[0]}–{seeds[-1]})." + ) + logger.info(f"Results will be written to: {output_path}") + + with open(output_path, "w", newline="") as csv_file: + writer = csv.DictWriter(csv_file, fieldnames=_FIELDNAMES) + writer.writeheader() + + for model_path_str in model_paths: + model_path = Path(hydra.utils.to_absolute_path(model_path_str)) + logger.info(f"Evaluating model: {model_path.name}") + + try: + metadata = load_metadata(model_path) + except FileNotFoundError as e: + logger.warning(f"Skipping model — {e}") + continue + + training = metadata_to_configs(metadata) + + # Determine morphologies to evaluate + # If comparison_morphologies is empty, use the model's training morphology + morphologies = [None] + if eval_cfg.comparison_morphologies: + morphologies = [ + Path(hydra.utils.to_absolute_path(m)) for m in eval_cfg.comparison_morphologies + ] + + for morph_path in morphologies: + morph_label = morph_path.name if morph_path else "training" + logger.info(f" Morphology: {morph_label}") + + bundle = build_eval_env( + model_path=model_path, + training=training, + metadata=metadata, + morphology_override_path=morph_path, + ) + + for seed in seeds: + t0 = time.time() + result = rollout_headless( + env=bundle.env, + policy=bundle.policy, + seed=seed, + max_steps=max_steps, + action_low=bundle.action_low, + action_high=bundle.action_high, + action_mask=bundle.action_mask, + ) + elapsed = time.time() - t0 + + velocity = _approx_max_velocity(result) + + logger.debug( + f" seed={seed:3d} | " + f"reached={str(result.reached_target):<5} | " + f"return={result.return_:+8.3f} | " + f"steps={result.length:4d} | " + f"({elapsed:.1f}s)" + ) + + row = { + "model_path": model_path_str, + "architecture": bundle.architecture, + "num_active_arms": bundle.num_active_arms, + "seed": seed, + "reached_target": result.reached_target, + "episode_length": result.length, + "eval_return": result.return_, + "initial_target_distance": result.initial_target_distance, + "final_xy_dist": result.final_xy_dist, + "approx_max_velocity": velocity, + } + # Add per-arm segments + for i, segs in enumerate(bundle.segments_per_arm): + row[f"arm_{i}"] = segs + + writer.writerow(row) + csv_file.flush() + + bundle.env.close() + + logger.info(f"Done. Results saved to {output_path}") + + +if __name__ == "__main__": + register_configs() + main() diff --git a/scripts/evaluate_checkpoints.py b/scripts/evaluate_checkpoints.py new file mode 100644 index 0000000..a63796b --- /dev/null +++ b/scripts/evaluate_checkpoints.py @@ -0,0 +1,264 @@ +"""Re-evaluate saved checkpoints from a completed training run using MJX. + +This script scans the checkpoint directory of a training run (the `checkpoints/` +folder inside a Hydra output directory), loads each `.flax` checkpoint, runs +one deterministic evaluation episode with `build_eval_rollout_fn`, and appends +the result to the run's `metrics/checkpoint_evaluation.csv`. + +It is intended for post-training analysis when per-checkpoint evaluation was not +enabled during training (`evaluate_checkpoints: false`). + +Usage: + python scripts/evaluate_checkpoints.py \ + simulation.model_path=runs/2024-01-01/12-00-00/final_model.flax \ + evaluation.eval_max_steps=5000 \ + evaluation.eval_seed=0 + +The script resolves the run directory from `simulation.model_path`, discovers +all `*.flax` checkpoints under `checkpoints/`, and evaluates them in order. +""" + +from __future__ import annotations +from brittle_star_project.MLPs.mlps import ( + Actor, + GenericDenseLayersWithActivation, + MessagePasser, +) +from brittle_star_project.MLPs.adjancency_builder import build_adjacency +from brittle_star_project.environment import MorphMode +from brittle_star_project.MLPs.routing import apply_per_node +import logging +import re +from pathlib import Path + +import hydra +import jax +import numpy as np +import jax.numpy as jnp + +from omegaconf import DictConfig, OmegaConf + +from brittle_star_project.configs.main_config import BrittleStarConfig +from brittle_star_project.configs.register_configs import register_configs +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from brittle_star_project.environment.obs_processing import create_obs_processor +from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks +from brittle_star_project.evaluation.checkpoint import ( + load_metadata, + load_params, + metadata_to_configs, +) +from brittle_star_project.evaluation.evaluate_mjx import ( + append_checkpoint_eval_row, + build_eval_rollout_fn, + evaluate_checkpoint_mjx, +) +from brittle_star_project.trainers.PPOTrainer import reward_fn + + +def _parse_iteration(checkpoint_path: Path) -> int: + """Parse the iteration number from a checkpoint filename like `checkpoint_0042.flax`.""" + match = re.search(r"(\d+)", checkpoint_path.stem) + return int(match.group(1)) if match else -1 + + +@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") +def main(dict_cfg: DictConfig) -> None: + logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s") + logger = logging.getLogger(__name__) + + cfg: BrittleStarConfig = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg) + ) + sim_cfg = cfg.simulation + eval_cfg = cfg.evaluation + + # --- Resolve the model path to find the run directory --- + model_path_str = sim_cfg.model_path + if model_path_str is None: + raise ValueError( + "simulation.model_path must point to the final_model.flax of a training run." + ) + + model_path = Path(hydra.utils.to_absolute_path(model_path_str)) + run_dir = model_path.parent + + checkpoints_dir = run_dir / "checkpoints" + if not checkpoints_dir.exists(): + raise FileNotFoundError( + f"No checkpoints/ directory found in run directory: {run_dir}\n" + "Make sure simulation.model_path points to a completed training run." + ) + + checkpoints = sorted(checkpoints_dir.glob("*.flax"), key=_parse_iteration) + if not checkpoints: + raise FileNotFoundError(f"No .flax checkpoints found in {checkpoints_dir}") + + logger.info(f"Found {len(checkpoints)} checkpoint(s) in {checkpoints_dir}") + + # --- Load sidecar metadata + reconstruct training config --- + metadata_override = ( + Path(hydra.utils.to_absolute_path(sim_cfg.metadata_path)) + if sim_cfg.metadata_path is not None + else None + ) + metadata = load_metadata(model_path, metadata_override) + training = metadata_to_configs(metadata) + + padding_masks = compute_padding_masks( + segments_per_arm=training.morphology.segments_per_arm, + reference_segments_per_arm=training.morphology.segments_per_arm, + ) + + morph_mode = training.morphology.morph_mode + + segments_per_arm = jnp.asarray( + training.morphology.segments_per_arm, + dtype=jnp.int32, + ) + + num_arms = ( + jnp.where( + segments_per_arm > 0, + 1, + 0, + ) + .sum() + .item() + ) + + match morph_mode: + case MorphMode.CENTRALIZED: + needed_copies = 1 + agent_indices = [0, 1, 2, 3, 4] + + case MorphMode.FULLY_CONNECTED | MorphMode.RING: + agent_mask = segments_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0] + needed_copies = num_arms + + case MorphMode.SEGMENT: + agent_mask = segments_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0] + + needed_copies = (segments_per_arm.sum() + num_arms).item() + + obs_processor = create_obs_processor( + bounds_dict=training.obs_bounds.to_bounds_dict(), + padding_masks=padding_masks, + num_arms=num_arms, + needed_copies=needed_copies, + morph_mode=morph_mode, + segments_per_arm=segments_per_arm, + agent_indices=agent_indices, + ) + + env = BrittleStarJaxEnvWrapper( + morphology=training.morphology, + arena=training.arena, + env_config=training.environment, + num_envs=1, + ) + + action_low = np.asarray(env.single_action_space.low, dtype=np.float32) + action_high = np.asarray(env.single_action_space.high, dtype=np.float32) + + sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) + actor = Actor(action_dim=env.single_action_space.shape[0]) + sensor.apply = jax.jit(sensor.apply) + actor.apply = jax.jit(actor.apply) + + eval_fn = build_eval_rollout_fn( + env=env, + obs_processor=obs_processor, + sensor_apply=sensor.apply, + actor_apply=actor.apply, + action_low=action_low, + action_high=action_high, + reward_fn=reward_fn, + ) + + morph_mode = training.morphology.morph_mode + + segments_per_arm = jnp.asarray( + training.morphology.segments_per_arm, + dtype=jnp.int32, + ) + + match morph_mode: + case MorphMode.CENTRALIZED: + needed_copies = 1 + + case MorphMode.FULLY_CONNECTED | MorphMode.RING: + needed_copies = jnp.where(segments_per_arm > 0, 1, 0).sum().item() + + case MorphMode.SEGMENT: + needed_copies = ( + segments_per_arm.sum() + jnp.where(segments_per_arm > 0, 1, 0).sum() + ).item() + + adj = build_adjacency( + training.morphology.segments_per_arm, + morph_mode, + ) + + sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) + + actor = Actor(action_dim=env.single_action_space.shape[0] // needed_copies) + + message_passer = ( + MessagePasser( + hidden_dim=300, + num_propagation_steps=4, + adj_matrix=adj, + ) + if morph_mode != MorphMode.CENTRALIZED + else None + ) + + eval_fn = build_eval_rollout_fn( + env=env, + obs_processor=obs_processor, + sensor_apply=lambda p, x: apply_per_node(sensor.apply, p, x), + actor_apply=lambda p, x: apply_per_node(actor.apply, p, x), + message_passer_apply=(None if message_passer is None else message_passer.apply), + action_low=action_low, + action_high=action_high, + reward_fn=reward_fn, + ) + seed = int(eval_cfg.eval_seed) + max_steps = int(eval_cfg.eval_max_steps) + + logger.info(f"Evaluating each checkpoint (seed={seed}, max_steps={max_steps}).") + + for checkpoint_path in checkpoints: + iteration = _parse_iteration(checkpoint_path) + try: + params = load_params(checkpoint_path) + except Exception as e: + logger.warning(f"Could not load {checkpoint_path.name}: {e}") + continue + + result = evaluate_checkpoint_mjx(eval_fn, params, seed=seed, max_steps=max_steps) + csv_path = append_checkpoint_eval_row( + run_dir, + iteration=iteration, + trained_timesteps=0, # unknown without training logs + result=result, + ) + + logger.debug( + f"checkpoint={iteration:5d} | " + f"reached={str(result.reached_target):<5} | " + f"return={result.eval_return:+8.3f} | " + f"steps={result.steps:4d} | " + f"final_dist={result.final_xy_dist:.3f}" + ) + + logger.info(f"Done. CSV at: {csv_path}") + env.close() + + +if __name__ == "__main__": + register_configs() + main() diff --git a/scripts/simulate.py b/scripts/simulate.py index a56e970..9f3e550 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -13,28 +13,20 @@ from __future__ import annotations from pathlib import Path import hydra -import numpy as np from omegaconf import DictConfig, OmegaConf -import yaml -import jax.numpy as jnp -from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory from brittle_star_project.configs.main_config import BrittleStarConfig from brittle_star_project.configs.register_configs import register_configs -from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks -from brittle_star_project.environment.obs_processing import create_obs_processor -from brittle_star_project.environment.env_config import MorphMode, MorphologyConfig from brittle_star_project.evaluation.checkpoint import load_metadata, metadata_to_configs -from brittle_star_project.evaluation.policy import PolicyAgent +from brittle_star_project.evaluation.eval_env_builder import build_eval_env from brittle_star_project.evaluation.rollout import rollout_headless, rollout_viewer from brittle_star_project.evaluation.video import ( record_episode, create_evaluation_dir, save_evaluation_metadata, ) -from brittle_star_project.MLPs.adjancency_builder import build_adjacency @hydra.main(config_path="../configs", config_name="main_config", version_base="1.3") @@ -63,107 +55,28 @@ def main(dict_cfg: DictConfig) -> None: # 3. Reconstruct typed configs from metadata training = metadata_to_configs(metadata) - # 4. Determine environment morphology - if sim_cfg.morphology_override is not None: - override_path = Path(hydra.utils.to_absolute_path(sim_cfg.morphology_override)) - if not override_path.exists(): - raise FileNotFoundError(f"Could not find morphology override YAML at {override_path}") - with open(override_path, "r") as f: - override_dict = yaml.safe_load(f) - env_morphology = OmegaConf.to_object( - OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict) - ) - else: - env_morphology = training.morphology - - # 5. Build obs_processor with TRAINING morphology padding masks always - padding_masks = compute_padding_masks( - segments_per_arm=env_morphology.segments_per_arm, - reference_segments_per_arm=training.morphology.segments_per_arm, - ) - - segs_per_arm = jnp.array(env_morphology.segments_per_arm) - - needed_copies = 0 - agent_indices = [0, 1, 2, 3, 4] - match env_morphology.morph_mode: - case MorphMode.CENTRALIZED: - needed_copies = 1 - case MorphMode.FULLY_CONNECTED | MorphMode.RING: - agent_mask = segs_per_arm > 0 - agent_indices = jnp.where(agent_mask)[0] - needed_copies = jnp.where(segs_per_arm > 0, 1, 0).sum().item() - case MorphMode.SEGMENT: - agent_mask = segs_per_arm > 0 - agent_indices = jnp.where(agent_mask)[0] - needed_copies = jnp.where(segs_per_arm > 0, 1, 0).sum().item() - needed_copies = (segs_per_arm.sum() + jnp.where(segs_per_arm > 0, 1, 0).sum()).item() - - num_arms = jnp.where(segs_per_arm > 0, 1, 0).sum().item() - - obs_processor = create_obs_processor( - bounds_dict=training.obs_bounds.to_bounds_dict(), - padding_masks=padding_masks, - needed_copies=needed_copies, - num_arms=num_arms, - morph_mode=env_morphology.morph_mode, - segments_per_arm=env_morphology.segments_per_arm, - agent_indices=agent_indices, - ) - - # 6. Build environment - backend = Backend.MJC seed = int(cfg.experiment.seed) - factory = BrittleStarEnvFactory() - raw_env = factory.create_environment( - backend, - env_morphology, - training.arena, - training.environment, - ) - env = BrittleStarEnv( - raw_env, - backend=backend, - config=training.environment, - morphology_config=env_morphology, + # 4-7. Build evaluation environment and policy + override_path = None + if sim_cfg.morphology_override is not None: + override_path = Path(hydra.utils.to_absolute_path(sim_cfg.morphology_override)) + + bundle = build_eval_env( + model_path=model_path, + training=training, + metadata=metadata, + morphology_override_path=override_path, ) + env = bundle.env + policy = bundle.policy + action_low = bundle.action_low + action_high = bundle.action_high + action_mask = bundle.action_mask + state0 = env.reset(seed=seed) - # Calculate the action dimension the model was trained with - trained_action_dim = raw_env.action_space.shape[0] // needed_copies - - # 7. Load policy - message_passing_steps = (metadata.get("architecture", {}) or {}).get("message_passing_steps") - if message_passing_steps is None: - message_passing_steps = 4 - message_passing_steps = int(message_passing_steps) - - adj_matrix = None - if env_morphology.morph_mode != MorphMode.CENTRALIZED: - adj_matrix = build_adjacency(env_morphology.segments_per_arm, env_morphology.morph_mode) - - policy = PolicyAgent.from_checkpoint( - model_path, - action_dim=trained_action_dim, - obs_processor=obs_processor, - message_passing_steps=message_passing_steps, - adj_matrix=adj_matrix, - ) - - # Convert the JAX boolean mask to a numpy array for easy indexing - action_mask = np.asarray(padding_masks["mask_2x"]) - - # Match training's action clipping behavior. - action_space = getattr(raw_env, "action_space", None) - action_low = ( - None if action_space is None else np.asarray(action_space.low, dtype=np.float32).ravel() - ) - action_high = ( - None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel() - ) - # 8. Run simulation headless = bool(sim_cfg.headless) max_steps = sim_cfg.max_steps diff --git a/src/brittle_star_project/MLPs/routing.py b/src/brittle_star_project/MLPs/routing.py new file mode 100644 index 0000000..4033cb9 --- /dev/null +++ b/src/brittle_star_project/MLPs/routing.py @@ -0,0 +1,22 @@ +"""Shared JAX routing utilities for decentralized multi-agent models.""" + +import jax + + +def apply_per_node(apply_fn, params, x): + """Apply a Flax module independently to each node. + + Args: + apply_fn: The module's ``apply`` method (e.g. ``sensor.apply``). + params: Per-node parameters with shape ``(num_nodes, ...)``. + x: Input tensor with shape ``(batch, num_nodes, features)``. + + Returns: + Output tensor with shape ``(batch, num_nodes, out_features)``. + """ + + def apply_single_node(p, x_node): + # x_node: (batch, feat) — one node's input across the batch + return jax.vmap(lambda xi: apply_fn(p, xi))(x_node) + + return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x) diff --git a/src/brittle_star_project/configs/config_evaluation.py b/src/brittle_star_project/configs/config_evaluation.py index 4eeadbc..c9c248b 100644 --- a/src/brittle_star_project/configs/config_evaluation.py +++ b/src/brittle_star_project/configs/config_evaluation.py @@ -1,6 +1,6 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field @dataclass @@ -16,6 +16,20 @@ class EvaluationConfig: eval_max_steps: int = 5000 eval_seed: int = 0 + # Cross-model comparison settings. + # comparison_base_seed is the starting seed for generating episode seeds. + comparison_base_seed: int = 0 + # comparison_num_episodes controls how many target positions to evaluate for each model. + comparison_num_episodes: int = 5 + # comparison_models lists the paths (relative to workspace root) to the .cleanrl_model files. + comparison_models: list[str] = field(default_factory=list) + # Path where the comparison results CSV will be saved (relative to workspace root). + comparison_output_csv: str = "metrics/model_comparison.csv" + # Morphology override YAML paths for cross-morphology comparison. + # Each path points to a file in configs/morphology/ (e.g., "configs/morphology/3_arms.yaml"). + # When empty, each model is evaluated only on its training morphology. + comparison_morphologies: list[str] = field(default_factory=list) + def __post_init__(self) -> None: if self.evaluate_checkpoints and self.eval_max_steps <= 0: raise ValueError( diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index f38d52b..e169ee3 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -109,6 +109,7 @@ def create_obs_processor( joints_per_segment = 2 joints_per_arm = segs_per_arm * joints_per_segment for key, arr in obs.items(): + arr = jnp.asarray(arr) if arr.size == 0: continue @@ -121,7 +122,7 @@ def create_obs_processor( idx = segment_indices[i] taken = jnp.take(arr, idx, axis=0) pad_len = segs_per_arm - taken.shape[0] - padded = jnp.pad(taken, [(9, pad_len)] + [(0, 0)] * (taken.ndim - 1)) + padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1)) per_agent.append(padded.reshape(-1)) arr = jnp.stack(per_agent) @@ -159,7 +160,7 @@ def create_obs_processor( """ values = [] - for key in ordered_keys: + for key in sorted(ordered_keys): if key not in obs: continue diff --git a/src/brittle_star_project/evaluation/__init__.py b/src/brittle_star_project/evaluation/__init__.py index 2922d04..e529115 100644 --- a/src/brittle_star_project/evaluation/__init__.py +++ b/src/brittle_star_project/evaluation/__init__.py @@ -7,9 +7,11 @@ from .evaluate_mjx import ( build_eval_rollout_fn, evaluate_checkpoint_mjx, ) +from .evaluate import evaluate_policy from .policy import PolicyAgent, ControlPolicy from .rollout import rollout_headless, rollout_viewer, EpisodeResult from .video import record_episode, create_evaluation_dir, save_evaluation_metadata +from .eval_env_builder import EvalEnvBundle, build_eval_env __all__ = [ # checkpoint loading @@ -22,6 +24,8 @@ __all__ = [ "append_checkpoint_eval_row", "build_eval_rollout_fn", "evaluate_checkpoint_mjx", + # CPU evaluation + "evaluate_policy", # policy "PolicyAgent", "ControlPolicy", @@ -33,4 +37,7 @@ __all__ = [ "record_episode", "create_evaluation_dir", "save_evaluation_metadata", + # env builder + "EvalEnvBundle", + "build_eval_env", ] diff --git a/src/brittle_star_project/evaluation/eval_env_builder.py b/src/brittle_star_project/evaluation/eval_env_builder.py new file mode 100644 index 0000000..6d45966 --- /dev/null +++ b/src/brittle_star_project/evaluation/eval_env_builder.py @@ -0,0 +1,176 @@ +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path + +import jax.numpy as jnp +import numpy as np +import yaml +from omegaconf import OmegaConf + +from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory +from brittle_star_project.environment.env_config import MorphMode, MorphologyConfig +from brittle_star_project.environment.obs_processing import create_obs_processor +from brittle_star_project.environment.padded_obs_wrapper import compute_padding_masks +from brittle_star_project.evaluation.checkpoint import TrainingConfig +from brittle_star_project.evaluation.policy import PolicyAgent +from brittle_star_project.MLPs.adjancency_builder import build_adjacency + + +@dataclass +class EvalEnvBundle: + """Everything needed to run a headless evaluation episode.""" + + env: BrittleStarEnv + policy: PolicyAgent + action_low: np.ndarray | None + action_high: np.ndarray | None + action_mask: np.ndarray | None + segments_per_arm: list[int] + num_active_arms: int + architecture: str + + +def build_eval_env( + *, + model_path: Path, + training: TrainingConfig, + metadata: dict, + morphology_override_path: Path | str | None = None, +) -> EvalEnvBundle: + """Build environment + policy for evaluation, optionally with a morphology override.""" + + # 1. Determine environment morphology + if morphology_override_path is not None: + override_path = Path(morphology_override_path) + if not override_path.exists(): + raise FileNotFoundError(f"Could not find morphology override YAML at {override_path}") + with open(override_path, "r") as f: + override_dict = yaml.safe_load(f) + env_morphology = OmegaConf.to_object( + OmegaConf.merge(OmegaConf.structured(MorphologyConfig), override_dict) + ) + # Force morph_mode to be inherited from training since it's baked into weights + env_morphology.morph_mode = training.morphology.morph_mode + else: + env_morphology = training.morphology + + # 2. Build obs_processor with TRAINING morphology padding masks always + padding_masks = compute_padding_masks( + segments_per_arm=env_morphology.segments_per_arm, + reference_segments_per_arm=training.morphology.segments_per_arm, + ) + + training_segs_per_arm = jnp.array(training.morphology.segments_per_arm) + + needed_copies = 0 + agent_indices = [0, 1, 2, 3, 4] + match training.morphology.morph_mode: + case MorphMode.CENTRALIZED: + needed_copies = 1 + case MorphMode.FULLY_CONNECTED | MorphMode.RING: + agent_mask = training_segs_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0].tolist() + needed_copies = jnp.where(training_segs_per_arm > 0, 1, 0).sum().item() + case MorphMode.SEGMENT: + agent_mask = training_segs_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0].tolist() + needed_copies = ( + training_segs_per_arm.sum() + jnp.where(training_segs_per_arm > 0, 1, 0).sum() + ).item() + + num_arms_training = jnp.where(training_segs_per_arm > 0, 1, 0).sum().item() + + obs_processor = create_obs_processor( + bounds_dict=training.obs_bounds.to_bounds_dict(), + padding_masks=padding_masks, + needed_copies=needed_copies, + num_arms=num_arms_training, + morph_mode=training.morphology.morph_mode, + segments_per_arm=env_morphology.segments_per_arm, + agent_indices=agent_indices, + ) + + # 3. Build environment + backend = Backend.MJC + factory = BrittleStarEnvFactory() + raw_env = factory.create_environment( + backend, + env_morphology, + training.arena, + training.environment, + ) + env = BrittleStarEnv( + raw_env, + backend=backend, + config=training.environment, + morphology_config=env_morphology, + ) + + # Calculate the action dimension the model was trained with + training_total_actions = sum(training.morphology.segments_per_arm) * 2 + trained_action_dim = training_total_actions // needed_copies + + # 4. Load policy + message_passing_steps = (metadata.get("architecture", {}) or {}).get("message_passing_steps") + if message_passing_steps is None: + message_passing_steps = 4 + message_passing_steps = int(message_passing_steps) + + adj_matrix = None + if training.morphology.morph_mode != MorphMode.CENTRALIZED: + adj_matrix = build_adjacency( + training.morphology.segments_per_arm, training.morphology.morph_mode + ) + + override_segs = env_morphology.segments_per_arm + if training.morphology.morph_mode in (MorphMode.FULLY_CONNECTED, MorphMode.RING): + for i, segs in enumerate(override_segs): + if segs == 0 and i < adj_matrix.shape[0]: + adj_matrix = adj_matrix.at[i, :].set(0) + adj_matrix = adj_matrix.at[:, i].set(0) + elif training.morphology.morph_mode == MorphMode.SEGMENT: + for i, segs in enumerate(override_segs): + if segs == 0 and i < num_arms_training: + adj_matrix = adj_matrix.at[i, :].set(0) + adj_matrix = adj_matrix.at[:, i].set(0) + + idx = 0 + for arm_idx, seg_count in enumerate(training.morphology.segments_per_arm): + if override_segs[arm_idx] == 0: + for i in range(seg_count): + seg_node = num_arms_training + idx + i + if seg_node < adj_matrix.shape[0]: + adj_matrix = adj_matrix.at[seg_node, :].set(0) + adj_matrix = adj_matrix.at[:, seg_node].set(0) + idx += seg_count + + policy = PolicyAgent.from_checkpoint( + model_path, + action_dim=trained_action_dim, + obs_processor=obs_processor, + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, + ) + + # 5. Build action clipping and masks + action_mask = np.asarray(padding_masks["mask_2x"]) + + action_space = getattr(raw_env, "action_space", None) + action_low = ( + None if action_space is None else np.asarray(action_space.low, dtype=np.float32).ravel() + ) + action_high = ( + None if action_space is None else np.asarray(action_space.high, dtype=np.float32).ravel() + ) + + return EvalEnvBundle( + env=env, + policy=policy, + action_low=action_low, + action_high=action_high, + action_mask=action_mask, + segments_per_arm=env_morphology.segments_per_arm, + num_active_arms=sum(1 for s in env_morphology.segments_per_arm if s > 0), + architecture=env_morphology.morph_mode.name, + ) diff --git a/src/brittle_star_project/evaluation/evaluate.py b/src/brittle_star_project/evaluation/evaluate.py new file mode 100644 index 0000000..9d62880 --- /dev/null +++ b/src/brittle_star_project/evaluation/evaluate.py @@ -0,0 +1,58 @@ +"""MJC-based (CPU) checkpoint evaluation. + +This module provides the CPU-bound evaluation path using the standard MJC backend. +It is primarily used by the `evaluate_checkpoints` CLI to compute metrics and +render videos. +""" + +from pathlib import Path + +import numpy as np + +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from brittle_star_project.environment.obs_processing import create_obs_processor +from brittle_star_project.evaluation.policy import PolicyAgent +from brittle_star_project.evaluation.rollout import EpisodeResult, rollout_headless + + +def evaluate_policy( + env: BrittleStarJaxEnvWrapper, + policy_path: str | Path, + seed: int, + max_steps: int, +) -> EpisodeResult: + """Evaluate a trained policy in a CPU-bound environment. + + Args: + env: Initialised CPU environment (MJC backend). + policy_path: Path to the `.cleanrl_model` weights file. + seed: Random seed for environment reset. + max_steps: Maximum number of control steps. + + Returns: + Structured result containing return, length, and distance metrics. + """ + obs_processor = create_obs_processor( + bounds_dict=env.cfg.obs_bounds.to_bounds_dict(), + padding_masks=env.padding_masks, + ) + + action_dim = env.single_action_space.shape[0] + + policy = PolicyAgent.from_checkpoint( + model_path=Path(policy_path), + action_dim=action_dim, + obs_processor=obs_processor, + ) + + action_low = np.asarray(env.single_action_space.low, dtype=np.float32) + action_high = np.asarray(env.single_action_space.high, dtype=np.float32) + + return rollout_headless( + env=env, + policy=policy, + seed=seed, + max_steps=max_steps, + action_low=action_low, + action_high=action_high, + ) diff --git a/src/brittle_star_project/evaluation/evaluate_mjx.py b/src/brittle_star_project/evaluation/evaluate_mjx.py index 317fa98..f557d2d 100644 --- a/src/brittle_star_project/evaluation/evaluate_mjx.py +++ b/src/brittle_star_project/evaluation/evaluate_mjx.py @@ -12,7 +12,7 @@ The key functions are: - `evaluate_checkpoint_mjx` — runs that function for a given set of parameters and returns a typed `CheckpointEvalResult`. - `append_checkpoint_eval_row` — persists the result to the run's - ``metrics/checkpoint_evaluation.csv``, migrating old schemas automatically. + `metrics/checkpoint_evaluation.csv`, migrating old schemas automatically. """ from __future__ import annotations @@ -62,20 +62,20 @@ def build_eval_rollout_fn( All outputs are JAX arrays. Convert to Python scalars before logging. Args: - env: The training environment wrapper. Must expose ``env.raw`` with - ``reset`` and ``step`` methods compatible with ``jax.vmap``. + env: The training environment wrapper. Must expose `env.raw` with + `reset` and `step` methods compatible with `jax.vmap`. obs_processor: Observation normalisation / padding callable, as - returned by ``create_obs_processor``. - sensor_apply: The sensor network's ``apply`` method (JIT-compiled). - actor_apply: The actor network's ``apply`` method (JIT-compiled). + returned by `create_obs_processor`. + sensor_apply: The sensor network's `apply` method (JIT-compiled). + actor_apply: The actor network's `apply` method (JIT-compiled). message_passer_apply: Optional message-passing module apply method. When provided, it is applied between the sensor and actor, using - ``params["message_passer_params"]``. - action_low: Per-joint action lower bound (JAX array, shape ``(action_dim,)``). - action_high: Per-joint action upper bound (JAX array, shape ``(action_dim,)``). + `params["message_passer_params"]`. + action_low: Per-joint action lower bound (JAX array, shape `(action_dim,)`). + action_high: Per-joint action upper bound (JAX array, shape `(action_dim,)`). reward_fn: Shaped reward function with signature - ``reward_fn(env_state, next_env_state) -> jnp.ndarray``. - Typically the module-level ``reward_fn`` from ``PPOTrainer``. + `reward_fn(env_state, next_env_state) -> jnp.ndarray`. + Typically, the module-level `reward_fn` from `PPOTrainer`. Returns: A JIT-compiled callable that runs one deterministic evaluation episode. @@ -217,9 +217,9 @@ def append_checkpoint_eval_row( trained_timesteps: int, result: CheckpointEvalResult, ) -> Path: - """Append one evaluation row to ``/metrics/checkpoint_evaluation.csv``. + """Append one evaluation row to `/metrics/checkpoint_evaluation.csv`. - Creates the file (including the ``metrics/`` directory) if it does not yet + Creates the file (including the `metrics/` directory) if it does not yet exist. Migrates the file to the current schema if the header has changed. Args: diff --git a/src/brittle_star_project/evaluation/policy.py b/src/brittle_star_project/evaluation/policy.py index 5b12d07..00c1e06 100644 --- a/src/brittle_star_project/evaluation/policy.py +++ b/src/brittle_star_project/evaluation/policy.py @@ -7,6 +7,7 @@ import jax import jax.numpy as jnp import numpy as np +from brittle_star_project.MLPs.routing import apply_per_node from brittle_star_project.evaluation.checkpoint import load_params @@ -147,22 +148,12 @@ class PolicyAgent: obs_processor=obs_processor, ) - def _apply_per_node(self, net, params, x): - # params: (nodes, ...) - # x: (batch, nodes, feat) - - def apply_single_node(p, x_node): - # x_node: (batch, feat) - return jax.vmap(lambda xi: net.apply(p, xi))(x_node) - - return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x) - def act(self, *, observations: dict[str, Any]) -> np.ndarray: """Return deterministic action (actor mean, no exploration noise).""" batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations) obs = self._obs_processor(batched_obs) - hidden = self._apply_per_node(self._sensor, self._params["sensor_params"], obs) + hidden = apply_per_node(self._sensor.apply, self._params["sensor_params"], obs) if self._message_passer is not None: mp_params = self._params.get("message_passer_params") @@ -172,6 +163,6 @@ class PolicyAgent: ) hidden = jax.vmap(lambda x: self._message_passer.apply(mp_params, x))(hidden) - mean, _log_std = self._apply_per_node(self._actor, self._params["actor_params"], hidden) + mean, _log_std = apply_per_node(self._actor.apply, self._params["actor_params"], hidden) return np.asarray(mean, dtype=np.float32).ravel() diff --git a/src/brittle_star_project/evaluation/rollout.py b/src/brittle_star_project/evaluation/rollout.py index 982a564..f292138 100644 --- a/src/brittle_star_project/evaluation/rollout.py +++ b/src/brittle_star_project/evaluation/rollout.py @@ -17,6 +17,7 @@ class EpisodeResult: length: int reached_target: bool final_xy_dist: float | None + initial_target_distance: float | None def _get_observations(state: Any) -> dict[str, Any] | None: @@ -61,6 +62,7 @@ def rollout_headless( ep_return = 0.0 observations = _get_observations(state) prev_dist = _get_xy_distance_to_target(observations) if observations else None + initial_target_distance = prev_dist reached_target = _target_reached(state=state) steps = 0 @@ -91,6 +93,7 @@ def rollout_headless( length=steps, reached_target=reached_target, final_xy_dist=final_dist, + initial_target_distance=initial_target_distance, ) diff --git a/src/brittle_star_project/evaluation/video.py b/src/brittle_star_project/evaluation/video.py index 174326b..b4b44db 100644 --- a/src/brittle_star_project/evaluation/video.py +++ b/src/brittle_star_project/evaluation/video.py @@ -97,10 +97,10 @@ def record_episode( data = state.mj_data renderer = mujoco.Renderer(model, width=width, height=height) - ep_return = 0.0 observations = _get_observations(state) prev_dist = _get_xy_distance_to_target(observations) if observations else None + initial_dist = prev_dist reached_target = _target_reached(state=state) frames = [] @@ -145,4 +145,5 @@ def record_episode( length=steps, reached_target=reached_target, final_xy_dist=final_dist, + initial_target_distance=initial_dist, ) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 76bb128..98d4045 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -23,6 +23,7 @@ from brittle_star_project.evaluation.evaluate_mjx import ( build_eval_rollout_fn, evaluate_checkpoint_mjx, ) +from brittle_star_project.MLPs.routing import apply_per_node from brittle_star_project.MLPs.mlps import ( Actor, AgentParams, @@ -71,7 +72,7 @@ def _get_action_and_value_noise( action_high, ): # (B, n_nodes, feat) - hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs) + hidden = apply_per_node(sensor.apply, agent_state.params["sensor_params"], next_obs) if message_passer is not None: params = agent_state.params["message_passer_params"] @@ -82,7 +83,7 @@ def _get_action_and_value_noise( feature_extractor, agent_state.params["feature_extractor_params"], next_obs ) - mean, log_std = apply_per_node(actor, agent_state.params["actor_params"], hidden) + mean, log_std = apply_per_node(actor.apply, agent_state.params["actor_params"], hidden) log_std = jnp.clip(log_std, -5, 2) key, subkey = jax.random.split(key) noise = jax.random.normal(subkey, shape=mean.shape) @@ -272,17 +273,6 @@ def _step_env_wrapped( ) -def apply_per_node(net, params, x): - # params: (nodes, ...) - # x: (batch, nodes, feat) - - def apply_single_node(p, x_node): - # x_node: (batch, feat) - return jax.vmap(lambda xi: net.apply(p, xi))(x_node) - - return jax.vmap(apply_single_node, in_axes=(0, 1), out_axes=1)(params, x) - - def apply_shared(net, params, x): # x: (batch, nodes, feat) # If the critic expects a single vector per environment: @@ -516,10 +506,10 @@ class PPOTrainer: ) def apply_sensor(p, x): - return apply_per_node(self.sensor, p, x) + return apply_per_node(self.sensor.apply, p, x) def apply_actor(p, x): - return apply_per_node(self.actor, p, x) + return apply_per_node(self.actor.apply, p, x) def apply_critic(p, x): return apply_shared(self.critic, p, x) @@ -903,8 +893,8 @@ class PPOTrainer: self._eval_fn = build_eval_rollout_fn( env=self.env, obs_processor=self.obs_processor, - sensor_apply=lambda p, x: apply_per_node(self.sensor, p, x), - actor_apply=lambda p, x: apply_per_node(self.actor, p, x), + sensor_apply=lambda p, x: apply_per_node(self.sensor.apply, p, x), + actor_apply=lambda p, x: apply_per_node(self.actor.apply, p, x), message_passer_apply=( None if self.message_passer is None else self.message_passer.apply ), diff --git a/tests/test_evaluation.py b/tests/test_evaluation.py index bbb03f8..d8be381 100644 --- a/tests/test_evaluation.py +++ b/tests/test_evaluation.py @@ -14,6 +14,7 @@ from brittle_star_project.environment.env_config import ( ArenaConfig, EnvConfig, ObservationBoundsConfig, + MorphMode, ) from brittle_star_project.environment.env_types import Task @@ -112,3 +113,106 @@ def test_load_metadata_with_override(tmp_path: Path): non_existent = tmp_path / "missing.yaml" with pytest.raises(FileNotFoundError, match="Could not find metadata YAML at"): load_metadata(model_path, metadata_override_path=non_existent) + + +@pytest.fixture +def mock_training_config(): + return TrainingConfig( + morphology=MorphologyConfig( + segments_per_arm=[1, 1, 1, 1, 1], morph_mode=MorphMode.CENTRALIZED + ), + arena=ArenaConfig(), + environment=EnvConfig(), + obs_bounds=ObservationBoundsConfig(), + ) + + +@pytest.fixture +def mock_metadata(): + return {"architecture": {"message_passing_steps": 2}} + + +def test_build_eval_env_training_morphology(tmp_path, mock_training_config, mock_metadata): + from brittle_star_project.evaluation.eval_env_builder import build_eval_env + from unittest.mock import patch + + model_path = tmp_path / "model.flax" + patch_target = "brittle_star_project.evaluation.eval_env_builder.PolicyAgent.from_checkpoint" + with patch(patch_target) as mock_agent: + mock_agent.return_value = "mock_policy" + bundle = build_eval_env( + model_path=model_path, + training=mock_training_config, + metadata=mock_metadata, + morphology_override_path=None, + ) + assert bundle.segments_per_arm == [1, 1, 1, 1, 1] + assert bundle.num_active_arms == 5 + assert bundle.architecture == "CENTRALIZED" + assert bundle.policy == "mock_policy" + + +def test_build_eval_env_override_morphology(tmp_path, mock_training_config, mock_metadata): + from brittle_star_project.evaluation.eval_env_builder import build_eval_env + from unittest.mock import patch + + model_path = tmp_path / "model.flax" + override_path = tmp_path / "override.yaml" + override_path.write_text(yaml.dump({"segments_per_arm": [1, 0, 1, 0, 1]})) + + with patch("brittle_star_project.evaluation.eval_env_builder.PolicyAgent.from_checkpoint"): + bundle = build_eval_env( + model_path=model_path, + training=mock_training_config, + metadata=mock_metadata, + morphology_override_path=override_path, + ) + assert bundle.segments_per_arm == [1, 0, 1, 0, 1] + assert bundle.num_active_arms == 3 + # Should be smaller than 5*N + assert sum(bundle.action_mask) < len(bundle.action_mask) + + +def test_build_eval_env_action_mask_shape(tmp_path, mock_training_config, mock_metadata): + from brittle_star_project.evaluation.eval_env_builder import build_eval_env + from unittest.mock import patch + + model_path = tmp_path / "model.flax" + override_path = tmp_path / "override.yaml" + override_path.write_text(yaml.dump({"segments_per_arm": [1, 0, 1, 0, 0]})) + + with patch("brittle_star_project.evaluation.eval_env_builder.PolicyAgent.from_checkpoint"): + bundle = build_eval_env( + model_path=model_path, + training=mock_training_config, + metadata=mock_metadata, + morphology_override_path=override_path, + ) + # For each segment with P-control, there's 2 actions (pitch and yaw). + # Total segments = 5 -> 10 actions for training. + assert len(bundle.action_mask) == 10 + # Active segments = 2 -> 4 actions active. + assert sum(bundle.action_mask) == 4 + + +def test_build_eval_env_morph_mode_inherited(tmp_path, mock_training_config, mock_metadata): + from brittle_star_project.evaluation.eval_env_builder import build_eval_env + from brittle_star_project.environment.env_config import MorphMode + from unittest.mock import patch + + model_path = tmp_path / "model.flax" + override_path = tmp_path / "override.yaml" + # No morph_mode in the override YAML + override_path.write_text(yaml.dump({"segments_per_arm": [1, 0, 1, 0, 1]})) + + # Change training config to be RING + mock_training_config.morphology.morph_mode = MorphMode.RING + + with patch("brittle_star_project.evaluation.eval_env_builder.PolicyAgent.from_checkpoint"): + bundle = build_eval_env( + model_path=model_path, + training=mock_training_config, + metadata=mock_metadata, + morphology_override_path=override_path, + ) + assert bundle.architecture == "RING" diff --git a/tests/test_network_shapes.py b/tests/test_network_shapes.py index 7bec744..ce189f4 100644 --- a/tests/test_network_shapes.py +++ b/tests/test_network_shapes.py @@ -36,12 +36,14 @@ def test_centralized_forward_pass_with_padding(): ) global_state = obs_processor(amputated_obs) - # 40 + 40 + 20 + padding = 145 dimensions - assert global_state.shape == (batch_size, 1, 145), ( - f"Expected global state shape (2, 1, 145), got {global_state.shape}" + # joint_position: 5 arms × 8 joints (padded) = 40 + # joint_velocity: 5 arms × 8 joints (padded) = 40 + # segment_contact: 5 arms × 4 segs (padded) = 20 + # Total = 100 (no disk or direction keys supplied) + assert global_state.shape == (batch_size, 1, 100), ( + f"Expected global state shape (2, 1, 100), got {global_state.shape}" ) - # 4. Initialize dummy networks (40 actuators for the max morphology output) actor = Actor(action_dim=40) critic = OneDenseLayerMLP() # Acts as the centralized critic diff --git a/tests/test_obs_processor.py b/tests/test_obs_processor.py index 17c5c8f..7008648 100644 --- a/tests/test_obs_processor.py +++ b/tests/test_obs_processor.py @@ -7,17 +7,14 @@ from brittle_star_project.environment.env_config import MorphMode, ObservationBo obs_bounds = ObservationBoundsConfig().to_bounds_dict() -""" -Test for obs_processor. - -Centralized: 40 features per agent: - disk_z_tilt → scalar → reshaped to (1,) → 1 feat - joint_actuator_force → 8 joints padded to 8 → 8 feat - joint_position → 8 joints padded to 8 → 8 feat - joint_velocity → 8 joints padded to 8 → 8 feat - robot_direction_to_target→ (x, y) → 2 feat - segment_contact → 4 segs, pre-padded by 9 → 13 feat (9 leading + 4) -""" +# Features per decentralized agent (one arm's data): +# disk_z_tilt → scalar → 1 feat +# joint_actuator_force → 4 segs × 2 joints → 8 feat +# joint_position → 4 segs × 2 joints → 8 feat +# joint_velocity → 4 segs × 2 joints → 8 feat +# robot_direction_to_target→ (x, y) → 2 feat +# segment_contact → 4 segs → 4 feat +# Total per agent: 1+8+8+8+2+4 = 31 NUM_ARMS = 5 SEGS_PER_ARM = 4 # healthy segments per arm @@ -28,7 +25,17 @@ SEGS_DAMAGED = [4, 4, 4, 4, 0] # arm 4 fully disabled SEGS_DAMAGED_2 = [4, 0, 4, 2, 4] # arm 3 fully disabled AGENT_INDICES = [0, 1, 2, 3, 4] -FEAT_PER_AGENT = 1 + 8 + 8 + 8 + 2 + 13 # = 40 +FEAT_PER_AGENT = 1 + 8 + 8 + 8 + 2 + 4 # = 31 + +# Centralized flattening (needed_copies=1, one copy of global features): +# disk_z_tilt → repeated once → 1 feat +# joint_actuator_force → 5 arms × 8 joints → 40 feat +# joint_position → 5 arms × 8 joints → 40 feat +# joint_velocity → 5 arms × 8 joints → 40 feat +# robot_direction_to_target→ repeated once → 2 feat +# segment_contact → 5 arms × 4 segs → 20 feat +# Total: 1+40+40+40+2+20 = 143 +FEAT_CENTRALIZED = 1 + 40 + 40 + 40 + 2 + 20 # = 143 def make_obs(segs_per_arm: list[int]) -> dict: @@ -73,10 +80,8 @@ def test_centralized_no_damage(): obs = batch_obs(obs) global_state = proc(obs) - # shape test - assert global_state.shape == (1, 1, 188) - - # TODO: more? + # Centralized: 5 agents flattened into 1 → shape (1, 1, 155) + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) def test_centralized_damaged_1_arm(): @@ -86,7 +91,7 @@ def test_centralized_damaged_1_arm(): global_state = proc(obs) # shape test - assert global_state.shape == (1, 1, 188) + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) def test_centralized_damaged_2_arms(): @@ -96,7 +101,7 @@ def test_centralized_damaged_2_arms(): global_state = proc(obs) # shape test - assert global_state.shape == (1, 1, 188) + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) def test_decentralized_fully_connected_no_damage():