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..45bdf63 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)**. @@ -30,6 +45,9 @@ For detailed instructions on how to use the project, please refer to the **[API 3. **Simulate a trained model:** See [Simulation & Evaluation](docs/api/simulation.md). +4. **Compare fault tolerance of models:** + See [Checkpoint & Model Evaluation](docs/api/evaluation.md) + ## HPC See **[docs/HPC.md](docs/HPC.md)** for the full guide, including environment setup, cluster selection, interactive debugging, and job submission. diff --git a/configs/architecture/decentralized.yaml b/configs/architecture/decentralized.yaml index 4b9fffb..70b0fb2 100644 --- a/configs/architecture/decentralized.yaml +++ b/configs/architecture/decentralized.yaml @@ -26,7 +26,7 @@ critic: activation: "tanh" # Synchronous message-passing rounds per control step -message_passing_steps: 1 +message_passing_steps: 4 # Connectivity topology (e.g., ring, fully_connected) -topology_type: "ring" +topology_type: "fully_connected" 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/dir_loc_further.yaml b/configs/environment/dir_loc_further.yaml new file mode 100644 index 0000000..236f0c5 --- /dev/null +++ b/configs/environment/dir_loc_further.yaml @@ -0,0 +1,2 @@ +simulation_time: 50000.0 +target_distance: 3.0 \ No newline at end of file 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 new file mode 100644 index 0000000..a7c2a00 --- /dev/null +++ b/configs/evaluation/default.yaml @@ -0,0 +1,8 @@ +# Default Evaluation Configuration +# Settings used for checkpoint evaluation during training. + +evaluate_checkpoints: false +# Max number of control steps during evaluation rollout. +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/experiment/long_2arm.yaml b/configs/experiment/long_2arm.yaml new file mode 100644 index 0000000..ae4f4b0 --- /dev/null +++ b/configs/experiment/long_2arm.yaml @@ -0,0 +1,6 @@ +# Testing chicken dinner 4 but further distance. + +exp_name: "long2arm" +seed: 123 +torch_deterministic: true +cuda: true \ No newline at end of file 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/logging/default.yaml b/configs/logging/default.yaml index 1ecf7fd..2f7de8c 100644 --- a/configs/logging/default.yaml +++ b/configs/logging/default.yaml @@ -10,4 +10,4 @@ save_checkpoints: true checkpoint_frequency: 100 upload_final_model: false upload_checkpoints: false -hf_entity: "" +hf_entity: "" \ No newline at end of file diff --git a/configs/main_config.yaml b/configs/main_config.yaml index b7cb620..44451f4 100644 --- a/configs/main_config.yaml +++ b/configs/main_config.yaml @@ -6,6 +6,7 @@ defaults: - brittle_star_config - experiment: base - logging: default + - evaluation: default - ppo: default - architecture: centralized - morphology: 5_arms_full diff --git a/configs/morphology/2_arms_decentralized.yaml b/configs/morphology/2_arms_decentralized.yaml new file mode 100644 index 0000000..9fe3160 --- /dev/null +++ b/configs/morphology/2_arms_decentralized.yaml @@ -0,0 +1,6 @@ +# 2 Arms Morphology Configuration + +segments_per_arm: [4, 0, 4, 0, 0] +use_p_control: true +use_torque_control: false +morph_mode: FULLY_CONNECTED \ No newline at end of file diff --git a/configs/morphology/5_arms_damaged.yaml b/configs/morphology/5_arms_damaged.yaml new file mode 100644 index 0000000..2afd643 --- /dev/null +++ b/configs/morphology/5_arms_damaged.yaml @@ -0,0 +1,6 @@ +# 5 Arms Full Morphology Configuration +# Baseline 5-arm brittle star. + +segments_per_arm: [4, 4, 0, 4, 4] +use_p_control: true +use_torque_control: false diff --git a/configs/morphology/5_arms_full_fullconnected.yaml b/configs/morphology/5_arms_full_fullconnected.yaml new file mode 100644 index 0000000..ab08e59 --- /dev/null +++ b/configs/morphology/5_arms_full_fullconnected.yaml @@ -0,0 +1,7 @@ +# 5 Arms Full Morphology Configuration +# Baseline 5-arm brittle star. + +segments_per_arm: [4, 4, 4, 4, 4] +use_p_control: true +use_torque_control: false +morph_mode: FULLY_CONNECTED diff --git a/configs/ppo/chickendinnerwinner.yaml b/configs/ppo/chickendinnerwinner.yaml new file mode 100644 index 0000000..cf5e3d6 --- /dev/null +++ b/configs/ppo/chickendinnerwinner.yaml @@ -0,0 +1,16 @@ +anneal_lr: true +clip_coef: 0.2 +clip_vloss: true +ent_coef: 0.001 +gae_lambda: 0.95 +gamma: 0.99 +learning_rate: 0.0001 +max_grad_norm: 0.5 +norm_adv: true +num_envs: 32 +num_minibatches: 32 +num_steps: 64 +target_kl: 0.02 +total_timesteps: 12288000 +update_epochs: 4 +vf_coef: 1.0 \ No newline at end of file diff --git a/configs/ppo/smoke_test.yaml b/configs/ppo/smoke_test.yaml index 40e4c07..4b7bb84 100644 --- a/configs/ppo/smoke_test.yaml +++ b/configs/ppo/smoke_test.yaml @@ -2,9 +2,9 @@ # Lower timestep count for quick iterations/testing. learning_rate: 0.0005 -total_timesteps: 65536 -num_envs: 512 -num_steps: 128 +total_timesteps: 1024 +num_envs: 32 +num_steps: 32 anneal_lr: true gamma: 0.99 gae_lambda: 0.95 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/README.md b/docs/README.md index 3c76684..521cb28 100644 --- a/docs/README.md +++ b/docs/README.md @@ -1,6 +1,6 @@ # Documentation -## Design & architecture ([`/design`](./design/)) +## Design & architecture (`/design`) If you are interested in the "why did you do it like this?" @@ -11,7 +11,7 @@ If you are interested in the "why did you do it like this?" - [Learning algorithm](./design/learning_algorithm.md): RL techniques, i.e. PPO. - [Reward function](./design/learning_algorithm.md): Goals, fitness tracking, and reward structures. -## API reference ([`/api`](./api/)) +## API reference (`/api`) If you are interested in the "how do I use it?" @@ -19,3 +19,5 @@ If you are interested in the "how do I use it?" - [Tracking & Monitoring](./api/tracking.md): Setting up WandB and TensorBoard to monitor runs. - [Simulation](./api/simulation.md): Visualizing and evaluating models. - [Environment](./api/environment.md): MuJoCo environment interaction and configuration. +- [Analysis](./api/analysis.md): Comparing checkpoints and generating plots. +- [Evaluation](./api/evaluation.md): Evaluating checkpoints and comparing fault tolerance. \ No newline at end of file diff --git a/docs/api/analysis.md b/docs/api/analysis.md new file mode 100644 index 0000000..e852295 --- /dev/null +++ b/docs/api/analysis.md @@ -0,0 +1,84 @@ +# Analysis & Plotting Tools + +This guide outlines the tools available for analyzing experimental data and generating poster-quality visualizations for the Brittle Star project. + +## Shared Configuration + +All plotting scripts share a central configuration in `scripts/plots/plot_config.py`. This file defines: +- **Color Palette:** A color-blind friendly, high-contrast palette for different architectures. +- **Typography:** Consistent font sizes and styles tailored for A0 posters. +- **Markers:** Shared visual indicators, such as the ★ used for best performers. + +## Comparison Visualization + +The `scripts/plots/analyze_comparisons.py` script generates grouped bar charts comparing the performance of different architectures across various morphologies. + +### Usage + +Run the script from the root of the project, providing the path to your evaluation CSV: + +```bash +# Basic usage (saves PNG and SVG to runs/evaluation/plots/) +uv run python scripts/plots/analyze_comparisons.py path/to/results.csv + +# Advanced usage for Figma/Poster integration +uv run python scripts/plots/analyze_comparisons.py path/to/results.csv \ + --output_dir docs/assets/plots/ \ + --font_size 30 \ + --fig_width 14 \ + --fig_height 10 +``` + +### CLI Arguments + +- `input_csv`: (Required) Path to the CSV file containing evaluation results. +- `--output_dir`, `-o`: Directory where plots will be saved (default: `runs/evaluation/plots`). +- `--show_titles`: Include titles in the plots. Default is **False**, as titles are typically added natively in design tools like Figma. +- `--font_size`: Base font size in points (default: 28). +- `--fig_width` / `--fig_height`: Physical dimensions of the plot in inches. Match these to your Figma layout to maintain exact font sizes. + +### Outputs + +The script generates four key plots, each saved as both `.png` and `.svg`: +1. **Forward Velocity:** Grouped bar chart (cm/s). +2. **Accumulated Reward:** Mean cumulative reward. +3. **Success Rate:** Target acquisition percentage. +4. **Distance Remaining:** Navigational accuracy. + +--- + +## Convergence Analysis + +The `scripts/plots/analyze_convergence.py` script determines the convergence point of training runs. + +### Usage + +```bash +uv run python scripts/plots/analyze_convergence.py --output_dir runs/convergence/ +``` + +### Configuration + +- **File Mapping:** The script uses hardcoded paths in the `FILE_MAPPING` dictionary. Update these paths to point to your specific run evaluation files. +- **CLI Arguments:** Supports the same `--show_titles`, `--font_size`, and `--fig_width/height` flags as the comparison script. + +### Outputs + +Generates three plots (PNG & SVG): +1. `convergence_comparison`: Grouped horizontal bar chart. +2. `progress_reward_curves`: Line plots of reward over time. +3. `progress_velocity_curves`: Line plots of velocity over time. + +--- + +## Poster Integration (Figma) + +### SVG & Scaling +We recommend using the **SVG** outputs for poster design in Figma: +1. **No Resolution Loss:** SVGs are vector-based and will remain sharp at any size. +2. **Native Text:** Text in the SVG imports as native text layers in Figma. +3. **Exact Font Matching:** To ensure a `28pt` font in the plot matches a `28pt` font in your poster, set the `--fig_width` and `--fig_height` to match the physical dimensions of the plot box in your Figma layout. +4. **Editable:** You can "Ungroup" the SVG in Figma to manually move labels, adjust colors, or tweak individual bars. + +### Image Placeholders +The comparison charts include light-gray square placeholders below the X-axis. These are designed as guides; in Figma, you can drop your morphology renders or illustrations directly on top of these squares. diff --git a/docs/api/evaluation.md b/docs/api/evaluation.md new file mode 100644 index 0000000..45b880e --- /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 & Fault 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/experiments/debug-experiment-10042026/used_variables.md b/experiments/debug-experiment-10042026/used_variables.md deleted file mode 100644 index 92f23a1..0000000 --- a/experiments/debug-experiment-10042026/used_variables.md +++ /dev/null @@ -1,107 +0,0 @@ -## Default envconfig -task: Task = Task.DIRECTED_LOCOMOTION -simulation_time: float = 500.0 -num_physics_steps_per_control_step: int = 10 -time_scale: int = 2 -camera_ids: list[int] = field(default_factory=lambda: [0, 1]) -render_size: tuple[int, int] = (480, 640) -joint_randomization_noise_scale: float = 0.0 -target_distance: float = 3.0 -light_perlin_noise_scale: int = 0 - - -## Default ppoargs -seed: int = 1 -torch_deterministic: bool = True -cuda: bool = True -track: bool = False -checkpoint_frequency: int = 100 -learning_rate: float = 2.5e-4 -anneal_lr: bool = True -gamma: float = 0.99 -gae_lambda: float = 0.95 -update_epochs: int = 4 -norm_adv: bool = True -clip_vloss: bool = True -max_grad_norm: float = 0.5 -target_kl: float | None = None -batch_size: int = 0 -minibatch_size: int = 0 -num_iterations: int = 0 - -## Used config file: (hpc/debug.yaml) -exp_name: "debug-experiment" -seed: 42 -track: true -wandb_project_name: "Let's-find-that-bug" -wandb_entity: "SEL3-2026-Groep-4" -run_dir: "/data/gent/465/vsc46589" -num_envs: 32 -num_steps: 32 -num_minibatches: 32 -total_timesteps: 409600 -num_arms: 2 -cuda: true - -ent_coef: 0.005 -vf_coef: 1.0 -clip_coef: 0.2 - -anneal_lr: true -learning_rate: 0.0003 - -## Arena config: -size: tuple[float, float] = (10.0, 5.0) -sand_ground_color: bool = True -attach_target: bool = True -wall_height: float = 1.5 -wall_thickness: float = 0.1 - -## Morphology: -num_segments_per_arm: int = 4 -use_p_control: bool = True -use_torque_control: bool = False - -## MLPs: -### Sensor & Feature_extractor: -Both with 3 layers of 300 neurons per layer. - -class GenericDenseLayersWithActivation(nn.Module): - layer_sizes: Sequence[int] = field(default_factory=lambda: [64, 64]) - activation: Callable = nn.tanh - - @nn.compact - def __call__(self, x): - for size in self.layer_sizes: - x = nn.Dense(size, kernel_init=orthogonal(jnp.sqrt(2)))(x) - x = self.activation(x) - return x - -### Actor: -class Actor(nn.Module): - action_dim: int - @nn.compact - def __call__(self, x): - mean = nn.Dense(self.action_dim, kernel_init=orthogonal(0.01), bias_init=constant(0.0))(x) - log_std = self.param("log_std", nn.initializers.zeros, (self.action_dim,)) - return mean, log_std - -### Critic: -class OneDenseLayerMLP(nn.Module): - @nn.compact - def __call__(self, x): - return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x) - -### Observations: -_ALLOWED_OBS_KEYS = { - "joint_position", - "joint_velocity", - "joint_actuator_force", - "actuator_force", - "disk_position", - "disk_rotation", - "disk_linear_velocity", - "disk_angular_velocity", - "unit_xy_direction_to_target", - "xy_distance_to_target", -} \ No newline at end of file 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/plots/analyze_comparisons.py b/scripts/plots/analyze_comparisons.py new file mode 100644 index 0000000..3ac6e8a --- /dev/null +++ b/scripts/plots/analyze_comparisons.py @@ -0,0 +1,445 @@ +""" +Poster Comparison Visualizations + +This script generates a Forward Velocity plot and three secondary plots (Accumulated Reward, Success +Rate, Distance Remaining). +""" + +import os +import pandas as pd +import numpy as np +import matplotlib.pyplot as plt + +from plot_config import ( + COLORS, + apply_style, + BEST_PERFORMER_MARKER, + BEST_PERFORMER_TEXT, + BEST_PERFORMER_COLOR, + create_common_parser, + LEGEND_KWARGS, +) + + +def load_and_preprocess_data(filepath): + """Loads CSV and prepares the metrics for plotting.""" + df = pd.read_csv(filepath) + + # Ensure success rate can be averaged numerically + if "reached_target" in df.columns: + df["reached_target"] = df["reached_target"].astype(int) + + return df + + +def _add_square_placeholders(ax, x_positions, labels): + """Adds square placeholders for images below the x-axis.""" + for x, label in zip(x_positions, labels): + # Create a roughly square rectangle in a mix of data/axes coords + # Shifted down to avoid overlapping with x-tick labels + rect = plt.Rectangle( + (x - 0.25, -0.40), + 0.5, + 0.18, + transform=ax.get_xaxis_transform(), + facecolor="#F0F0F0", + edgecolor="#A9A9A9", + linestyle="--", + zorder=1, + clip_on=False, + ) + ax.add_patch(rect) + ax.text( + x, + -0.31, + f"[ Insert {label}\nImage ]", + transform=ax.get_xaxis_transform(), + ha="center", + va="center", + fontsize=10, + color="#888888", + zorder=2, + ) + + +def plot_grouped_bar( + df, + metric_col, + ylabel, + title, + output_filename, + output_dir, + higher_is_better=True, + show_titles=False, + figsize=(12, 8), +): + """Generates and saves a highly customized grouped bar chart (grouped by Morphology).""" + grouped = ( + df.groupby(["num_active_arms", "architecture"])[metric_col] + .agg(["mean", "std"]) + .reset_index() + ) + morphologies = sorted(grouped["num_active_arms"].unique(), reverse=True) + architectures = grouped["architecture"].unique() + + fig, ax = plt.subplots(figsize=figsize) + bar_width = 0.35 + x_indices = np.arange(len(morphologies)) + all_bars = {} + all_means = [] + + for i, arch in enumerate(architectures): + arch_data = grouped[grouped["architecture"] == arch] + means = [ + arch_data[arch_data["num_active_arms"] == m]["mean"].values[0] + if not arch_data[arch_data["num_active_arms"] == m].empty + else 0 + for m in morphologies + ] + stds = [ + arch_data[arch_data["num_active_arms"] == m]["std"].values[0] + if not arch_data[arch_data["num_active_arms"] == m].empty + else 0 + for m in morphologies + ] + all_means.extend(means) + x_pos = x_indices + (i * bar_width) - (bar_width / 2 if len(architectures) == 2 else 0) + color = COLORS.get(arch, "#888888") + clean_label = arch.replace("_", " ").title() + bars = ax.bar( + x_pos, + means, + bar_width, + yerr=stds, + label=clean_label, + color=color, + capsize=8, + error_kw={"elinewidth": 2, "alpha": 0.7}, + ) + all_bars[arch] = (x_pos, means, stds, bars) + + for m_idx, m in enumerate(morphologies): + m_means = {arch: all_bars[arch][1][m_idx] for arch in architectures} + best_arch = ( + max(m_means, key=m_means.get) if higher_is_better else min(m_means, key=m_means.get) + ) + best_x = all_bars[best_arch][0][m_idx] + best_y = all_bars[best_arch][1][m_idx] + best_std = all_bars[best_arch][2][m_idx] + offset = best_std + (abs(max(m_means.values())) * 0.05) if m_means.values() else 0 + ax.text( + best_x, + best_y + offset, + BEST_PERFORMER_TEXT, + ha="center", + va="bottom", + fontsize=28, + color=BEST_PERFORMER_COLOR, + ) + + # Aesthetics + ax.set_ylabel(ylabel, labelpad=15) + if show_titles: + ax.set_title(title, pad=25, fontweight="bold") + + x_ticks_pos = ( + x_indices + + (bar_width / 2 if len(architectures) % 2 == 0 else 0) + - (bar_width / 2 if len(architectures) == 2 else 0) + ) + ax.set_xticks(x_ticks_pos) + ax.set_xticklabels([f"{m} Arms" for m in morphologies]) + ax.tick_params(axis="x", pad=25) # More padding for the squares + + # X-axis at zero + ax.axhline(0, color="black", linewidth=1.5) + ax.spines["bottom"].set_visible(False) + + # Y-axis limits explicitly including 0 + if all_means: + min_val = min([*all_means, 0]) + max_val = max([*all_means, 0]) + margin = (max_val - min_val) * 0.15 if max_val != min_val else 0.1 + ax.set_ylim(min_val - margin, max_val + margin * 1.5) # Extra top margin for stars + # Format y-ticks to not have excessive decimals, include 0 + ticks = ( + [min_val, max_val] + if min_val == 0 and max_val == 0 + else sorted(list(set([min_val, 0, max_val]))) + ) + ax.set_yticks(ticks) + ax.yaxis.set_major_formatter( + plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") + ) + + _add_square_placeholders(ax, x_ticks_pos, [f"{m} Arms" for m in morphologies]) + + # Add custom legend entry for best performer + ax.plot( + [], + [], + marker=BEST_PERFORMER_MARKER, + color="w", + markerfacecolor=BEST_PERFORMER_COLOR, + markersize=15, + label="Best Performance", + ls="", + ) + ax.legend(**LEGEND_KWARGS, ncol=len(architectures) + 1) + ax.set_facecolor("white") + fig.patch.set_facecolor("white") + + os.makedirs(output_dir, exist_ok=True) + base_path = os.path.join(output_dir, os.path.splitext(output_filename)[0]) + plt.savefig(f"{base_path}.png", dpi=300, bbox_inches="tight") + plt.savefig(f"{base_path}.svg", format="svg", bbox_inches="tight") + plt.close() + + +def plot_grouped_bar_alt( + df, + metric_col, + ylabel, + title, + output_filename, + output_dir, + higher_is_better=True, + show_titles=False, + figsize=(12, 8), +): + """Generates and saves a highly customized grouped bar chart (grouped by Architecture).""" + grouped = ( + df.groupby(["architecture", "num_active_arms"])[metric_col] + .agg(["mean", "std"]) + .reset_index() + ) + architectures = sorted(grouped["architecture"].unique()) + morphologies = sorted(grouped["num_active_arms"].unique(), reverse=True) + + fig, ax = plt.subplots(figsize=figsize) + bar_width = 0.8 / len(morphologies) + x_indices = np.arange(len(architectures)) + all_bars = {} + all_means = [] + + for i, m in enumerate(morphologies): + m_data = grouped[grouped["num_active_arms"] == m] + means = [ + m_data[m_data["architecture"] == arch]["mean"].values[0] + if not m_data[m_data["architecture"] == arch].empty + else 0 + for arch in architectures + ] + stds = [ + m_data[m_data["architecture"] == arch]["std"].values[0] + if not m_data[m_data["architecture"] == arch].empty + else 0 + for arch in architectures + ] + all_means.extend(means) + + # Offset bars based on morphology index + offset = (i - len(morphologies) / 2 + 0.5) * bar_width + x_pos = x_indices + offset + + # We can use a color gradient or different colors for morphologies + # For simplicity, using a colormap + color = plt.cm.viridis(i / max(1, len(morphologies) - 1)) + + bars = ax.bar( + x_pos, + means, + bar_width, + yerr=stds, + label=f"{m} Arms", + color=color, + capsize=4, + error_kw={"elinewidth": 1.5, "alpha": 0.7}, + ) + all_bars[m] = (x_pos, means, stds, bars) + + for a_idx, arch in enumerate(architectures): + a_means = {m: all_bars[m][1][a_idx] for m in morphologies} + best_m = ( + max(a_means, key=a_means.get) if higher_is_better else min(a_means, key=a_means.get) + ) + best_x = all_bars[best_m][0][a_idx] + best_y = all_bars[best_m][1][a_idx] + best_std = all_bars[best_m][2][a_idx] + offset = best_std + (abs(max(a_means.values())) * 0.05) if a_means.values() else 0 + ax.text( + best_x, + best_y + offset, + BEST_PERFORMER_TEXT, + ha="center", + va="bottom", + fontsize=20, + color=BEST_PERFORMER_COLOR, + ) + + # Aesthetics + ax.set_ylabel(ylabel, labelpad=15) + if show_titles: + ax.set_title(title + " (Alt)", pad=25, fontweight="bold") + + ax.set_xticks(x_indices) + ax.set_xticklabels([arch.replace("_", " ").title() for arch in architectures]) + ax.tick_params(axis="x", pad=25) + + # X-axis at zero + ax.axhline(0, color="black", linewidth=1.5) + ax.spines["bottom"].set_visible(False) + + if all_means: + min_val = min([*all_means, 0]) + max_val = max([*all_means, 0]) + margin = (max_val - min_val) * 0.15 if max_val != min_val else 0.1 + ax.set_ylim(min_val - margin, max_val + margin * 1.5) + ticks = ( + [min_val, max_val] + if min_val == 0 and max_val == 0 + else sorted(list(set([min_val, 0, max_val]))) + ) + ax.set_yticks(ticks) + ax.yaxis.set_major_formatter( + plt.FuncFormatter(lambda x, _: f"{x:.2f}" if abs(x) < 10 else f"{x:.0f}") + ) + + # In this alt plot, placeholders might be per architecture + _add_square_placeholders( + ax, x_indices, [arch.replace("_", "\n").title() for arch in architectures] + ) + + ax.plot( + [], + [], + marker=BEST_PERFORMER_MARKER, + color="w", + markerfacecolor=BEST_PERFORMER_COLOR, + markersize=15, + label="Best Performance", + ls="", + ) + ax.legend(**LEGEND_KWARGS, ncol=len(morphologies) + 1) + ax.set_facecolor("white") + fig.patch.set_facecolor("white") + + os.makedirs(output_dir, exist_ok=True) + base_path = os.path.join(output_dir, os.path.splitext(output_filename)[0]) + plt.savefig(f"{base_path}.png", dpi=300, bbox_inches="tight") + plt.savefig(f"{base_path}.svg", format="svg", bbox_inches="tight") + plt.close() + + +if __name__ == "__main__": + parser = create_common_parser(description="Generate comparison poster plots.") + parser.add_argument( + "input_csv", help="Path to the input CSV file containing evaluation results." + ) + args = parser.parse_args() + + INPUT_CSV = args.input_csv + OUTPUT_DIR = args.output_dir + + if not os.path.exists(INPUT_CSV): + print(f"Error: Could not find {INPUT_CSV}. Please ensure the file exists.") + else: + df = load_and_preprocess_data(INPUT_CSV) + print("Data loaded successfully. Generating poster plots...") + + apply_style(font_size=args.font_size) + kwargs = {"show_titles": args.show_titles, "figsize": (args.fig_width, args.fig_height)} + + # Velocity Conversion: m/s to cm/s + if "approx_max_velocity" in df.columns: + df["approx_max_velocity"] = df["approx_max_velocity"] * 100 + + # 1. Primary Plot: Forward Velocity + plot_grouped_bar( + df=df, + metric_col="approx_max_velocity", + ylabel="Max Forward Velocity (cm/s)", + title="Graceful Degradation: Velocity Across Morphologies", + output_filename="poster_plot_velocity.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + plot_grouped_bar_alt( + df=df, + metric_col="approx_max_velocity", + ylabel="Max Forward Velocity (cm/s)", + title="Graceful Degradation: Velocity Across Morphologies", + output_filename="poster_plot_velocity_alt.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + + # 2. Secondary Plot: Accumulated Reward + plot_grouped_bar( + df=df, + metric_col="eval_return", + ylabel="Mean Cumulative Reward", + title="Overall Efficiency Across Morphologies", + output_filename="poster_plot_reward.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + plot_grouped_bar_alt( + df=df, + metric_col="eval_return", + ylabel="Mean Cumulative Reward", + title="Overall Efficiency Across Morphologies", + output_filename="poster_plot_reward_alt.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + + # 3. Secondary Plot: Success Rate + plot_grouped_bar( + df=df, + metric_col="reached_target", + ylabel="Success Rate (%)", + title="Target Acquisition Consistency", + output_filename="poster_plot_success_rate.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + plot_grouped_bar_alt( + df=df, + metric_col="reached_target", + ylabel="Success Rate (%)", + title="Target Acquisition Consistency", + output_filename="poster_plot_success_rate_alt.png", + output_dir=OUTPUT_DIR, + higher_is_better=True, + **kwargs, + ) + + # 4. Secondary Plot: Final Distance Remaining + plot_grouped_bar( + df=df, + metric_col="final_xy_dist", + ylabel="Distance to Target Remaining", + title="Navigational Accuracy (Lower is Better)", + output_filename="poster_plot_distance.png", + output_dir=OUTPUT_DIR, + higher_is_better=False, # For distance, a lower score is better + **kwargs, + ) + plot_grouped_bar_alt( + df=df, + metric_col="final_xy_dist", + ylabel="Distance to Target Remaining", + title="Navigational Accuracy (Lower is Better)", + output_filename="poster_plot_distance_alt.png", + output_dir=OUTPUT_DIR, + higher_is_better=False, + **kwargs, + ) + + print(f"All plots generated in the '{OUTPUT_DIR}/' directory.") diff --git a/scripts/plots/analyze_convergence.py b/scripts/plots/analyze_convergence.py new file mode 100644 index 0000000..2612bb9 --- /dev/null +++ b/scripts/plots/analyze_convergence.py @@ -0,0 +1,335 @@ +""" +Convergence Analysis Script for Poster Visualizations + +This script analyzes evaluation metrics from multiple training runs to determine +the convergence point of different reinforcement learning architectures. + +Workflow: +1. Loads evaluation data from the CSV files defined in FILE_MAPPING. +2. Calculates a rolling average of the reward and velocity to smooth noise. +3. Determines the convergence timestep for each metric (first time 95% of peak is reached). +4. Generates a grouped bar chart comparing convergence speed and line plots of the raw curves. + +Usage: + uv run python scripts/analysis/analyze_convergence.py + +Note: For these metrics to be valid, the evaluation CSVs must be generated with +exploration noise strictly disabled (e.g., taking the mean of the action distribution). +""" + +import logging +import os + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from enum import Enum + +from plot_config import COLORS, apply_style, create_common_parser, LEGEND_KWARGS + +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s [%(levelname)s] %(name)s: %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", +) +logger = logging.getLogger(__name__) + +# --- Globals & Configuration --- +USING_DUMMY_DATA = False +SMOOTHING_WINDOW = 3 +CONVERGENCE_THRESHOLD = 0.95 + + +class Columns(str, Enum): + # ... (rest of the file remains same, just need to update plotting functions and obtain_data) + """Column names expected in every evaluation CSV.""" + + ARCH = "architecture" + TIMESTEPS = "total_trained_timesteps" + REWARD = "accumulated_reward" + VELOCITY = "velocity" + + +# Maps architecture display names to the path of their evaluation CSV. +# Update these paths once real evaluation data is available. +FILE_MAPPING: dict[str, str] = { + "centralized 2 arms": "runs/dummy/dummy_centralized_2_arms.csv", + "centralized 5 arms": "runs/dummy/dummy_centralized_5_arms.csv", + "decentralized fully connected": "runs/dummy/dummy_decentralized_fully_connected.csv", + "decentralized ring-level": "runs/dummy/dummy_decentralized_ring-level.csv", + "decentralized segment-level": "runs/dummy/dummy_decentralized_segment-level.csv", +} + +# Architecture profiles for dummy data generation: (max_reward, max_velocity, sigmoid_speed) +_DUMMY_PROFILES: dict[str, tuple[float, float, float]] = { + "centralized 2 arms": (300, 0.8, 1.2), + "centralized 5 arms": (450, 1.1, 1.0), + "decentralized fully connected": (500, 1.3, 0.7), + "decentralized ring-level": (480, 1.2, 0.8), + "decentralized segment-level": (520, 1.4, 0.6), +} + + +def generate_dummy_csvs(file_mapping: dict[str, str]): + """ + Generates one dummy CSV per architecture in FILE_MAPPING at their expected locations. + Skips any architecture without a defined profile. + """ + checkpoints = list(range(100, 1100, 100)) + timesteps = [cp * 10_000 for cp in checkpoints] + + for arch, path in file_mapping.items(): + if arch not in _DUMMY_PROFILES: + logger.warning(f"No dummy profile for '{arch}'. Skipping.") + continue + + m_reward, m_vel, speed = _DUMMY_PROFILES[arch] + + rows = [] + for i, ts in enumerate(timesteps): + progress = 1 / (1 + np.exp(-speed * (i - 4))) + rows.append( + { + Columns.TIMESTEPS: ts, + Columns.REWARD: m_reward * progress + np.random.normal(0, 5), + Columns.VELOCITY: m_vel * progress + np.random.normal(0, 0.02), + } + ) + + # Create parent directories if they don't exist + os.makedirs(os.path.dirname(path), exist_ok=True) + + pd.DataFrame(rows).to_csv(path, index=False) + logger.info(f"Generated dummy CSV at expected path: {path}") + + +def load_metrics(file_mapping: dict[str, str]) -> pd.DataFrame: + """ + Loads one CSV per architecture, injects the architecture name as a column, + and returns the combined DataFrame with only the required columns. + """ + required = [Columns.TIMESTEPS, Columns.REWARD, Columns.VELOCITY] + dfs = [] + + for arch_name, filepath in file_mapping.items(): + if not os.path.exists(filepath): + logger.warning(f"File not found: '{filepath}'. Skipping.") + continue + + df = pd.read_csv(filepath) + + missing = [c for c in required if c not in df.columns] + if missing: + logger.warning(f"Missing columns {missing} in '{filepath}'. Skipping.") + continue + + df = df[required].copy() + df[Columns.ARCH] = arch_name + dfs.append(df) + + return pd.concat(dfs, ignore_index=True) if dfs else pd.DataFrame() + + +def _convergence_timestep(series: pd.Series, timesteps: pd.Series) -> float: + """Returns the first timestep where the smoothed series reaches 95% of its peak.""" + smoothed = series.rolling(window=SMOOTHING_WINDOW, min_periods=1).mean() + threshold = smoothed.max() * CONVERGENCE_THRESHOLD + return timesteps[smoothed >= threshold].iloc[0] + + +def analyze_convergence(df: pd.DataFrame) -> pd.DataFrame: + """ + For each architecture, determines the convergence timestep based on both + reward and velocity, returning one summary row per architecture. + """ + results = [] + + for arch in df[Columns.ARCH].unique(): + arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS) + + results.append( + { + "Architecture": arch, + "Reward_Convergence_Timestep": _convergence_timestep( + arch_data[Columns.REWARD], arch_data[Columns.TIMESTEPS] + ), + "Velocity_Convergence_Timestep": _convergence_timestep( + arch_data[Columns.VELOCITY], arch_data[Columns.TIMESTEPS] + ), + } + ) + + return pd.DataFrame(results) + + +def _add_bar_labels(bars, max_val: float): + """Annotates each bar with its value in white bold text, positioned inside.""" + for bar in bars: + width = bar.get_width() + label = f"{width / 1e6:.1f}M" if width >= 1e6 else f"{width:,.0f}" + plt.text( + width - (max_val * 0.02), + bar.get_y() + bar.get_height() / 2, + label, + ha="right", + va="center", + fontsize=11, + color="white", + fontweight="bold", + ) + + +def plot_grouped_convergence_chart( + results_df: pd.DataFrame, output_filename: str, output_dir: str, **kwargs +): + """ + Saves a grouped horizontal bar chart comparing Reward and Velocity convergence timesteps + across all architectures. + """ + sorted_df = results_df.sort_values("Reward_Convergence_Timestep", ascending=True) + architectures = sorted_df["Architecture"].tolist() + y_pos = np.arange(len(architectures)) + bar_height = 0.35 + max_val = sorted_df[ + ["Reward_Convergence_Timestep", "Velocity_Convergence_Timestep"] + ].values.max() + + fig, ax = plt.subplots(figsize=kwargs.get("figsize", (12, 8))) + + bars_reward = ax.barh( + y_pos + bar_height / 2, + sorted_df["Reward_Convergence_Timestep"], + height=bar_height, + label="Reward Convergence", + color="#1f77b4", + ) + bars_velocity = ax.barh( + y_pos - bar_height / 2, + sorted_df["Velocity_Convergence_Timestep"], + height=bar_height, + label="Velocity Convergence", + color="#ff7f0e", + ) + + title_suffix = " (DUMMY DATA)" if USING_DUMMY_DATA else "" + if kwargs.get("show_titles", True): + ax.set_title( + f"Comparison of Training Convergence Timesteps{title_suffix}", fontsize=20, pad=20 + ) + ax.set_xlabel("Timesteps to Convergence (95% of peak)", fontsize=16) + ax.set_ylabel("Architecture", fontsize=16) + ax.set_yticks(y_pos) + ax.set_yticklabels(architectures, fontsize=14) + ax.tick_params(axis="x", labelsize=14) + ax.legend(**LEGEND_KWARGS, ncol=2) + ax.set_xlim(left=0) + ax.spines["top"].set_visible(False) + ax.spines["right"].set_visible(False) + + _add_bar_labels(bars_reward, max_val) + _add_bar_labels(bars_velocity, max_val) + + plt.tight_layout() + os.makedirs(output_dir, exist_ok=True) + base_path = os.path.join(output_dir, os.path.splitext(output_filename)[0]) + plt.savefig(f"{base_path}.png", dpi=300, bbox_inches="tight") + plt.savefig(f"{base_path}.svg", format="svg", bbox_inches="tight") + plt.close() + + +def plot_metric_curves( + df: pd.DataFrame, metric_col: str, title: str, output_filename: str, output_dir: str, **kwargs +): + """ + Saves a line plot of the given metric over training timesteps for every architecture. + """ + fig, ax = plt.subplots(figsize=kwargs.get("figsize", (12, 7))) + + for arch in df[Columns.ARCH].unique(): + arch_data = df[df[Columns.ARCH] == arch].sort_values(Columns.TIMESTEPS) + color_key = arch.split()[0].upper() if isinstance(arch, str) else "UNKNOWN" + color = COLORS.get(color_key, "#888888") + ax.plot( + arch_data[Columns.TIMESTEPS], + arch_data[metric_col], + label=arch, + marker="o", + markersize=4, + alpha=0.8, + color=color, + ) + + title_suffix = " (DUMMY DATA)" if USING_DUMMY_DATA else "" + if kwargs.get("show_titles", True): + ax.set_title(f"{title}{title_suffix}", fontsize=18, pad=20) + ax.set_xlabel("Training Timesteps", fontsize=14) + ax.set_ylabel(metric_col.replace("_", " ").title(), fontsize=14) + ax.legend(**LEGEND_KWARGS, ncol=len(df[Columns.ARCH].unique())) + ax.grid(True, linestyle="--", alpha=0.6) + ax.set_xlim(left=0) + ax.set_ylim(bottom=0) + + plt.tight_layout() + os.makedirs(output_dir, exist_ok=True) + base_path = os.path.join(output_dir, os.path.splitext(output_filename)[0]) + plt.savefig(f"{base_path}.png", dpi=300, bbox_inches="tight") + plt.savefig(f"{base_path}.svg", format="svg", bbox_inches="tight") + plt.close() + + +def plot_results(df: pd.DataFrame, results: pd.DataFrame, output_dir: str, **kwargs): + """Generates and saves all analysis plots.""" + plot_grouped_convergence_chart( + results, output_filename="convergence_comparison.png", output_dir=output_dir, **kwargs + ) + plot_metric_curves( + df, + Columns.REWARD, + "Training Progress: Accumulated Reward", + "progress_reward_curves.png", + output_dir=output_dir, + **kwargs, + ) + plot_metric_curves( + df, + Columns.VELOCITY, + "Training Progress: Velocity", + "progress_velocity_curves.png", + output_dir=output_dir, + **kwargs, + ) + + +def obtain_data() -> pd.DataFrame: + """Resolves the file mapping, falling back to generated dummy CSVs if needed.""" + global USING_DUMMY_DATA + if not any(os.path.exists(p) for p in FILE_MAPPING.values()): + logger.info("No real evaluation files found. Generating dummy CSVs at expected locations.") + generate_dummy_csvs(FILE_MAPPING) + USING_DUMMY_DATA = True + + return load_metrics(FILE_MAPPING) + + +def run_analysis(output_dir: str, **kwargs): + """Orchestrates data loading, convergence analysis, and plot generation.""" + df = obtain_data() + if df.empty: + logger.error("No data found to analyze.") + return + + results = analyze_convergence(df) + plot_results(df, results, output_dir, **kwargs) + logger.info("Analysis complete. Plots saved to disk.") + + +if __name__ == "__main__": + parser = create_common_parser(description="Analyze training convergence.") + args = parser.parse_args() + + apply_style(font_size=args.font_size) + run_analysis( + output_dir=args.output_dir, + show_titles=args.show_titles, + figsize=(args.fig_width, args.fig_height), + ) diff --git a/scripts/plots/plot_config.py b/scripts/plots/plot_config.py new file mode 100644 index 0000000..5fd6ffd --- /dev/null +++ b/scripts/plots/plot_config.py @@ -0,0 +1,77 @@ +import argparse +import matplotlib.pyplot as plt + +# Shared Color Palette (Colorblind friendly, high contrast) +# Matches poster design +COLORS = { + "CENTRALIZED": "#2B4162", # Deep Slate Blue + "FULLY_CONNECTED": "#FA9F42", # Vibrant Orange + "RING_LEVEL": "#4E937A", # Muted Teal + "SEGMENT_LEVEL": "#B4436C", # Soft Red + "DECENTRALIZED": "#4E937A", # Default decentralized fallback +} + + +def apply_style(font_size=28): + """ + Applies the shared typography and aesthetic settings to Matplotlib. + """ + plt.rcParams.update( + { + "font.size": font_size, + "axes.labelsize": font_size + 4, + "axes.titlesize": font_size + 8, + "xtick.labelsize": font_size - 4, + "ytick.labelsize": font_size - 4, + "legend.fontsize": font_size - 6, + "axes.linewidth": 2, + "axes.spines.top": False, + "axes.spines.right": False, + "axes.spines.left": False, + "figure.facecolor": "white", + "axes.facecolor": "white", + "savefig.bbox": "tight", + "savefig.dpi": 300, + } + ) + + +# Star marker for best performer +BEST_PERFORMER_TEXT = "★" +BEST_PERFORMER_MARKER = "*" +BEST_PERFORMER_COLOR = "#D4AF37" # Gold + +# Centralized Legend Configuration +LEGEND_KWARGS = { + "loc": "upper center", + "bbox_to_anchor": (0.5, -0.5), + "frameon": False, +} + + +def create_common_parser(description: str) -> argparse.ArgumentParser: + """ + Creates an argparse parser with common plotting arguments. + """ + parser = argparse.ArgumentParser(description=description) + parser.add_argument( + "--output_dir", + "-o", + default="runs/evaluation/plots", + help="Directory to save the generated plots.", + ) + parser.add_argument( + "--show_titles", + action="store_true", + help="Include titles in the plots. Default is False for easier poster integration.", + ) + parser.add_argument( + "--font_size", type=int, default=28, help="Base font size in points. Default is 28." + ) + parser.add_argument( + "--fig_width", type=float, default=12.0, help="Figure width in inches. Default is 12.0." + ) + parser.add_argument( + "--fig_height", type=float, default=8.0, help="Figure height in inches. Default is 8.0." + ) + return parser diff --git a/scripts/simulate.py b/scripts/simulate.py index a2b121e..9f3e550 100644 --- a/scripts/simulate.py +++ b/scripts/simulate.py @@ -13,19 +13,14 @@ from __future__ import annotations from pathlib import Path import hydra -import numpy as np from omegaconf import DictConfig, OmegaConf -import yaml -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 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, @@ -60,69 +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, - ) - obs_processor = create_obs_processor( - bounds_dict=training.obs_bounds.to_bounds_dict(), - padding_masks=padding_masks, - ) - - # 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 = sum(training.morphology.segments_per_arm) * 2 - - # 7. Load policy - policy = PolicyAgent.from_checkpoint( - model_path, action_dim=trained_action_dim, obs_processor=obs_processor - ) - - # 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/scripts/simulate.sh b/scripts/simulate.sh new file mode 100755 index 0000000..1e6901d --- /dev/null +++ b/scripts/simulate.sh @@ -0,0 +1,9 @@ +#!/usr/bin/env bash + +path=$1 + +uv run simulate.py \ + simulation.model_path="$path"/final_model.flax \ + simulation.record_video=True \ + simulation.video_output_path=../vids/simulation.mp4 \ + simulation.max_steps=10000 diff --git a/src/brittle_star_project/MLPs/__init__.py b/src/brittle_star_project/MLPs/__init__.py new file mode 100644 index 0000000..8f9c7f7 --- /dev/null +++ b/src/brittle_star_project/MLPs/__init__.py @@ -0,0 +1,19 @@ +from .mlps import ( + GenericDenseLayersWithActivation, + OneDenseLayerMLP, + Actor, + MessagePasser, + AgentParams, + Storage, +) +from .adjancency_builder import build_adjacency + +__all__ = [ + "GenericDenseLayersWithActivation", + "OneDenseLayerMLP", + "Actor", + "MessagePasser", + "AgentParams", + "Storage", + "build_adjacency", +] diff --git a/src/brittle_star_project/MLPs/adjancency_builder.py b/src/brittle_star_project/MLPs/adjancency_builder.py new file mode 100644 index 0000000..878a02c --- /dev/null +++ b/src/brittle_star_project/MLPs/adjancency_builder.py @@ -0,0 +1,67 @@ +from brittle_star_project.environment.env_config import MorphMode +import jax.numpy as jnp + + +def build_adjacency(segments_per_arm, mode: MorphMode): + num_arms = sum(1 for s in segments_per_arm if s > 0) + num_segments = sum(segments_per_arm) + + # FOR NOW SEMI HARDCODE: + # CENTRALIZED: 1 agent, no stress, adja = 1,1 = [[1]] + # FULLY CONNECTED: 5 agents: adj = alle 1 + # CENTRAL DISK:#arms= 5 agents, only neighbor as adjacent so diagonal kinda.. + # ARM = #segments agents: diago kinda, but extra, center ring too, put center mlps first or.. + + if mode == MorphMode.CENTRALIZED: + return jnp.ones((1, 1)) + + if mode == MorphMode.FULLY_CONNECTED: + adj = jnp.ones((num_arms, num_arms)) # everybody adjacent everybody + return adj + + if mode == MorphMode.RING: # ring + adj = jnp.zeros((num_arms, num_arms)) + for i in range(num_arms): + adj = adj.at[i, i].set(1) # self + adj = adj.at[i, (i - 1) % num_arms].set(1) + adj = adj.at[i, (i + 1) % num_arms].set(1) # left and right.. + return adj + + if mode == MorphMode.SEGMENT: + num_nodes = num_arms + num_segments + adj = jnp.zeros((num_nodes, num_nodes)) + + # first ring + for i in range(num_arms): + # self + adj = adj.at[i, i].set(1) + + # ring neighbors + adj = adj.at[i, (i - 1) % num_arms].set(1) + adj = adj.at[i, (i + 1) % num_arms].set(1) + + # then segment chains + idx = 0 + for arm_idx, seg_count in enumerate(segments_per_arm): + for i in range(seg_count): + seg_node = num_arms + idx + i + + adj = adj.at[seg_node, seg_node].set(1) + if i > 0: + adj = adj.at[seg_node, seg_node - 1].set(1) + if i < seg_count - 1: + adj = adj.at[seg_node, seg_node + 1].set(1) + + idx += seg_count + + idx = 0 + for arm_idx, seg_count in enumerate(segments_per_arm): + first_seg = num_arms + idx # first segment of this arm + + # connect ring node first segment + adj = adj.at[arm_idx, first_seg].set(1) + adj = adj.at[first_seg, arm_idx].set(1) + + idx += seg_count + + return adj diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 5e2deb5..2c5989f 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -1,11 +1,11 @@ from dataclasses import dataclass, fields, field -import flax import flax.linen as nn import jax.numpy as jnp import jax.tree_util from typing import Sequence, Callable from flax.linen.initializers import constant, orthogonal +from flax.core import FrozenDict # semi generic so we can easily make a config for it in experiments @@ -37,30 +37,56 @@ class Actor(nn.Module): return mean, log_std +class MessagePasser(nn.Module): + hidden_dim: int + num_propagation_steps: int + adj_matrix: jnp.ndarray + + @nn.compact + def __call__(self, x: jnp.ndarray): + for _ in range(self.num_propagation_steps): + # (n_nodes, feat) + messages = nn.Dense(self.hidden_dim)(x) + messages = nn.tanh(messages) + + # note: if mean is wanted: adj_matrix / (adj.sum(axis=-1, keepdims=True) + 1e-8) + agg = self.adj_matrix + aggregated = agg @ messages + + x_concat = jnp.concatenate([x, aggregated], axis=-1) + + gate = nn.sigmoid(nn.Dense(self.hidden_dim)(x_concat)) + candidate = nn.tanh(nn.Dense(self.hidden_dim)(x_concat)) + x = gate * x + (1 - gate) * candidate + + return x + + @jax.tree_util.register_dataclass @dataclass class AgentParams: - sensor_params: flax.core.FrozenDict - actor_params: flax.core.FrozenDict - critic_params: flax.core.FrozenDict - feature_extractor_params: flax.core.FrozenDict + sensor_params: FrozenDict | dict + actor_params: FrozenDict | dict + critic_params: FrozenDict | dict + feature_extractor_params: FrozenDict | dict + message_passer_params: FrozenDict | dict @jax.tree_util.register_dataclass @dataclass class Storage: - obs: jnp.array - actions: jnp.array - logprobs: jnp.array - dones: jnp.array - values: jnp.array - advantages: jnp.array - returns: jnp.array - rewards: jnp.array + obs: jnp.ndarray + actions: jnp.ndarray + logprobs: jnp.ndarray + dones: jnp.ndarray + values: jnp.ndarray + advantages: jnp.ndarray + returns: jnp.ndarray + rewards: jnp.ndarray - raw_actions: jnp.ndarray = None # before clipping - means: jnp.ndarray = None # policy mean - stds: jnp.ndarray = None # policy std + raw_actions: jnp.ndarray | None = None # before clipping + means: jnp.ndarray | None = None # policy mean + stds: jnp.ndarray | None = None # policy std def replace(self, **kwargs) -> "Storage": fs = fields(self) 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 new file mode 100644 index 0000000..c9c248b --- /dev/null +++ b/src/brittle_star_project/configs/config_evaluation.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from dataclasses import dataclass, field + + +@dataclass +class EvaluationConfig: + """Evaluation settings. + + Currently used for synchronous checkpoint evaluation during training. + """ + + # When enabled, each saved checkpoint is evaluated headlessly and the results + # are appended to a CSV in the run's metrics/ folder. + evaluate_checkpoints: bool = False + 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( + "Configuration Error: 'eval_max_steps' must be > 0 when " + "'evaluate_checkpoints' is enabled." + ) diff --git a/src/brittle_star_project/configs/main_config.py b/src/brittle_star_project/configs/main_config.py index 10fd22e..7caa1aa 100644 --- a/src/brittle_star_project/configs/main_config.py +++ b/src/brittle_star_project/configs/main_config.py @@ -2,6 +2,7 @@ from dataclasses import dataclass, field from experiment_logger.config_logger import LoggingConfig from brittle_star_project.configs.config_experiment import ExperimentConfig +from brittle_star_project.configs.config_evaluation import EvaluationConfig from brittle_star_project.configs.config_ppo import PPOConfig from brittle_star_project.configs.config_architecture import ArchitectureConfig from brittle_star_project.configs.config_simulation import SimulationSettings @@ -23,6 +24,7 @@ class BrittleStarConfig: experiment: ExperimentConfig = field(default_factory=ExperimentConfig) logging: LoggingConfig = field(default_factory=LoggingConfig) + evaluation: EvaluationConfig = field(default_factory=EvaluationConfig) ppo: PPOConfig = field(default_factory=PPOConfig) # This field is polymorphic; defaults to the base class to allow subclasses # (CentralizedConfig, DecentralizedConfig) to be merged in via Hydra. diff --git a/src/brittle_star_project/configs/register_configs.py b/src/brittle_star_project/configs/register_configs.py index 147da10..9522212 100644 --- a/src/brittle_star_project/configs/register_configs.py +++ b/src/brittle_star_project/configs/register_configs.py @@ -2,6 +2,7 @@ from hydra.core.config_store import ConfigStore from experiment_logger.config_logger import LoggingConfig from brittle_star_project.configs.config_experiment import ExperimentConfig +from brittle_star_project.configs.config_evaluation import EvaluationConfig from brittle_star_project.configs.config_ppo import PPOConfig from brittle_star_project.configs.config_architecture import ( CentralizedConfig, @@ -32,6 +33,7 @@ def register_configs() -> None: # Sub-config groups — each group corresponds to a configs/ subdirectory. cs.store(group="experiment", name="base_experiment", node=ExperimentConfig) cs.store(group="logging", name="base_logging", node=LoggingConfig) + cs.store(group="evaluation", name="base_evaluation", node=EvaluationConfig) cs.store(group="ppo", name="base_ppo", node=PPOConfig) # Architecture variants — swap via CLI: architecture=decentralized diff --git a/src/brittle_star_project/dataclasses/EpisodeStatistics.py b/src/brittle_star_project/dataclasses/EpisodeStatistics.py index 0b2832f..ff2982e 100644 --- a/src/brittle_star_project/dataclasses/EpisodeStatistics.py +++ b/src/brittle_star_project/dataclasses/EpisodeStatistics.py @@ -4,7 +4,7 @@ import jax.numpy as jnp @flax.struct.dataclass class EpisodeStatistics: - episode_returns: jnp.array - episode_lengths: jnp.array - returned_episode_returns: jnp.array - returned_episode_lengths: jnp.array + episode_returns: jnp.ndarray + episode_lengths: jnp.ndarray + returned_episode_returns: jnp.ndarray + returned_episode_lengths: jnp.ndarray diff --git a/src/brittle_star_project/environment/__init__.py b/src/brittle_star_project/environment/__init__.py index 77d106c..7a9a9eb 100644 --- a/src/brittle_star_project/environment/__init__.py +++ b/src/brittle_star_project/environment/__init__.py @@ -1,4 +1,4 @@ -from .env_config import ArenaConfig, EnvConfig, MorphologyConfig +from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, MorphMode from .env_types import Backend, Task from .env_wrapper import BrittleStarEnv from .factory import BrittleStarEnvFactory @@ -13,6 +13,7 @@ __all__ = [ "Task", "BrittleStarEnv", "BrittleStarEnvFactory", + "MorphMode", "create_obs_processor", "compute_padding_masks", ] diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index 22b8967..33e8bdf 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -1,10 +1,18 @@ from __future__ import annotations from dataclasses import dataclass, field +from enum import Enum from .env_types import Task +class MorphMode(Enum): + CENTRALIZED = 0 + FULLY_CONNECTED = 1 + RING = 2 + SEGMENT = 3 + + @dataclass class MorphologyConfig: """Brittle star morphology configuration. @@ -20,6 +28,7 @@ class MorphologyConfig: segments_per_arm: list[int] = field(default_factory=lambda: [4, 4, 4, 4, 4]) use_p_control: bool = True use_torque_control: bool = False + morph_mode: MorphMode = MorphMode.CENTRALIZED @property def num_arms(self) -> int: diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index b18ee8c..e169ee3 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -2,6 +2,12 @@ import jax import jax.numpy as jnp from typing import Dict, Tuple, Optional +from brittle_star_project.environment.env_config import MorphMode + +from experiment_logger import get_logger + +logger = get_logger() + _JOINT_SCALED_KEYS = frozenset( { "joint_position", @@ -18,9 +24,53 @@ _SEGMENT_SCALED_KEYS = frozenset( ) +def _build_joint_indices(segments_per_arm, indices_mlp): + indices = [] + start = 0 + for i, segs in enumerate(segments_per_arm): + # 2 joints per segment + if i in indices_mlp: + count = segs * 2 + idx = jnp.arange(start, start + count) + indices.append(idx) + start += count + return indices + + +def _build_segment_indices(segments_per_arm, indices_mlp): + indices = [] + start = 0 + for i, segs in enumerate(segments_per_arm): + if i in indices_mlp: + idx = jnp.arange(start, start + segs) + indices.append(idx) + start += segs + return indices + + def create_obs_processor( - bounds_dict: Dict[str, Tuple[float, float]], padding_masks: Optional[Dict] = None + bounds_dict: Dict[str, Tuple[float, float]], + num_arms: int, + needed_copies: int, + padding_masks: Optional[Dict] = None, + morph_mode: MorphMode = MorphMode.CENTRALIZED, + segments_per_arm=[4, 4, 4, 4, 4], + agent_indices=[0, 1, 2, 3, 4], ): + # made a set to allow O(1) search + ordered_keys = frozenset( + [ + "disk_z_tilt", + "joint_actuator_force", + "joint_position", + "joint_velocity", + "robot_direction_to_target", + "segment_contact", + ] + ) + segment_indices = _build_segment_indices(segments_per_arm, agent_indices) + joint_indices = _build_joint_indices(segments_per_arm, agent_indices) + def _add_derived_features(obs: dict) -> dict: new_obs = dict(obs) if "disk_rotation" in new_obs: @@ -51,41 +101,92 @@ def create_obs_processor( normalized[key] = arr return normalized - def _pad_features(obs: dict) -> dict: - padded = {} + def _split_to_agents(obs: dict, morph_mode) -> dict: + output = {} + num_agents = needed_copies # IMPORTANT: number of MLPs + + segs_per_arm = 4 + joints_per_segment = 2 + joints_per_arm = segs_per_arm * joints_per_segment for key, arr in obs.items(): - if key in _JOINT_SCALED_KEYS: - padded_arr = jnp.zeros(padding_masks["target_size_2x"], dtype=arr.dtype) - padded[key] = padded_arr.at[padding_masks["mask_2x"]].set(arr) - elif key in _SEGMENT_SCALED_KEYS: - padded_arr = jnp.zeros(padding_masks["target_size_1x"], dtype=arr.dtype) - padded[key] = padded_arr.at[padding_masks["mask_1x"]].set(arr) + arr = jnp.asarray(arr) + if arr.size == 0: + continue + + if arr.ndim == 0: + arr = arr.reshape(1) + + if key in _SEGMENT_SCALED_KEYS: + per_agent = [] + for i, _ in enumerate(agent_indices): + idx = segment_indices[i] + taken = jnp.take(arr, idx, axis=0) + pad_len = segs_per_arm - taken.shape[0] + padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1)) + + per_agent.append(padded.reshape(-1)) + arr = jnp.stack(per_agent) + elif key in _JOINT_SCALED_KEYS: + per_agent = [] + for i, _ in enumerate(agent_indices): + idx = joint_indices[i] + taken = jnp.take(arr, idx, axis=0) + pad_len = joints_per_arm - taken.shape[0] + padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1)) + + per_agent.append(padded.reshape(-1)) + arr = jnp.stack(per_agent) else: - padded[key] = arr - return padded + arr = jnp.repeat(arr[None, :], num_agents, axis=0) + + if morph_mode == MorphMode.CENTRALIZED: + output[key] = arr.reshape(1, -1) + elif key in _JOINT_SCALED_KEYS: + output[key] = arr.reshape(num_agents, -1) + elif key in _SEGMENT_SCALED_KEYS: + output[key] = arr[:, None] + else: + output[key] = arr + + return output def _flatten_features(obs: dict) -> jnp.ndarray: - ordered_keys = [ - "disk_z_tilt", - "joint_actuator_force", - "joint_position", - "joint_velocity", - "robot_direction_to_target", - "segment_contact", - ] + """ + Input: + key -> (num_arms, feat_per_key) + + Output: + (num_arms, total_features) + """ values = [] - for key in ordered_keys: - if key in obs: - arr = jnp.asarray(obs[key]).flatten() - if arr.size > 0: - values.append(arr) - return jnp.concatenate(values) + + for key in sorted(ordered_keys): + if key not in obs: + continue + + arr = jnp.asarray(obs[key]) # (num_arms, feat) + + if arr.size == 0: + continue + + if arr.ndim == 1: + arr = arr[:, None] + + arr = arr.reshape(arr.shape[0], -1) + + values.append(arr) + + return jnp.concatenate(values, axis=-1) # (num_arms, total_feat) def _process_single(obs_dict: dict) -> jnp.ndarray: processed = _add_derived_features(obs_dict) processed = _normalize_features(processed) - if padding_masks is not None: - processed = _pad_features(processed) - return _flatten_features(processed) + processed = _split_to_agents(processed, morph_mode) + flat = _flatten_features(processed) # (num_arms, total_feat) + + logger.debug(f"[FLATTENED FINAL] shape: {flat.shape}") + logger.debug(f"[PER AGENT] example row 0 shape: {flat[0].shape}") + + return flat # (agents, feat) return jax.jit(jax.vmap(_process_single)) diff --git a/src/brittle_star_project/environment/padded_obs_wrapper.py b/src/brittle_star_project/environment/padded_obs_wrapper.py index ea00713..3216dfe 100644 --- a/src/brittle_star_project/environment/padded_obs_wrapper.py +++ b/src/brittle_star_project/environment/padded_obs_wrapper.py @@ -30,6 +30,12 @@ def compute_padding_masks( mask_2x = [] for arm_idx, (actual, ref) in enumerate(zip(segments_per_arm, reference_segments_per_arm)): + if not isinstance(actual, int): + actual = actual.item() + + if not isinstance(ref, int): + ref = ref.item() + if not (0 <= actual <= ref): raise ValueError( f"Invalid amputation at arm {arm_idx}: " diff --git a/src/brittle_star_project/evaluation/__init__.py b/src/brittle_star_project/evaluation/__init__.py index 38a1c8b..e529115 100644 --- a/src/brittle_star_project/evaluation/__init__.py +++ b/src/brittle_star_project/evaluation/__init__.py @@ -1,21 +1,43 @@ from __future__ import annotations from .checkpoint import load_metadata, load_params, metadata_to_configs, TrainingConfig +from .evaluate_mjx import ( + CheckpointEvalResult, + append_checkpoint_eval_row, + 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 "load_metadata", "load_params", "metadata_to_configs", "TrainingConfig", + # MJX evaluation + "CheckpointEvalResult", + "append_checkpoint_eval_row", + "build_eval_rollout_fn", + "evaluate_checkpoint_mjx", + # CPU evaluation + "evaluate_policy", + # policy "PolicyAgent", "ControlPolicy", + # rollout "rollout_headless", "rollout_viewer", "EpisodeResult", + # video "record_episode", "create_evaluation_dir", "save_evaluation_metadata", + # env builder + "EvalEnvBundle", + "build_eval_env", ] diff --git a/src/brittle_star_project/evaluation/checkpoint.py b/src/brittle_star_project/evaluation/checkpoint.py index d10046b..9b868d5 100644 --- a/src/brittle_star_project/evaluation/checkpoint.py +++ b/src/brittle_star_project/evaluation/checkpoint.py @@ -4,6 +4,8 @@ import yaml from dataclasses import dataclass from pathlib import Path +from collections.abc import Mapping + import flax from omegaconf import OmegaConf @@ -32,15 +34,19 @@ def load_params(path: Path) -> dict: sensor_params = None actor_params = None + message_passer_params = None # Extract params from restored checkpoint - if isinstance(restored, dict): + if isinstance(restored, Mapping): params_sub = restored.get("params", {}) sensor_params = restored.get("sensor_params") or params_sub.get("sensor_params") actor_params = restored.get("actor_params") or params_sub.get("actor_params") + message_passer_params = restored.get("message_passer_params") or params_sub.get( + "message_passer_params" + ) elif isinstance(restored, (list, tuple)) and len(restored) >= 2: params_part = restored[1] - if isinstance(params_part, dict): + if isinstance(params_part, Mapping): sensor_params = params_part.get("0", params_part.get(0)) actor_params = params_part.get("1", params_part.get(1)) elif isinstance(params_part, (list, tuple)) and len(params_part) >= 2: @@ -53,6 +59,7 @@ def load_params(path: Path) -> dict: return { "sensor_params": sensor_params, "actor_params": actor_params, + "message_passer_params": message_passer_params, } 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 new file mode 100644 index 0000000..f557d2d --- /dev/null +++ b/src/brittle_star_project/evaluation/evaluate_mjx.py @@ -0,0 +1,258 @@ +"""MJX-based headless checkpoint evaluation. + +This module provides a fast, JIT-compiled evaluation path using the MJX +(JAX-accelerated MuJoCo) backend. It is intended for evaluating checkpoints +*during* or *after* a training run, where the environment and policy are +already fully initialised. + +The key functions are: + +- `build_eval_rollout_fn` — builds and JIT-compiles a single-episode rollout function from the + training environment and policy components. +- `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. +""" + +from __future__ import annotations + +import csv +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable + +import jax +import jax.numpy as jnp + + +@dataclass +class CheckpointEvalResult: + """Structured result from a single MJX checkpoint evaluation episode.""" + + steps: int + """Number of control steps taken (≤ max_steps).""" + + reached_target: bool + """Whether the robot reached the target (terminated) before max_steps.""" + + eval_return: float + """Accumulated shaped reward over the episode.""" + + final_xy_dist: float + """XY distance to target at episode end. 0.0 when ``reached_target`` is True.""" + + initial_xy_dist: float + """XY distance to target at episode start.""" + + +def build_eval_rollout_fn( + *, + env: Any, + obs_processor: Callable, + sensor_apply: Callable, + actor_apply: Callable, + message_passer_apply: Callable | None = None, + action_low: jnp.ndarray, + action_high: jnp.ndarray, + reward_fn: Callable, +) -> Callable: + """Build and JIT-compile a single-episode MJX evaluation rollout. + + 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`. + 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). + 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,)`). + reward_fn: Shaped reward function with signature + `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. + """ + # vmap over a batch of 1 so the MJX API is satisfied without any + # extra bookkeeping in the caller. + reset_1 = jax.vmap(env.raw.reset) + step_1 = jax.vmap(env.raw.step) + + def _eval_rollout(params: dict, seed: int, max_steps: int): + rng = jax.random.PRNGKey(seed) + rngs = jnp.asarray(jax.random.split(rng, 1)) + state = reset_1(rng=rngs) + + initial_xy_dist = jnp.squeeze(state.observations["xy_distance_to_target"]) + + t0 = jnp.asarray(0, dtype=jnp.int32) + done0 = jnp.squeeze(state.terminated | state.truncated) + return0 = jnp.asarray(0.0, dtype=jnp.float32) + + def cond(carry): + t, _state, done, _return_ = carry + return jnp.logical_and(t < max_steps, jnp.logical_not(done)) + + def body(carry): + t, state, _done, return_ = carry + + obs = obs_processor(state.observations) + hidden = sensor_apply(params["sensor_params"], obs) + if message_passer_apply is not None: + mp_params = params["message_passer_params"] + hidden = jax.vmap(lambda x: message_passer_apply(mp_params, x))(hidden) + mean, _log_std = actor_apply(params["actor_params"], hidden) + + # Deterministic action: use the actor mean, no exploration noise. + flat_mean = mean.reshape(mean.shape[0], -1) + action = jnp.clip(flat_mean, action_low, action_high) + next_state = step_1(state=state, action=action) + + shaped_reward = reward_fn(state, next_state) + return_ = return_ + jnp.squeeze(shaped_reward) + + done_next = jnp.squeeze(next_state.terminated | next_state.truncated) + return (t + 1, next_state, done_next, return_) + + t, final_state, _done, return_ = jax.lax.while_loop(cond, body, (t0, state, done0, return0)) + + reached_target = jnp.squeeze(final_state.terminated) + final_xy_dist_raw = jnp.squeeze(final_state.observations["xy_distance_to_target"]) + # Clamp to 0 when the target was reached so downstream consumers + # don't have to special-case "terminated" themselves. + final_xy_dist = jnp.where(reached_target, 0.0, final_xy_dist_raw) + + return t, reached_target, return_, final_xy_dist, initial_xy_dist + + return jax.jit(_eval_rollout) + + +def evaluate_checkpoint_mjx( + eval_fn: Callable, + params: dict, + *, + seed: int, + max_steps: int, +) -> CheckpointEvalResult: + """Run one deterministic evaluation episode and return typed metrics. + + Args: + eval_fn: A JIT-compiled function as returned by `build_eval_rollout_fn`. + params: Agent parameter dict (e.g. ``agent_state.params``). + seed: Random seed for environment reset (controls target placement). + max_steps: Maximum number of control steps before the episode is cut off. + + Returns: + A `CheckpointEvalResult` with all JAX arrays converted to + plain Python scalars. + """ + steps, reached, eval_return, final_xy_dist, initial_xy_dist = eval_fn(params, seed, max_steps) + return CheckpointEvalResult( + steps=int(steps), + reached_target=bool(reached), + eval_return=float(eval_return), + final_xy_dist=float(final_xy_dist), + initial_xy_dist=float(initial_xy_dist), + ) + + +_FIELDNAMES = [ + "checkpoint", + "trained_timesteps", + "eval_steps", + "eval_return", + "final_xy_dist", + "initial_xy_dist", + "reached_target", +] + + +def _migrate_csv_if_needed(csv_path: Path) -> None: + """Rewrite the CSV with the canonical field names if the schema changed. + + Best-effort: any exception is silently swallowed so that a schema mismatch + never causes a training crash. + """ + try: + with open(csv_path, "r", newline="") as f: + header = next(csv.reader(f), None) + + if header is None or list(header) == _FIELDNAMES: + return # Nothing to migrate. + + migrated_rows: list[dict[str, Any]] = [] + with open(csv_path, "r", newline="") as f: + for row in csv.DictReader(f): + migrated_rows.append( + { + "checkpoint": row.get("checkpoint", row.get("iteration")), + "trained_timesteps": row.get("trained_timesteps"), + "eval_steps": row.get("eval_steps", row.get("steps_to_target")), + "eval_return": row.get("eval_return"), + "final_xy_dist": row.get("final_xy_dist"), + "initial_xy_dist": row.get("initial_xy_dist"), + "reached_target": row.get("reached_target"), + } + ) + + with open(csv_path, "w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=_FIELDNAMES) + writer.writeheader() + writer.writerows(migrated_rows) + except Exception: + pass # Never crash training on a migration issue. + + +def append_checkpoint_eval_row( + run_dir: str | Path, + *, + iteration: int, + trained_timesteps: int, + result: CheckpointEvalResult, +) -> Path: + """Append one evaluation row to `/metrics/checkpoint_evaluation.csv`. + + 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: + run_dir: Root directory of the training run (Hydra's output dir). + iteration: Training iteration number, used as the checkpoint identifier. + trained_timesteps: Total environment steps taken at this checkpoint. + result: Evaluation result as returned by `evaluate_checkpoint_mjx`. + + Returns: + Absolute path to the CSV file (useful for W&B sync). + """ + metrics_dir = Path(run_dir) / "metrics" + metrics_dir.mkdir(parents=True, exist_ok=True) + csv_path = metrics_dir / "checkpoint_evaluation.csv" + + if csv_path.exists(): + _migrate_csv_if_needed(csv_path) + + file_exists = csv_path.exists() + with open(csv_path, "a", newline="") as f: + writer = csv.DictWriter(f, fieldnames=_FIELDNAMES) + if not file_exists: + writer.writeheader() + writer.writerow( + { + "checkpoint": int(iteration), + "trained_timesteps": int(trained_timesteps), + "eval_steps": result.steps, + "eval_return": result.eval_return, + "final_xy_dist": result.final_xy_dist, + "initial_xy_dist": result.initial_xy_dist, + "reached_target": result.reached_target, + } + ) + + return csv_path diff --git a/src/brittle_star_project/evaluation/policy.py b/src/brittle_star_project/evaluation/policy.py index 5aa2d6e..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 @@ -24,10 +25,17 @@ class PolicyAgent: *, sensor_params: Any, actor_params: Any, + message_passer_params: Any | None = None, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, action_dim: int, obs_processor: Any, ) -> None: - from brittle_star_project.MLPs.mlps import Actor, GenericDenseLayersWithActivation + from brittle_star_project.MLPs.mlps import ( + Actor, + GenericDenseLayersWithActivation, + MessagePasser, + ) # Infer layer sizes from params try: @@ -45,7 +53,8 @@ class PolicyAgent: key = f"Dense_{idx}" if key not in dense_params: break - layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[1])) + + layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[-1])) idx += 1 if not layer_sizes: @@ -53,14 +62,69 @@ class PolicyAgent: self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes) self._actor = Actor(action_dim=action_dim) - self._sensor_apply = jax.jit(self._sensor.apply) - self._actor_apply = jax.jit(self._actor.apply) + + self._message_passer = None + if message_passer_params is not None and not ( + isinstance(message_passer_params, dict) and len(message_passer_params) == 0 + ): + if message_passing_steps is None or adj_matrix is None: + raise ValueError( + "Checkpoint contains message_passer_params but PolicyAgent was not given " + "message_passing_steps and adj_matrix. Pass these when constructing the agent " + "so decentralized evaluation matches training." + ) + + hidden_dim = int(layer_sizes[-1]) + self._message_passer = MessagePasser( + hidden_dim=hidden_dim, + num_propagation_steps=int(message_passing_steps), + adj_matrix=jnp.asarray(adj_matrix), + ) + self._message_passer.apply = jax.jit(self._message_passer.apply) + self._sensor.apply = jax.jit(self._sensor.apply) + self._actor.apply = jax.jit(self._actor.apply) self._params = { "sensor_params": sensor_params, "actor_params": actor_params, + "message_passer_params": message_passer_params, } self._obs_processor = obs_processor + @classmethod + def from_params( + cls, + *, + sensor_params: Any, + actor_params: Any, + message_passer_params: Any | None = None, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, + action_dim: int, + obs_processor: Any, + ) -> "PolicyAgent": + """Construct a PolicyAgent directly from in-memory parameters.""" + return cls( + sensor_params=sensor_params, + actor_params=actor_params, + message_passer_params=message_passer_params, + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, + action_dim=action_dim, + obs_processor=obs_processor, + ) + + def set_params( + self, + *, + sensor_params: Any, + actor_params: Any, + message_passer_params: Any | None = None, + ) -> None: + """Update parameters for evaluation without rebuilding the model.""" + self._params["sensor_params"] = sensor_params + self._params["actor_params"] = actor_params + self._params["message_passer_params"] = message_passer_params + @classmethod def from_checkpoint( cls, @@ -68,6 +132,8 @@ class PolicyAgent: *, action_dim: int, obs_processor: Any, + message_passing_steps: int | None = None, + adj_matrix: Any | None = None, ) -> "PolicyAgent": """Load params from .flax and construct the agent.""" params = load_params(model_path) @@ -75,6 +141,9 @@ class PolicyAgent: return cls( sensor_params=params["sensor_params"], actor_params=params["actor_params"], + message_passer_params=params.get("message_passer_params"), + message_passing_steps=message_passing_steps, + adj_matrix=adj_matrix, action_dim=action_dim, obs_processor=obs_processor, ) @@ -82,8 +151,18 @@ class PolicyAgent: 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)[0] - hidden = self._sensor_apply(self._params["sensor_params"], obs) - mean, _log_std = self._actor_apply(self._params["actor_params"], hidden) + obs = self._obs_processor(batched_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") + if mp_params is None or (isinstance(mp_params, dict) and len(mp_params) == 0): + raise ValueError( + "PolicyAgent has a message passer but message_passer_params are missing/empty." + ) + hidden = jax.vmap(lambda x: self._message_passer.apply(mp_params, x))(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 79c86ed..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, ) @@ -114,18 +117,20 @@ def rollout_viewer( episode_return = 0.0 observations = _get_observations(state) + prev_dist = _get_xy_distance_to_target(observations) if observations else None reached_target = _target_reached(state=state) steps = 0 with mujoco.viewer.launch_passive(model, data) as viewer: step_iter = range(int(max_steps)) if max_steps is not None else itertools.count() - for _step_idx in step_iter: + for _ in step_iter: if not viewer.is_running(): break step_start = time.time() obs_dict = observations or {} + action = policy.act(observations=obs_dict) if action_mask is not None: action = action[action_mask] 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/ppo.py b/src/brittle_star_project/ppo.py index 1ea37a8..fb27ccd 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -1,14 +1,27 @@ from functools import partial -import flax import jax import jax.numpy as jnp +from jax import debug +from flax.core import FrozenDict +from experiment_logger import get_logger +from brittle_star_project.utils import logged_jit + +logger = get_logger() # Chose to use a class as it seemed the easiest way to integrate the CleanRL code style # with our need to seperate concerns class PPO: - def __init__(self, args, sensor, actor, critic, feature_extractor, message_passer=None): + def __init__( + self, + args, + sensor_apply, + actor_apply, + critic_apply, + feature_extractor_apply, + message_passer=None, + ): self.args = args if not message_passer: @@ -18,10 +31,10 @@ class PPO: partial( ppo_loss, args=args, - sensor_apply=sensor.apply, - actor_apply=actor.apply, - critic_apply=critic.apply, - feature_extractor_apply=feature_extractor.apply, + sensor_apply=sensor_apply, + actor_apply=actor_apply, + critic_apply=critic_apply, + feature_extractor_apply=feature_extractor_apply, message_passer=message_passer, ), has_aux=True, @@ -29,8 +42,14 @@ class PPO: # This PPO class should be initialized only once, # or this function will need to recompile - @partial(jax.jit, static_argnums=0) + @partial(logged_jit, static_argnums=0) def update_ppo(self, agent_state, storage, key): + debug.callback(logger.debug, f"[PPO] storage.obs shape: {storage.obs.shape}") + debug.callback(logger.debug, f"[PPO] storage.actions shape: {storage.actions.shape}") + debug.callback(logger.debug, f"[PPO] storage.logprobs shape: {storage.logprobs.shape}") + debug.callback(logger.debug, f"[PPO] storage.advantages shape: {storage.advantages.shape}") + debug.callback(logger.debug, f"[PPO] storage.returns shape: {storage.returns.shape}") + args = self.args ppo_loss_grad_fn = self.ppo_loss_grad_fn @@ -49,6 +68,16 @@ class PPO: shuffled_storage = jax.tree.map(convert_data, flatten_storage) def update_minibatch(agent_state, minibatch): + debug.callback(logger.debug, f"[PPO] minibatch.obs: {minibatch.obs.shape}") + debug.callback(logger.debug, f"[PPO] minibatch.actions: {minibatch.actions.shape}") + debug.callback( + logger.debug, f"[PPO] minibatch.logprobs: {minibatch.logprobs.shape}" + ) + debug.callback( + logger.debug, f"[PPO] minibatch.advantages: {minibatch.advantages.shape}" + ) + debug.callback(logger.debug, f"[PPO] minibatch.returns: {minibatch.returns.shape}") + (loss, (pg_loss, v_loss, entropy_loss, approx_kl)), grads = ppo_loss_grad_fn( agent_state.params, minibatch.obs, @@ -58,19 +87,12 @@ class PPO: minibatch.returns, ) agent_state = agent_state.apply_gradients(grads=grads) - return agent_state, ( - loss, - pg_loss, - v_loss, - entropy_loss, - approx_kl, - grads, - ) + return agent_state, (loss, pg_loss, v_loss, entropy_loss, approx_kl) agent_state, metrics = jax.lax.scan(update_minibatch, agent_state, shuffled_storage) return (agent_state, key), metrics - (agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl, grads) = jax.lax.scan( + (agent_state, key), (loss, pg_loss, v_loss, entropy_loss, approx_kl) = jax.lax.scan( update_epoch, (agent_state, key), (), length=args.update_epochs ) return agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key @@ -84,27 +106,45 @@ that are now not in the same scope """ -@partial(jax.jit, static_argnums=(0, 1, 2, 3, 4)) +@partial(logged_jit, static_argnums=(0, 1, 2, 3, 4)) def get_action_and_value( sensor_apply, actor_apply, message_passer, critic_apply, feature_extractor_apply, - params: flax.core.FrozenDict, + params: FrozenDict, x: jnp.ndarray, action: jnp.ndarray, ): hidden_sensor = sensor_apply(params["sensor_params"], x) hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) - hidden_sensor = message_passer(hidden_sensor) + + # only apply message passing in decentralized context + if message_passer is not None: + hidden_sensor = message_passer(params["message_passer_params"], hidden_sensor) + + debug.callback(logger.debug, f"[SHAPE] hidden_sensor: {hidden_sensor.shape}") + debug.callback(logger.debug, f"[SHAPE] hidden_critic: {hidden_critic.shape}") + mean, log_std = actor_apply(params["actor_params"], hidden_sensor) + + debug.callback(logger.debug, f"[SHAPE] mean: {mean.shape}") + debug.callback(logger.debug, f"[SHAPE] log_std: {log_std.shape}") + debug.callback(logger.debug, f"[SHAPE] action: {action.shape}") + log_std = jnp.clip(log_std, -5, 2) std = jnp.exp(log_std) - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(-1) + logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)) + debug.callback(logger.debug, f"[SHAPE] logprob pre-sum: {logprob.shape}") + + logprob = logprob.sum(axis=(-2, -1)) + debug.callback(logger.debug, f"[SHAPE] logprob final: {logprob.shape}") + + entropy = (0.5 + 0.5 * jnp.log(2 * jnp.pi) + log_std).sum(axis=(-2, -1)) value = critic_apply(params["critic_params"], hidden_critic).squeeze(-1) + debug.callback(logger.debug, f"[SHAPE] value: {value.shape}") return logprob, entropy, value @@ -150,7 +190,7 @@ def ppo_loss( return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl)) -def identity(hidden): +def identity(_, hidden): """ Used for seamless jax integration, avoids having branching inside jitted function, diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index a3e4eac..98d4045 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -3,12 +3,13 @@ import random import time from dataclasses import asdict, dataclass from functools import partial -from typing import Any +from typing import Any, Optional import jax import jax.numpy as jnp import numpy as np import optax +import flax.linen as nn from flax.training.train_state import TrainState from experiment_logger import get_logger @@ -17,26 +18,31 @@ from brittle_star_project.configs.main_config import BrittleStarConfig from brittle_star_project.dataclasses import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.obs_processing import create_obs_processor +from brittle_star_project.evaluation.evaluate_mjx import ( + append_checkpoint_eval_row, + 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, GenericDenseLayersWithActivation, + MessagePasser, OneDenseLayerMLP, Storage, ) +from brittle_star_project.MLPs.adjancency_builder import build_adjacency from brittle_star_project.ppo import PPO +from brittle_star_project.environment import MorphMode +from brittle_star_project.utils import logged_jit + +from brittle_star_project.environment.env_types import Backend # TODO: clip scaled reward? -@jax.jit -def _get_xy_distance_to_target(obs_dict: dict) -> jnp.ndarray: - """Extract xy_distance_to_target for all environments.""" - # obs_dict is a dict of arrays with leading batch dimension (num_envs, ...) - return obs_dict["xy_distance_to_target"].squeeze(-1) # shape: (num_envs,) - - -@jax.jit +@logged_jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: return jnp.clip(action, low, high) @@ -47,60 +53,113 @@ def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> fl return float(explained_var) -@jax.jit +@logged_jit def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations return learning_rate * frac def _get_action_and_value_noise( - sensor: GenericDenseLayersWithActivation, - feature_extractor: GenericDenseLayersWithActivation, - actor: Actor, - critic: OneDenseLayerMLP, + sensor: nn.Module, + feature_extractor: nn.Module, + actor: nn.Module, + critic: nn.Module, + message_passer: Optional[nn.Module], agent_state: TrainState, next_obs: jnp.ndarray, - key: jax.random.PRNGKey, + key, action_low, action_high, ): - hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) - hidden_critic = feature_extractor.apply( - agent_state.params["feature_extractor_params"], next_obs + # (B, n_nodes, feat) + 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"] + # (n_nodes, feat) --> let each node talk with its neighbours ==> vmap over B dimension + hidden = jax.vmap(lambda x: message_passer.apply(params, x))(hidden) + + hidden_critic = apply_shared( + feature_extractor, agent_state.params["feature_extractor_params"], next_obs ) - mean, log_std = actor.apply(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) std = jnp.exp(log_std) - raw_action = mean + noise * std - clipped_action = _clip_action(raw_action, action_low, action_high) - logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden_critic) - return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key + raw_action = mean + noise * std + flat_action = raw_action.reshape( + raw_action.shape[0], -1 + ) # concat the per agent, keep the envs dim (batch, agent * action) + flat_clipped_action = _clip_action(flat_action, action_low, action_high) + + logprob = -0.5 * (((raw_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum( + axis=(-2, -1) + ) + value = apply_shared(critic, agent_state.params["critic_params"], hidden_critic) + + return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key def _step_once( carry, _, env_step_fn, - sensor: GenericDenseLayersWithActivation, - feature_extractor: GenericDenseLayersWithActivation, - actor: Actor, - critic: OneDenseLayerMLP, + num_envs: int, + sensor: nn.Module, + feature_extractor: nn.Module, + actor: nn.Module, + critic: nn.Module, + message_passer: Optional[nn.Module], action_low, action_high, ): - agent_state, episode_stats, obs, done, key, env_state = carry - clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise( - sensor, feature_extractor, actor, critic, agent_state, obs, key, action_low, action_high + agent_state, episode_stats, obs, done, key, env_state, terminated_any, truncated_any = carry + flat_clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise( + sensor, + feature_extractor, + actor, + critic, + message_passer, + agent_state, + obs, + key, + action_low, + action_high, + ) + logger = get_logger() + + logger.debug(f"[_step_once] raw_action: {raw_action.shape}") + logger.debug(f"[_step_once] clipped_action: {flat_clipped_action.shape}") + + # Supporting signals (often where mismatch originates) + logger.debug(f"[_step_once] logprob: {logprob.shape}") + logger.debug(f"[_step_once] value: {value.shape}") + logger.debug(f"[_step_once] mean: {mean.shape}") + logger.debug(f"[_step_once] std: {std.shape}") + + key, reset_key = jax.random.split(key) + reset_rngs = jax.random.split(reset_key, num_envs) + + # ---- ENV STEP ---- + key, reset_key = jax.random.split(key) + reset_rngs = jax.random.split(reset_key, num_envs) + + episode_stats, env_state, (next_obs, reward, next_done, terminated, truncated) = env_step_fn( + episode_stats, + env_state, + flat_clipped_action, + reset_rngs, ) - episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( - episode_stats, env_state, clipped_action - ) + terminated_any = terminated_any | terminated + truncated_any = truncated_any | truncated + + logger.debug(f"[_step_once] next_obs: {next_obs.shape}") + logger.debug(f"[_step_once] reward: {reward.shape}") + logger.debug(f"[_step_once] next_done: {next_done.shape}") storage = Storage( obs=obs, @@ -115,11 +174,25 @@ def _step_once( returns=jnp.zeros_like(reward), advantages=jnp.zeros_like(reward), ) - return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage + return ( + agent_state, + episode_stats, + next_obs, + next_done, + key, + env_state, + terminated_any, + truncated_any, + ), storage -def _reward_fn(env_state, next_env_state): - # if delta distance positive ==> brittle star walking away from target +def reward_fn(env_state, next_env_state): + """Shaped reward used during training and checkpoint evaluation. + + Public so that ``evaluation.evaluate_mjx`` can import it and produce + metrics that are directly comparable to training-time returns. + """ + # Positive delta_distance means the brittle star is moving *away* from target. delta_distance = ( next_env_state.observations["xy_distance_to_target"] - env_state.observations["xy_distance_to_target"] @@ -135,12 +208,20 @@ def _reward_fn(env_state, next_env_state): return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty) -def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_processor): - next_env_state = env_step_fn(env_state, action) +def _step_env_wrapped( + episode_stats, + env_state, + action, + reset_rngs, + env_step_fn, + reset_single_fn, + obs_processor, +): + next_env_state_pre_reset = env_step_fn(env_state, action) - reward = _reward_fn(env_state, next_env_state) - terminated = next_env_state.terminated - truncated = next_env_state.truncated + reward = reward_fn(env_state, next_env_state_pre_reset) + terminated = next_env_state_pre_reset.terminated + truncated = next_env_state_pre_reset.truncated done = terminated | truncated new_episode_return = episode_stats.episode_returns + reward @@ -156,13 +237,50 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_process done, new_episode_length, episode_stats.returned_episode_lengths ), ) + + def _maybe_reset(state_i, rng_i, do_reset_i): + def _do(_): + reset_state = reset_single_fn(rng=rng_i) + + def _cast_leaf(new_leaf, like_leaf): + if like_leaf is None or new_leaf is None: + return new_leaf + + # Use jnp.asarray(...) to robustly get dtype for both JAX arrays and Python scalars. + like_dtype = jnp.asarray(like_leaf).dtype + + # Avoid unnecessary work when already matching. + if hasattr(new_leaf, "dtype") and new_leaf.dtype == like_dtype: + return new_leaf + + return jnp.asarray(new_leaf, dtype=like_dtype) + + # `lax.cond` requires both branches to return identical PyTree types/dtypes. + return jax.tree_util.tree_map(_cast_leaf, reset_state, state_i) + + def _dont(_): + return state_i + + return jax.lax.cond(do_reset_i, _do, _dont, operand=None) + + # Auto-reset done envs so rollouts continue with fresh episode initial states. + next_env_state = jax.vmap(_maybe_reset)(next_env_state_pre_reset, reset_rngs, done) + return ( episode_stats, next_env_state, - (obs_processor(next_env_state.observations), reward, done), + (obs_processor(next_env_state.observations), reward, done, terminated, truncated), ) +def apply_shared(net, params, x): + # x: (batch, nodes, feat) + # If the critic expects a single vector per environment: + batch_size = x.shape[0] + x_flattened = x.reshape(batch_size, -1) + return jax.vmap(lambda xi: net.apply(params, xi))(x_flattened) + + def _rollout_jit( agent_state, episode_stats, @@ -172,29 +290,67 @@ def _rollout_jit( key, max_steps, step_env_fn, - sensor: GenericDenseLayersWithActivation, - feature_extractor: GenericDenseLayersWithActivation, - actor: Actor, - critic: OneDenseLayerMLP, + num_envs: int, + sensor: nn.Module, + feature_extractor: nn.Module, + actor: nn.Module, + critic: nn.Module, + message_passer: Optional[nn.Module], action_low, action_high, ): - (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( + terminated_any0 = jnp.zeros((num_envs,), dtype=jnp.bool_) + truncated_any0 = jnp.zeros((num_envs,), dtype=jnp.bool_) + + ( + ( + agent_state, + episode_stats, + next_obs, + next_done, + key, + env_state, + terminated_any, + truncated_any, + ), + storage, + ) = jax.lax.scan( partial( _step_once, sensor=sensor, feature_extractor=feature_extractor, actor=actor, critic=critic, + message_passer=message_passer, env_step_fn=step_env_fn, + num_envs=num_envs, action_low=action_low, action_high=action_high, ), - (agent_state, episode_stats, next_obs, next_done, key, env_state), + ( + agent_state, + episode_stats, + next_obs, + next_done, + key, + env_state, + terminated_any0, + truncated_any0, + ), (), max_steps, ) - return agent_state, episode_stats, next_obs, next_done, storage, key, env_state + return ( + agent_state, + episode_stats, + next_obs, + next_done, + storage, + key, + env_state, + terminated_any, + truncated_any, + ) def _compute_gae_once(carry, inp, gamma, gae_lambda): @@ -217,9 +373,10 @@ def _compute_gae_jit( feature_extractor, critic, ): - next_value = critic.apply( + next_value = apply_shared( + critic, agent_state.params["critic_params"], - feature_extractor.apply(agent_state.params["feature_extractor_params"], next_obs), + apply_shared(feature_extractor, agent_state.params["feature_extractor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((num_envs,)) @@ -253,12 +410,17 @@ class TrainingMeasurements: class PPOTrainer: def __init__( - self, cfg: BrittleStarConfig, env: BrittleStarJaxEnvWrapper, run_dir: str, run_name: str + self, + cfg: BrittleStarConfig, + env: BrittleStarJaxEnvWrapper, + run_dir: str, + run_name: str, ): self.cfg = cfg self.ppo = cfg.ppo self.experiment = cfg.experiment self.logging_cfg = cfg.logging + self.evaluation_cfg = cfg.evaluation self.env = env self.run_dir = run_dir self.run_name = run_name @@ -270,39 +432,69 @@ class PPOTrainer: self.key = jax.random.PRNGKey(self.experiment.seed) + self.morph_mode = self.cfg.morphology.morph_mode + + self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32) + self.num_segments = self.segments_per_arm.sum().item() + self.num_arms = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item() + + self.logger.info(f"[INIT]: Used morphology mode {self.morph_mode}") + self.adj = build_adjacency(cfg.morphology.segments_per_arm, self.morph_mode) + + ( + self.sensor, + self.message_passer, + self.actor, + self.feature_extractor, + self.critic, + self.needed_copies, + self.agent_indices, + ) = self._init_agent() + + self.sensor.apply = logged_jit(self.sensor.apply) + self.feature_extractor.apply = logged_jit(self.feature_extractor.apply) + self.actor.apply = logged_jit(self.actor.apply) + self.critic.apply = logged_jit(self.critic.apply) + # Build the centralized observation processor: derive -> normalize -> pad -> flatten. self.obs_processor = create_obs_processor( bounds_dict=self.cfg.obs_bounds.to_bounds_dict(), + needed_copies=self.needed_copies, + num_arms=self.num_arms, + morph_mode=self.morph_mode, padding_masks=self.env.padding_masks, + segments_per_arm=self.segments_per_arm, + agent_indices=self.agent_indices, ) - self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent() - self.sensor.apply = jax.jit(self.sensor.apply) - self.feature_extractor.apply = jax.jit(self.feature_extractor.apply) - self.actor.apply = jax.jit(self.actor.apply) - self.critic.apply = jax.jit(self.critic.apply) + self.logger.debug(f"needed copies = {self.needed_copies}") action_low = jnp.asarray(self.env.single_action_space.low, dtype=jnp.float32) action_high = jnp.asarray(self.env.single_action_space.high, dtype=jnp.float32) + self._action_low = action_low + self._action_high = action_high - self._rollout_jit = jax.jit( + self._rollout_jit = logged_jit( partial( _rollout_jit, max_steps=self.ppo.num_steps, step_env_fn=partial( _step_env_wrapped, env_step_fn=self.env.step, + reset_single_fn=self.env.raw.reset, obs_processor=self.obs_processor, ), + num_envs=self.ppo.num_envs, sensor=self.sensor, feature_extractor=self.feature_extractor, actor=self.actor, critic=self.critic, + message_passer=self.message_passer, action_low=action_low, action_high=action_high, ) ) - self._compute_gae_jit = jax.jit( + self._compute_gae_jit = logged_jit( partial( _compute_gae_jit, num_envs=self.ppo.num_envs, @@ -313,13 +505,38 @@ class PPOTrainer: ) ) - self._ppo = PPO(self.ppo, self.sensor, self.actor, self.critic, self.feature_extractor) + def apply_sensor(p, x): + return apply_per_node(self.sensor.apply, p, x) + + def apply_actor(p, x): + return apply_per_node(self.actor.apply, p, x) + + def apply_critic(p, x): + return apply_shared(self.critic, p, x) + + def apply_feature(p, x): + return apply_shared(self.feature_extractor, p, x) + + def apply_message_passer(p, x): + assert self.message_passer is not None + return jax.vmap(lambda x_in: self.message_passer.apply(p, x_in))(x) + + self._ppo = PPO( + self.ppo, + apply_sensor, + apply_actor, + apply_critic, + apply_feature, + apply_message_passer if self.message_passer is not None else None, + ) self.agent_state = self._init_agent_state() self.episode_stats = self._init_episode_stats() self._init_random() + # Lazily-built JIT-compiled MJX eval rollout, created on first evaluation. + self._eval_fn = None def _init_random(self): self.logger.info(f"[RANDOM]: Setting random seed to {self.experiment.seed}") @@ -329,33 +546,136 @@ class PPOTrainer: def _init_agent(self): self.logger.info("[AGENT]: Initializing agent...") + agent_indices = [0, 1, 2, 3, 4] + match self.morph_mode: + case MorphMode.CENTRALIZED: + needed_copies = 1 + case MorphMode.FULLY_CONNECTED | MorphMode.RING: + agent_mask = self.segments_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0] + needed_copies = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item() + case MorphMode.SEGMENT: + agent_mask = self.segments_per_arm > 0 + agent_indices = jnp.where(agent_mask)[0] + needed_copies = ( + self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum() + ).item() + # scale actor output with size of model --> more models ==> less actions needed per model + actor = Actor(action_dim=self.env.single_action_space.shape[0] // needed_copies) sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) + message_passer: Optional[nn.Module] = ( + MessagePasser( + hidden_dim=300, + num_propagation_steps=self.cfg.architecture.message_passing_steps or 4, + adj_matrix=self.adj, + ) + if self.morph_mode != MorphMode.CENTRALIZED + else None + ) + feature_extractor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300]) - actor = Actor(action_dim=self.env.single_action_space.shape[0]) critic = OneDenseLayerMLP() - return sensor, feature_extractor, actor, critic + return ( + sensor, + message_passer, + actor, + feature_extractor, + critic, + needed_copies, + agent_indices, + ) def _init_agent_state(self) -> TrainState: self.logger.info("[AGENT STATE]: Initializing agent state...") - self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split( - self.key, 5 + self.key, sensor_key, actor_key, critic_key, feature_extractor_key, message_passer_key = ( + jax.random.split(self.key, 6) ) dummy_reset = self.env.reset(seed=0) + + for k, v in dummy_reset.observations.items(): + self.logger.debug(k, v.shape) + sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env - sensor_params = self.sensor.init(sensor_key, sample_obs) - feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs) - actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs)) - critic_params = self.critic.init( - critic_key, self.feature_extractor.apply(feature_extractor_params, sample_obs) + + self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}") + self.obs_mean = jnp.zeros((sample_obs.shape[-1],)) + self.obs_var = jnp.ones((sample_obs.shape[-1],)) + self.obs_count = 1e-4 + self.logger.debug(f"[_init_agent_state] obs_mean: {self.obs_mean.shape}") + self.logger.debug(f"[_init_agent_state] obs_var: {self.obs_var.shape}") + + self.logger.debug(f"[_init_agent_state]: Needed copies: {self.needed_copies}") + sensor_keys = jax.random.split(sensor_key, self.needed_copies) + actor_keys = jax.random.split(actor_key, self.needed_copies) + + # (needed_copies, X) + sensor_params = jax.vmap(lambda k: self.sensor.init(k, sample_obs))(sensor_keys) + self.logger.debug( + f"[_init_agent_state] sensor_params: {jax.tree.map(lambda x: x.shape, sensor_params)}" + ) + + single_sensor_param = jax.tree.map(lambda x: x[0], sensor_params) + self.logger.debug( + f"[_init_agent_state] single_sensor_param: { + jax.tree.map(lambda x: x.shape, single_sensor_param) + }" + ) + + sensor_params_sample = self.sensor.apply(single_sensor_param, sample_obs) + self.logger.debug( + f"[_init_agent_state] sensor_params_sample shape: {sensor_params_sample.shape}" + ) + + actor_params = jax.vmap(lambda k: self.actor.init(k, sensor_params_sample))(actor_keys) + self.logger.debug( + f"[_init_agent_state] actor_params: {jax.tree.map(lambda x: x.shape, actor_params)}" + ) + + message_passer_params = {} + if self.morph_mode != MorphMode.CENTRALIZED: + assert self.message_passer is not None, "decentralized modes require a message passer" + + message_passer_params = self.message_passer.init( + message_passer_key, + self.sensor.apply(single_sensor_param, sample_obs), + ) + self.logger.debug( + f"[_init_agent_state] message_passer_params: { + jax.tree.map(lambda x: x.shape, message_passer_params) + }" + ) + + flat_obs = sample_obs.reshape(-1) # BECAUSE 1 centralized critic + self.logger.debug(f"[_init_agent_state] flat_obs: {flat_obs.shape}") + + feature_extractor_params = self.feature_extractor.init(feature_extractor_key, flat_obs) + self.logger.debug( + f"[_init_agent_state] feature_extractor_params: { + jax.tree.map(lambda x: x.shape, feature_extractor_params) + }" + ) + + critic_input = self.feature_extractor.apply(feature_extractor_params, flat_obs) + self.logger.debug(f"[_init_agent_state] critic_input: {critic_input.shape}") + + critic_params = self.critic.init(critic_key, critic_input) + self.logger.debug( + f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}" ) return TrainState.create( apply_fn=None, params=asdict( - AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) + AgentParams( + sensor_params, + actor_params, + critic_params, + feature_extractor_params, + message_passer_params, + ) ), tx=optax.chain( optax.clip_by_global_norm(self.ppo.max_grad_norm), @@ -458,7 +778,7 @@ class PPOTrainer: def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple: if iteration == 1: self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}") - + self.logger.debug(f"[_step] next_obs (in): {next_obs.shape}") ( self.agent_state, self.episode_stats, @@ -467,13 +787,15 @@ class PPOTrainer: storage, self.key, next_env_state, + terminated_any, + truncated_any, ) = self._rollout(env_state, next_obs, next_done) - + self.logger.debug(f"[_step] next_obs (post-rollout): {next_obs.shape}") if iteration == 1: self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}") storage = self._compute_gae(storage, next_obs, next_done) - + self.logger.debug(f"[_step] storage.obs (post-gae): {storage.obs.shape}") if iteration == 1: self.logger.log_non_interactive(f"Starting first PPO update (JIT): {time.ctime()}") @@ -490,8 +812,8 @@ class PPOTrainer: explained_var = _compute_explained_variance(storage.values, storage.returns) - terminated = next_env_state.terminated - truncated = next_env_state.truncated + terminated = terminated_any + truncated = truncated_any episode_lengths = self.episode_stats.returned_episode_lengths num_terminated = int(jnp.sum(terminated).item()) @@ -538,6 +860,65 @@ class PPOTrainer: params=self.agent_state.params, step=iteration, metadata=asdict(self.cfg) ) + def _evaluate_checkpoint(self, iteration: int, *, trained_timesteps: int) -> None: + """Evaluate the current checkpoint and persist metrics to CSV. + + Delegates all evaluation logic to `evaluation.evaluate_mjx`. + Best-effort: a failure here must never abort training. + """ + if not self.evaluation_cfg.evaluate_checkpoints: + return + + max_steps = int(self.evaluation_cfg.eval_max_steps) + seed = int(self.evaluation_cfg.eval_seed) + + if max_steps <= 0: + self.logger.warning("[EVAL]: eval_max_steps must be > 0; skipping evaluation") + return + + if not self.logging_cfg.save_checkpoints or self.logging_cfg.checkpoint_frequency <= 0: + self.logger.warning( + "[EVAL]: evaluate_checkpoints is enabled but checkpoint saving is disabled; " + "skipping evaluation" + ) + return + + try: + if self._eval_fn is None: + if getattr(self.env, "backend", None) != Backend.MJX: + self.logger.warning( + f"[EVAL]: Training env backend is {self.env.backend}; " + "MJX evaluation may be unavailable/slow." + ) + 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.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 + ), + action_low=self._action_low, + action_high=self._action_high, + reward_fn=reward_fn, + ) + + result = evaluate_checkpoint_mjx( + self._eval_fn, + self.agent_state.params, + seed=seed, + max_steps=max_steps, + ) + csv_path = append_checkpoint_eval_row( + self.run_dir, + iteration=iteration, + trained_timesteps=int(trained_timesteps), + result=result, + ) + self.logger.sync_file(csv_path) + except Exception as e: + self.logger.warning(f"[EVAL]: Checkpoint evaluation failed: {e}") + def train(self): """ Train the PPO agent for a specified number of iterations. @@ -549,7 +930,10 @@ class PPOTrainer: self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}") env_state = self.env.reset(seed=self.experiment.seed) + next_obs = self.obs_processor(env_state.observations) + self.logger.debug(f"[train] next_obs: {next_obs.shape}") + next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_) self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}") @@ -591,6 +975,7 @@ class PPOTrainer: if self.logging_cfg.save_checkpoints and self.logging_cfg.checkpoint_frequency > 0: if iteration % self.logging_cfg.checkpoint_frequency == 0: self._save_checkpoint(iteration) + self._evaluate_checkpoint(iteration, trained_timesteps=global_step) if getattr(self.cfg.experiment, "debug_sanity", False): self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch") diff --git a/src/brittle_star_project/utils/__init__.py b/src/brittle_star_project/utils/__init__.py new file mode 100644 index 0000000..ed72e5a --- /dev/null +++ b/src/brittle_star_project/utils/__init__.py @@ -0,0 +1,3 @@ +from .logged_jit import logged_jit + +__all__ = ["logged_jit"] diff --git a/src/brittle_star_project/utils/logged_jit.py b/src/brittle_star_project/utils/logged_jit.py new file mode 100644 index 0000000..3d29ba6 --- /dev/null +++ b/src/brittle_star_project/utils/logged_jit.py @@ -0,0 +1,17 @@ +import jax +from experiment_logger import get_logger + + +def logged_jit(fn, **jit_kwargs): + logger = get_logger() + name = getattr(fn, "__name__", getattr(fn, "__qualname__", repr(fn))) + + def decorator(func): + def traced_func(*args, **kwargs): + logger.debug(f"[JIT] Compiling {name}...") + return func(*args, **kwargs) + + jitted = jax.jit(traced_func, **jit_kwargs) + return jitted + + return decorator(fn) diff --git a/src/experiment_logger/config_logger.py b/src/experiment_logger/config_logger.py index c34c162..fd77a28 100644 --- a/src/experiment_logger/config_logger.py +++ b/src/experiment_logger/config_logger.py @@ -31,3 +31,6 @@ class LoggingConfig: "Configuration Error: 'upload_checkpoints' is True, but it requires " "both 'track' and 'save_checkpoints' to also be True." ) + + # NOTE: Checkpoint evaluation settings live under the project's + # `evaluation` config group (see brittle_star_project.configs). diff --git a/src/experiment_logger/simple_logger.py b/src/experiment_logger/simple_logger.py index 7e4a816..71844dc 100644 --- a/src/experiment_logger/simple_logger.py +++ b/src/experiment_logger/simple_logger.py @@ -71,6 +71,10 @@ class SimpleLogger: def save_final_model(self, params: Any, metadata: Optional[Dict[str, Any]] = None): print("[SAVE] Final model would be saved (SimpleLogger: No-Op)") + def sync_file(self, path: Any): + """No-op for SimpleLogger.""" + pass + def finish(self): print(f"[FINISH] SimpleLogger finished for run: {self.run_name}") diff --git a/src/experiment_logger/unified_logger.py b/src/experiment_logger/unified_logger.py index f37136b..e308d12 100644 --- a/src/experiment_logger/unified_logger.py +++ b/src/experiment_logger/unified_logger.py @@ -432,6 +432,21 @@ class UnifiedLogger: except Exception as e: self.error(f"Error saving final model: {e}") + def sync_file(self, path: Path) -> None: + """Upload a file to W&B if tracking is enabled. + + Best-effort: logs a warning on failure, never raises. + """ + if self.wandb_run is None: + return + try: + import wandb + + # "Simple sync" behavior: wandb will copy this file into the run. + wandb.save(str(path), base_path=str(path.parent)) + except Exception as e: + self.warning(f"Failed to sync file to W&B: {e}") + def finish(self): """Finalize logging and cleanup.""" # Flush remaining metrics diff --git a/tests/test_adjacency.py b/tests/test_adjacency.py new file mode 100644 index 0000000..27c8b43 --- /dev/null +++ b/tests/test_adjacency.py @@ -0,0 +1,98 @@ +import jax.numpy as jnp +import numpy as np + +from brittle_star_project.MLPs import build_adjacency +from brittle_star_project.environment.env_config import MorphMode + + +def assert_symmetric(adj): + assert jnp.all(adj == adj.T) + + +def test_centralized(): + adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.CENTRALIZED) + + assert adj.shape == (1, 1) + assert adj[0, 0] == 1 + + +def test_fully_connected(): + adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.FULLY_CONNECTED) + + assert adj.shape == (5, 5) + assert jnp.all(adj == 1) + + +def test_ring(): + adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.RING) + + assert adj.shape == (5, 5) + assert_symmetric(adj) + + # each node should connect to itself + 2 neighbors + for node in range(5): + assert adj[node, node] == 1 + assert jnp.sum(adj[node]) == 3 + neighbor1 = (node - 1) % 5 + neighbor2 = (node + 1) % 5 + assert adj[neighbor1, node] == 1 + assert adj[node, neighbor2] == 1 # Symmetrical + + +def test_segment_structure(): + segments = [4, 4, 4, 4, 4] + adj = build_adjacency(segments, MorphMode.SEGMENT) + + num_arms = 5 + num_segments = sum(segments) + num_nodes = num_arms + num_segments + + assert adj.shape == (num_nodes, num_nodes) + + # --- ring connectivity --- + for i in range(num_arms): + assert adj[i, i] == 1 + assert adj[i, (i - 1) % num_arms] == 1 + assert adj[i, (i + 1) % num_arms] == 1 + + # --- segment chain checks --- + offset = num_arms + for arm in range(5): + for i in range(4): + node = offset + arm * 4 + i + + # self + assert adj[node, node] == 1 + + # chain neighbors + if i > 0: + assert adj[node, node - 1] == 1 + if i < 3: + assert adj[node, node + 1] == 1 + + # --- ring ↔ segment connections --- + for arm in range(5): + first_seg = num_arms + arm * 4 + assert adj[arm, first_seg] == 1 + assert adj[first_seg, arm] == 1 + + save_adj(adj) + + +def save_adj(adj, name="adjacency_debug.txt"): + a = np.array(adj) + + with open(name, "w") as f: + f.write("\nAdjacency matrix:\n") + f.write(" " + " ".join([f"{i:2d}" for i in range(a.shape[0])]) + "\n") + + for i, row in enumerate(a): + line = f"{i:2d} " + " ".join(["█" if x > 0 else "." for x in row]) + f.write(line + "\n") + + +def test_no_isolated_nodes(): + adj = build_adjacency([4, 4, 4, 4, 4], MorphMode.SEGMENT) + + # no node should be completely isolated + assert jnp.all(jnp.sum(adj, axis=0) > 0) 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 e1a7a1a..ce189f4 100644 --- a/tests/test_network_shapes.py +++ b/tests/test_network_shapes.py @@ -1,5 +1,6 @@ import jax import jax.numpy as jnp +from brittle_star_project.environment.env_config import MorphMode from brittle_star_project.environment.padded_obs_wrapper import ( compute_padding_masks, ) @@ -20,17 +21,29 @@ def test_centralized_forward_pass_with_padding(): "segment_contact": jnp.zeros((batch_size, 14)), } + segments_per_arm = jnp.array((4, 0, 4, 2, 4)) + num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum().item() + # 2. Process and Pad Observation - masks = compute_padding_masks(segments_per_arm=(4, 0, 4, 2, 4)) - obs_processor = create_obs_processor(bounds_dict={}, padding_masks=masks) + masks = compute_padding_masks(segments_per_arm=list(segments_per_arm)) + obs_processor = create_obs_processor( + bounds_dict={}, + needed_copies=1, + num_arms=num_arms, + padding_masks=masks, + morph_mode=MorphMode.CENTRALIZED, + segments_per_arm=segments_per_arm, + ) global_state = obs_processor(amputated_obs) - # 40 + 40 + 20 = 100 dimensions - assert global_state.shape == (batch_size, 100), ( - f"Expected global state shape (2, 100), 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 @@ -45,9 +58,11 @@ def test_centralized_forward_pass_with_padding(): action_mean, action_log_std = actor.apply(actor_params, global_state) value = critic.apply(critic_params, global_state) - assert action_mean.shape == (batch_size, 40), f"Actor mean shape mismatch: {action_mean.shape}" + assert action_mean.shape == (batch_size, 1, 40), ( + f"Actor mean shape mismatch: {action_mean.shape}" + ) assert action_log_std.shape == (40,), f"Actor log_std shape mismatch: {action_log_std.shape}" - assert value.shape == (batch_size, 1) or value.shape == (batch_size,), ( + assert value.shape == (batch_size, 1, 1) or value.shape == (batch_size,), ( f"Critic value shape mismatch: {value.shape}" ) diff --git a/tests/test_obs_processor.py b/tests/test_obs_processor.py new file mode 100644 index 0000000..7008648 --- /dev/null +++ b/tests/test_obs_processor.py @@ -0,0 +1,124 @@ +import jax +import jax.numpy as jnp + +from brittle_star_project.environment.obs_processing import create_obs_processor +from brittle_star_project.environment.env_config import MorphMode, ObservationBoundsConfig + + +obs_bounds = ObservationBoundsConfig().to_bounds_dict() + +# 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 +JOINTS_PER_SEG = 2 # from _build_joint_indices: segs * 2 + +SEGS_HEALTHY = [4, 4, 4, 4, 4] +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 + 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: + total_segs = sum(segs_per_arm) + total_joints = JOINTS_PER_SEG * total_segs + + return { + "actuator_force": jnp.ones(total_joints), + "disk_angular_velocity": jnp.zeros(3), + "disk_linear_velocity": jnp.zeros(3), + "disk_position": jnp.zeros(3), + "disk_rotation": jnp.array([0.1, 0.1, 0.5]), # (roll, pitch, yaw) + "joint_actuator_force": jnp.full(total_joints, 1.0), + "joint_position": jnp.full(total_joints, 0.5), + "joint_velocity": jnp.full(total_joints, 2.0), + "segment_contact": jnp.ones(total_segs), + "tendon_position": jnp.zeros(0), + "tendon_velocity": jnp.zeros(0), + "unit_xy_direction_to_target": jnp.array([1.0, 0.0]), + "xy_distance_to_target": jnp.array([3.5]), + } + + +def batch_obs(obs: dict): + return jax.tree_util.tree_map(lambda x: x[None, :], obs) + + +def make_processor(morph_mode: MorphMode, needed_copies: int, segments_per_arm: list[int]): + return create_obs_processor( + bounds_dict=obs_bounds, + num_arms=NUM_ARMS, + needed_copies=needed_copies, + morph_mode=morph_mode, + segments_per_arm=segments_per_arm, + agent_indices=AGENT_INDICES, + ) + + +def test_centralized_no_damage(): + proc = make_processor(MorphMode.CENTRALIZED, 1, SEGS_HEALTHY) + obs = make_obs(SEGS_HEALTHY) + obs = batch_obs(obs) + global_state = proc(obs) + + # Centralized: 5 agents flattened into 1 → shape (1, 1, 155) + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) + + +def test_centralized_damaged_1_arm(): + proc = make_processor(MorphMode.CENTRALIZED, 1, SEGS_DAMAGED) + obs = make_obs(SEGS_DAMAGED) + obs = batch_obs(obs) + global_state = proc(obs) + + # shape test + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) + + +def test_centralized_damaged_2_arms(): + proc = make_processor(MorphMode.CENTRALIZED, 1, SEGS_DAMAGED_2) + obs = make_obs(SEGS_DAMAGED_2) + obs = batch_obs(obs) + global_state = proc(obs) + + # shape test + assert global_state.shape == (1, 1, FEAT_CENTRALIZED) + + +def test_decentralized_fully_connected_no_damage(): + proc = make_processor(MorphMode.FULLY_CONNECTED, NUM_ARMS, SEGS_HEALTHY) + obs = make_obs(SEGS_HEALTHY) + obs = batch_obs(obs) + global_state = proc(obs) + + # shape test + assert global_state.shape == (1, NUM_ARMS, FEAT_PER_AGENT) + + +def test_decentralized_fully_connected_damaged_1_arm(): + proc = make_processor(MorphMode.FULLY_CONNECTED, NUM_ARMS, SEGS_DAMAGED) + obs = make_obs(SEGS_DAMAGED) + obs = batch_obs(obs) + global_state = proc(obs) + + # shape test + assert global_state.shape == (1, NUM_ARMS, FEAT_PER_AGENT) diff --git a/tests/test_target_direction.py b/tests/test_target_direction.py index a81cdf1..8317fb2 100644 --- a/tests/test_target_direction.py +++ b/tests/test_target_direction.py @@ -1,6 +1,7 @@ import jax.numpy as jnp from brittle_star_project.configs.main_config import BrittleStarConfig +from brittle_star_project.environment.env_config import MorphMode from brittle_star_project.environment.env_types import Backend from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.obs_processing import create_obs_processor @@ -43,8 +44,16 @@ def test_processor_converts_to_egocentric_direction(): cfg = BrittleStarConfig() env = BrittleStarJaxEnvWrapper.default(num_envs=1, backend=Backend.MJX) + segments_per_arm = jnp.array((4, 4, 4, 4, 4)) + num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum().item() + obs_processor = create_obs_processor( - bounds_dict=cfg.obs_bounds.to_bounds_dict(), padding_masks=env.padding_masks + bounds_dict=cfg.obs_bounds.to_bounds_dict(), + needed_copies=1, + num_arms=num_arms, + padding_masks=env.padding_masks, + morph_mode=MorphMode.CENTRALIZED, + segments_per_arm=segments_per_arm, ) env_state = env.reset(seed=42) @@ -66,10 +75,13 @@ def test_processor_converts_to_egocentric_direction(): processed_2 = obs_processor(dummy_obs_2) # Find the indices of the elements that changed - diff_array = jnp.abs(processed_1[0] - processed_2[0]) + diff_array = jnp.abs(processed_1[0, 0] - processed_2[0, 0]) changed_indices = jnp.where(diff_array > 1e-4)[0] - local_target = processed_1[0, changed_indices] + # (143,) + local_target = processed_1[0, 0, changed_indices] + + # (2,) expected_local_target = jnp.array([0.0, -1.0]) assert jnp.sum(jnp.abs(local_target - expected_local_target)) < 1e-4, (