Deployed e4869e0 with MkDocs version: 1.6.1
This commit is contained in:
parent
fd3dbe898a
commit
26e0b9ee28
75 changed files with 13749 additions and 5 deletions
274
404.html
274
404.html
|
|
@ -732,6 +732,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="/scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="/src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -843,6 +843,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
@ -964,7 +1238,7 @@
|
|||
<ul>
|
||||
<li><strong>Git Practices:</strong> Commits must be frequent and small. Each commit should relate to exactly one piece of functionality.</li>
|
||||
<li><strong>Branching Strategy:</strong> The <code>dev</code> branch serves as the integration branch for pushing and merging code. Only stable releases may be pushed to the <code>main</code> branch.</li>
|
||||
<li><strong>Artifact Management:</strong> Data files, trained models, and large datasets must never be committed directly to Git. Git Large File Storage (LFS) must be used for tracking large files. <strong>All developers must have <code>git-lfs</code> installed locally</strong> (see <code>DEVELOPMENT.md</code> for setup).</li>
|
||||
<li><strong>Artifact Management:</strong> Data files, trained models, and large datasets must never be committed directly to Git. Git Large File Storage (LFS) must be used for tracking large files. <strong>All developers must have <code>git-lfs</code> installed locally</strong> (see <a href="../DEVELOPMENT/">DEVELOPMENT.md</a> for setup).</li>
|
||||
<li><strong>Repository Layout:</strong> The repository must maintain the following core directories: <code>src/</code> for algorithms, <code>env/</code> for MuJoCo wrappers, <code>config/</code> for experiment configurations, <code>experiments/</code> for scripts, <code>docs/</code> for Doxygen or ReadTheDocs documentation, and <code>tests/</code> for unit tests.</li>
|
||||
</ul>
|
||||
<h2 id="4-architecture-tooling">4. Architecture & Tooling</h2>
|
||||
|
|
|
|||
|
|
@ -910,6 +910,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
@ -1127,7 +1401,7 @@ In the devcontainer, this will succeed on both CPU and GPU. A <code>GpuDevice</c
|
|||
<h2 id="logging-monitoring">Logging & Monitoring</h2>
|
||||
<p>This project uses a unified logging system through the <code>experiment_logger</code> package. </p>
|
||||
<ul>
|
||||
<li><strong>Usage in Code</strong>: To use the logger in your scripts, refer to the <a href="../src/experiment_logger/README.md">package README</a> for the API reference.</li>
|
||||
<li><strong>Usage in Code</strong>: To use the logger in your scripts, refer to the <a href="../src/experiment_logger/">package README</a> for the API reference.</li>
|
||||
<li><strong>WandB/TensorBoard Setup</strong>: For information on how to configure tracking for experiments, see the <a href="../api/tracking/">Tracking & Monitoring API Guide</a>.</li>
|
||||
</ul>
|
||||
<p>The logger automatically detects if it is running in an interactive terminal or a non-interactive environment (like an HPC Slurm job), adjusting progress bars and fallback modes accordingly.</p>
|
||||
|
|
|
|||
274
HPC/index.html
274
HPC/index.html
|
|
@ -888,6 +888,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -940,6 +940,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -823,6 +823,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -857,6 +857,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -834,6 +834,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -857,6 +857,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
@ -998,7 +1272,7 @@
|
|||
<p>For quick diagnostics or to export data to CSV without launching the full TensorBoard UI, you can use the <code>explore_tensorboard.py</code> script:</p>
|
||||
<div class="highlight"><pre><span></span><code>uv<span class="w"> </span>run<span class="w"> </span>python<span class="w"> </span>scripts/analysis/explore_tensorboard.py<span class="w"> </span>runs/your_run_name/
|
||||
</code></pre></div>
|
||||
<p>See the detailed description in <a href="../../scripts/analysis/README.md"><code>/scripts/analysis/README.md</code></a>.</p>
|
||||
<p>See the detailed description in <a href="../../scripts/analysis/"><code>/scripts/analysis/README.md</code></a>.</p>
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -857,6 +857,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -801,6 +801,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -812,6 +812,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -801,6 +801,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -834,6 +834,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -812,6 +812,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -12,6 +12,8 @@
|
|||
<link rel="prev" href="../learning_algorithm/">
|
||||
|
||||
|
||||
<link rel="next" href="../../scripts/analysis/">
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
@ -821,6 +823,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="../../src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
|
|||
276
index.html
276
index.html
|
|
@ -808,6 +808,280 @@
|
|||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7" id="__nav_7_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_7_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Scripts
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_7_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_7_1" id="__nav_7_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_7_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_7_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Analysis
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="scripts/analysis/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Analysis Tools
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8" id="__nav_8_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="1" aria-labelledby="__nav_8_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Src
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item md-nav__item--nested">
|
||||
|
||||
|
||||
|
||||
<input class="md-nav__toggle md-toggle " type="checkbox" id="__nav_8_1" >
|
||||
|
||||
|
||||
<label class="md-nav__link" for="__nav_8_1" id="__nav_8_1_label" tabindex="0">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
</label>
|
||||
|
||||
<nav class="md-nav" data-md-level="2" aria-labelledby="__nav_8_1_label" aria-expanded="false">
|
||||
<label class="md-nav__title" for="__nav_8_1">
|
||||
<span class="md-nav__icon md-icon"></span>
|
||||
|
||||
|
||||
Experiment logger
|
||||
|
||||
|
||||
</label>
|
||||
<ul class="md-nav__list" data-md-scrollfix>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
<li class="md-nav__item">
|
||||
<a href="src/experiment_logger/" class="md-nav__link">
|
||||
|
||||
|
||||
|
||||
<span class="md-ellipsis">
|
||||
|
||||
|
||||
Experiment Logger
|
||||
|
||||
|
||||
|
||||
</span>
|
||||
|
||||
|
||||
|
||||
</a>
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
|
||||
</li>
|
||||
|
||||
|
||||
|
||||
</ul>
|
||||
</nav>
|
||||
</div>
|
||||
|
|
@ -882,7 +1156,7 @@
|
|||
<li><a href="design/controllers/">Controllers</a>: Macroscopig brain toplogy, centralized, arm-level, segment-level.</li>
|
||||
<li><a href="design/input_action_spaces/">Input/output</a>: Description of the model's input and output.</li>
|
||||
<li><a href="design/learning_algorithm/">Learning algorithm</a>: RL techniques, i.e. PPO.</li>
|
||||
<li><a href="design/learning_algorithm/">Reward function</a>: Goals, fitness tracking, and reward structures.</li>
|
||||
<li><a href="design/reward_function/">Reward function</a>: Goals, fitness tracking, and reward structures.</li>
|
||||
</ul>
|
||||
<h2 id="api-reference-api">API reference (<code>/api</code>)</h2>
|
||||
<p>If you are interested in the "how do I use it?"</p>
|
||||
|
|
|
|||
143
scripts/analysis/explore_tensorboard.py
Normal file
143
scripts/analysis/explore_tensorboard.py
Normal file
|
|
@ -0,0 +1,143 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Reproducible CLI tool to explore TensorBoard logs.
|
||||
Designed for both local development and HPC diagnostics.
|
||||
|
||||
Requirements:
|
||||
pip install tensorboard
|
||||
|
||||
Usage:
|
||||
python explore_tensorboard.py <path_to_run_directory> [--csv output.csv]
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import csv
|
||||
|
||||
try:
|
||||
from tensorboard.backend.event_processing import event_accumulator
|
||||
except ImportError:
|
||||
print("Error: Missing dependency. Please run: pip install tensorboard")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def explore_run(log_dir):
|
||||
"""
|
||||
Extracts and displays a summary of scalar metrics from a TensorBoard log directory.
|
||||
"""
|
||||
print(f"\n{'=' * 20} Exploring Run {'=' * 20}")
|
||||
print(f"Directory: {log_dir}")
|
||||
print(f"{'=' * 55}\n")
|
||||
|
||||
if not os.path.exists(log_dir):
|
||||
print(f"Error: Directory '{log_dir}' does not exist.")
|
||||
return None
|
||||
|
||||
# Initialize EventAccumulator
|
||||
# size_guidance=0 loads all data points for each tag.
|
||||
ea = event_accumulator.EventAccumulator(
|
||||
log_dir,
|
||||
size_guidance={
|
||||
event_accumulator.SCALARS: 0,
|
||||
event_accumulator.TENSORS: 0,
|
||||
},
|
||||
)
|
||||
|
||||
print("Loading event files (this may take a moment for large runs)...")
|
||||
ea.Reload()
|
||||
|
||||
tags = ea.Tags()
|
||||
scalar_tags = tags.get("scalars", [])
|
||||
|
||||
if not scalar_tags:
|
||||
print("No scalar metrics found in this directory.")
|
||||
return None
|
||||
|
||||
print(f"Found {len(scalar_tags)} scalar metrics.\n")
|
||||
|
||||
data = {}
|
||||
summary = []
|
||||
|
||||
# Process scalar values
|
||||
for tag in scalar_tags:
|
||||
events = ea.Scalars(tag)
|
||||
if not events:
|
||||
continue
|
||||
|
||||
values = [e.value for e in events]
|
||||
last_event = events[-1]
|
||||
data[tag] = values
|
||||
|
||||
summary.append(
|
||||
{
|
||||
"Metric": tag,
|
||||
"Steps": len(events),
|
||||
"Last Value": f"{last_event.value:.4f}",
|
||||
"Max": f"{max(values):.4f}",
|
||||
"Min": f"{min(values):.4f}",
|
||||
}
|
||||
)
|
||||
|
||||
# Display summary table formatted manually
|
||||
summary = sorted(summary, key=lambda x: x["Metric"])
|
||||
print(f"{'Metric':<30} {'Steps':>10} {'Last':>12} {'Max':>12} {'Min':>12}")
|
||||
print("-" * 80)
|
||||
for row in summary:
|
||||
print(
|
||||
f"{row['Metric']:<30} {row['Steps']:>10} {row['Last Value']:>12} "
|
||||
f"{row['Max']:>12} {row['Min']:>12}"
|
||||
)
|
||||
|
||||
# Calculate and display global metadata
|
||||
if "charts/SPS" in data:
|
||||
sps_events = ea.Scalars("charts/SPS")
|
||||
if len(sps_events) > 1:
|
||||
total_duration_hours = (sps_events[-1].wall_time - sps_events[0].wall_time) / 3600
|
||||
print(f"\nTotal Recorded Duration: {total_duration_hours:.2f} hours")
|
||||
|
||||
# Estimate completion if total_timesteps is available in hyperparameters
|
||||
try:
|
||||
hp_tags = [t for t in tags.get("tensors", []) if "hyperparameters" in t]
|
||||
if hp_tags:
|
||||
hp_event = ea.Tensors(hp_tags[0])[0]
|
||||
hp_text = hp_event.tensor_proto.string_val[0].decode("utf-8")
|
||||
if "total_timesteps" in hp_text:
|
||||
for line in hp_text.split("\n"):
|
||||
if "total_timesteps" in line:
|
||||
target = int(line.split("|")[2].strip())
|
||||
current = ea.Scalars(scalar_tags[0])[-1].step
|
||||
percent = (current / target) * 100
|
||||
print(f"Progress: {current:,} / {target:,} steps ({percent:.1f}%)")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return data
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Reproducible TensorBoard exploration tool.")
|
||||
parser.add_argument("log_dir", help="Path to the TensorBoard run directory.")
|
||||
parser.add_argument("--csv", help="Optional: Path to export scalar data to CSV.", default=None)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
scalar_data = explore_run(args.log_dir)
|
||||
|
||||
if args.csv and scalar_data:
|
||||
# Reloading for wall_time and steps
|
||||
ea = event_accumulator.EventAccumulator(args.log_dir).Reload()
|
||||
with open(args.csv, mode="w", newline="") as f:
|
||||
writer = csv.DictWriter(f, fieldnames=["tag", "step", "value", "wall_time"])
|
||||
writer.writeheader()
|
||||
for tag in scalar_data.keys():
|
||||
for e in ea.Scalars(tag):
|
||||
writer.writerow(
|
||||
{"tag": tag, "step": e.step, "value": e.value, "wall_time": e.wall_time}
|
||||
)
|
||||
|
||||
print(f"\nData exported to: {args.csv}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
1244
scripts/analysis/index.html
Normal file
1244
scripts/analysis/index.html
Normal file
File diff suppressed because it is too large
Load diff
182
scripts/compare_models.py
Normal file
182
scripts/compare_models.py
Normal file
|
|
@ -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()
|
||||
264
scripts/evaluate_checkpoints.py
Normal file
264
scripts/evaluate_checkpoints.py
Normal file
|
|
@ -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()
|
||||
83
scripts/hpc/export_requirements.py
Normal file
83
scripts/hpc/export_requirements.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Export HPC pip requirements from pyproject.toml.
|
||||
|
||||
This is a LOCAL DEVELOPER UTILITY — run it on your own machine before pushing
|
||||
code whenever pyproject.toml dependencies change. It reads the modules from
|
||||
env/hpc/modules.txt and the full dependency list from pyproject.toml, then
|
||||
writes the remainder to env/hpc/requirements.txt.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
|
||||
|
||||
def normalise(name: str) -> str:
|
||||
"""Normalise a PyPI package name for comparison."""
|
||||
return re.sub(r"[-_.]+", "-", name).lower()
|
||||
|
||||
|
||||
def pkg_name(dep: str) -> str:
|
||||
"""Extract the bare package name from a PEP 508 dependency string."""
|
||||
return re.split(r"[\[=><~!;]", dep)[0].strip()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
import tomllib
|
||||
|
||||
modules_path = ROOT / "env" / "hpc" / "modules.txt"
|
||||
if not modules_path.exists():
|
||||
print(f"Error: {modules_path} not found.", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
# Read normalized module names from base modules only
|
||||
# Library modules (like PyTorch) are kept in requirements for portability
|
||||
module_names = [
|
||||
normalise(line.split()[0].split("/")[0])
|
||||
for line in modules_path.read_text().splitlines()
|
||||
if line.strip() and not line.startswith("#")
|
||||
]
|
||||
|
||||
pyproject_path = ROOT / "pyproject.toml"
|
||||
with pyproject_path.open("rb") as f:
|
||||
data = tomllib.load(f)
|
||||
|
||||
# Collect all dependencies, merging 'cuda' extras into base dependencies
|
||||
dep_dict: dict[str, str] = {}
|
||||
for dep in data.get("project", {}).get("dependencies", []):
|
||||
dep_dict[normalise(pkg_name(dep))] = dep
|
||||
|
||||
# Add cuda extras (takes precedence for HPC)
|
||||
optional_deps = data.get("project", {}).get("optional-dependencies", {})
|
||||
for group in ["cuda"]:
|
||||
for dep in optional_deps.get(group, []):
|
||||
dep_dict[normalise(pkg_name(dep))] = dep
|
||||
|
||||
deps = list(dep_dict.values())
|
||||
|
||||
final_deps: list[str] = []
|
||||
print("Checking dependencies against HPC module list...", file=sys.stderr)
|
||||
for dep in deps:
|
||||
name = normalise(pkg_name(dep))
|
||||
# Smart check: if the package name is a substring of any loaded module name
|
||||
# (e.g. 'torch' in 'pytorch', 'scipy' in 'scipy-bundle')
|
||||
if any(name in mod for mod in module_names):
|
||||
print(f" [skip – module provider found] {dep}", file=sys.stderr)
|
||||
continue
|
||||
|
||||
final_deps.append(dep)
|
||||
print(f" [pip] {dep}", file=sys.stderr)
|
||||
|
||||
hpc_dir = ROOT / "env" / "hpc"
|
||||
output_path = hpc_dir / "requirements.txt"
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
output_path.write_text("\n".join(final_deps) + "\n")
|
||||
print(f"\nWrote {len(final_deps)} requirement(s) to {output_path}", file=sys.stderr)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
54
scripts/hpc/install.sh
Normal file
54
scripts/hpc/install.sh
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
#!/bin/bash -l
|
||||
# scripts/hpc/install.sh
|
||||
#
|
||||
# Usage (on any compute node):
|
||||
# bash scripts/hpc/install.sh
|
||||
#
|
||||
# Batch usage:
|
||||
# qsub scripts/hpc/install.sh
|
||||
|
||||
#PBS -N brittlestar-install
|
||||
#PBS -l walltime=00:15:00
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# Preliminary status echo
|
||||
echo ">>> Starting installation job $PBS_JOBID on $(hostname)..."
|
||||
|
||||
if [ -n "$PBS_O_WORKDIR" ]; then
|
||||
cd "$PBS_O_WORKDIR"
|
||||
fi
|
||||
|
||||
mkdir -p "${PBS_O_WORKDIR}/runs"
|
||||
|
||||
# Mirror configs to $VSC_DATA to avoid home quota limits (3GB)
|
||||
# vsc-venv manages environments relative to the requirements file
|
||||
PROJ_NAME=$(basename "$PWD")
|
||||
HPC_CONFIG_DIR="$VSC_DATA/$PROJ_NAME/env/hpc"
|
||||
mkdir -p "$HPC_CONFIG_DIR"
|
||||
cp env/hpc/*.txt "$HPC_CONFIG_DIR/"
|
||||
|
||||
# Keep caches off $VSC_HOME (quota ~3 GB).
|
||||
export PIP_CACHE_DIR="$VSC_SCRATCH/.cache/pip"
|
||||
export UV_CACHE_DIR="$VSC_SCRATCH/.cache/uv"
|
||||
mkdir -p "$PIP_CACHE_DIR" "$UV_CACHE_DIR"
|
||||
|
||||
module load vsc-venv
|
||||
|
||||
echo ">>> Synchronizing and activating environment (vsc-venv)..."
|
||||
# cd to $VSC_DATA so vsc-venv creates its venvs/ directory there, not in $HOME.
|
||||
mkdir -p "$VSC_DATA/$PROJ_NAME"
|
||||
cd "$VSC_DATA/$PROJ_NAME"
|
||||
set +euo pipefail
|
||||
source vsc-venv --activate \
|
||||
--modules "$HPC_CONFIG_DIR/modules.txt" \
|
||||
--requirements "$HPC_CONFIG_DIR/requirements.txt"
|
||||
set -euo pipefail
|
||||
cd "$PBS_O_WORKDIR"
|
||||
|
||||
echo '>>> Installing ipykernel...'
|
||||
CLUSTER_ID="${VSC_INSTITUTE_CLUSTER:-generic}"
|
||||
python -m ipykernel install --user --name="sel3_${CLUSTER_ID}" \
|
||||
--display-name "SEL3 (${CLUSTER_ID})"
|
||||
|
||||
echo '>>> Done'
|
||||
75
scripts/hpc/train.pbs
Normal file
75
scripts/hpc/train.pbs
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
# Production training (requires GPU at runtime):
|
||||
# qsub -l gpus=1 scripts/hpc/train.pbs
|
||||
# Debug/CPU training:
|
||||
# qsub scripts/hpc/train.pbs
|
||||
|
||||
#PBS -N brittlestar-ppo
|
||||
#PBS -l nodes=1:ppn=8
|
||||
#PBS -l walltime=24:00:00
|
||||
#PBS -o runs/brittlestar-ppo.o$PBS_JOBID
|
||||
#PBS -e runs/brittlestar-ppo.e$PBS_JOBID
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# Preliminary status echo
|
||||
echo ">>> Starting training job $PBS_JOBID on $(hostname)..."
|
||||
|
||||
if [ -n "$PBS_O_WORKDIR" ]; then
|
||||
cd "$PBS_O_WORKDIR"
|
||||
fi
|
||||
|
||||
# Set up storage paths dynamically
|
||||
PROJ_NAME=$(basename "$PWD")
|
||||
RUN_ID="brittlestar_${PBS_JOBID}"
|
||||
SCRATCH_RUNDIR="$VSC_SCRATCH/runs/$RUN_ID"
|
||||
DATA_RUNDIR="$VSC_DATA/runs/$RUN_ID"
|
||||
mkdir -p "$SCRATCH_RUNDIR" "$DATA_RUNDIR" runs/
|
||||
|
||||
# Keep caches off $VSC_HOME (quota ~3 GB).
|
||||
export PIP_CACHE_DIR="$VSC_SCRATCH/.cache/pip"
|
||||
export UV_CACHE_DIR="$VSC_SCRATCH/.cache/uv"
|
||||
mkdir -p "$PIP_CACHE_DIR" "$UV_CACHE_DIR"
|
||||
|
||||
module load vsc-venv
|
||||
|
||||
echo ">>> Synchronizing and activating environment (vsc-venv)..."
|
||||
HPC_CONFIG_DIR="$VSC_DATA/$PROJ_NAME/env/hpc"
|
||||
if [ ! -d "$HPC_CONFIG_DIR" ]; then
|
||||
echo "ERROR: HPC_CONFIG_DIR ($HPC_CONFIG_DIR) does not exist. Run install.sh first."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# cd to $VSC_DATA so vsc-venv finds its venvs/ directory there, not in $HOME.
|
||||
cd "$VSC_DATA/$PROJ_NAME"
|
||||
set +euo pipefail
|
||||
source vsc-venv --activate \
|
||||
--modules "$HPC_CONFIG_DIR/modules.txt" \
|
||||
--requirements "$HPC_CONFIG_DIR/requirements.txt"
|
||||
set -euo pipefail
|
||||
cd "$PBS_O_WORKDIR"
|
||||
|
||||
|
||||
echo ">>> Starting BrittleStar training..."
|
||||
export MUJOCO_GL=egl
|
||||
export WANDB_DIR="$SCRATCH_RUNDIR"
|
||||
|
||||
export PYTHONPATH="$PBS_O_WORKDIR/src:${PYTHONPATH:-}"
|
||||
|
||||
if [ -f "$VSC_DATA/$PROJ_NAME/.env" ]; then
|
||||
echo ">>> Sourcing API keys from .env..."
|
||||
export $(grep -v '^#' "$VSC_DATA/$PROJ_NAME/.env" | xargs)
|
||||
elif [ -f "$PBS_O_WORKDIR/.env" ]; then
|
||||
echo ">>> Sourcing API keys from .env..."
|
||||
export $(grep -v '^#' "$PBS_O_WORKDIR/.env" | xargs)
|
||||
fi
|
||||
|
||||
# Run training using Hydra overrides
|
||||
python scripts/train.py \
|
||||
hydra.run.dir="$SCRATCH_RUNDIR" \
|
||||
ppo=stable \
|
||||
logging=hpc
|
||||
|
||||
echo ">>> Staging out results to $DATA_RUNDIR..."
|
||||
cp -r "$SCRATCH_RUNDIR/." "$DATA_RUNDIR/"
|
||||
|
||||
echo ">>> Done"
|
||||
445
scripts/plots/analyze_comparisons.py
Normal file
445
scripts/plots/analyze_comparisons.py
Normal file
|
|
@ -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.")
|
||||
335
scripts/plots/analyze_convergence.py
Normal file
335
scripts/plots/analyze_convergence.py
Normal file
|
|
@ -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),
|
||||
)
|
||||
77
scripts/plots/plot_config.py
Normal file
77
scripts/plots/plot_config.py
Normal file
|
|
@ -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
|
||||
177
scripts/simulate.py
Normal file
177
scripts/simulate.py
Normal file
|
|
@ -0,0 +1,177 @@
|
|||
"""Simulate a trained policy in the MuJoCo viewer.
|
||||
|
||||
Automatically extracts the training configuration (morphology, environment, etc.)
|
||||
from the sidecar metadata YAML file to ensure simulation perfectly matches training.
|
||||
Override simulation settings via CLI, e.g.:
|
||||
uv run scripts/simulate.py \
|
||||
simulation.morphology_override=configs/morphology/3_arms.yaml \
|
||||
simulation.model_path=runs/.../final_model.flax
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
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.checkpoint import load_metadata, metadata_to_configs
|
||||
from brittle_star_project.evaluation.eval_env_builder import build_eval_env
|
||||
from brittle_star_project.evaluation.rollout import rollout_headless, rollout_viewer
|
||||
from brittle_star_project.evaluation.video import (
|
||||
record_episode,
|
||||
create_evaluation_dir,
|
||||
save_evaluation_metadata,
|
||||
)
|
||||
|
||||
|
||||
@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3")
|
||||
def main(dict_cfg: DictConfig) -> None:
|
||||
# 1. Hydra composes ONLY SimulationSettings
|
||||
cfg = OmegaConf.to_object(OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg))
|
||||
sim_cfg = cfg.simulation
|
||||
|
||||
model_path_str = sim_cfg.model_path
|
||||
if model_path_str is None:
|
||||
raise ValueError(
|
||||
"simulation.model_path must be set to a .flax checkpoint (e.g. final_model.flax)"
|
||||
)
|
||||
|
||||
model_path = Path(hydra.utils.to_absolute_path(model_path_str))
|
||||
if model_path.suffix != ".flax":
|
||||
raise ValueError(f"Expected a '.flax' checkpoint, got '{model_path.name}'.")
|
||||
|
||||
# 2. Discover + load sidecar metadata YAML
|
||||
metadata_override = None
|
||||
if sim_cfg.metadata_path is not None:
|
||||
metadata_override = Path(hydra.utils.to_absolute_path(sim_cfg.metadata_path))
|
||||
|
||||
metadata = load_metadata(model_path, metadata_override)
|
||||
|
||||
# 3. Reconstruct typed configs from metadata
|
||||
training = metadata_to_configs(metadata)
|
||||
|
||||
seed = int(cfg.experiment.seed)
|
||||
|
||||
# 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)
|
||||
|
||||
# 8. Run simulation
|
||||
headless = bool(sim_cfg.headless)
|
||||
max_steps = sim_cfg.max_steps
|
||||
|
||||
if sim_cfg.record_video:
|
||||
if max_steps is None:
|
||||
raise ValueError("simulation.max_steps is required when simulation.record_video=true")
|
||||
|
||||
max_steps_i = int(max_steps)
|
||||
if max_steps_i <= 0:
|
||||
raise ValueError("simulation.max_steps must be > 0")
|
||||
|
||||
if sim_cfg.video_output_path is None:
|
||||
eval_dir = create_evaluation_dir(model_path)
|
||||
output_path = eval_dir / "simulation.mp4"
|
||||
else:
|
||||
output_path = Path(hydra.utils.to_absolute_path(sim_cfg.video_output_path))
|
||||
eval_dir = output_path.parent
|
||||
eval_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
result = record_episode(
|
||||
env=env,
|
||||
policy=policy,
|
||||
seed=seed,
|
||||
max_steps=max_steps_i,
|
||||
action_low=action_low,
|
||||
action_high=action_high,
|
||||
action_mask=action_mask,
|
||||
output_path=output_path,
|
||||
camera_id=sim_cfg.camera_id,
|
||||
)
|
||||
|
||||
save_evaluation_metadata(
|
||||
eval_dir=eval_dir,
|
||||
morphology_override_path=sim_cfg.morphology_override,
|
||||
seed=seed,
|
||||
max_steps=max_steps_i,
|
||||
result=result,
|
||||
)
|
||||
final_dist_str = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}"
|
||||
print(f"Video saved to {output_path}")
|
||||
print(
|
||||
"episode done: "
|
||||
f"return={result.return_:.6f}, len={result.length}, "
|
||||
f"target_reached={result.reached_target}, final_xy_dist={final_dist_str}"
|
||||
)
|
||||
elif headless:
|
||||
if max_steps is None:
|
||||
raise ValueError("simulation.max_steps is required when simulation.headless=true")
|
||||
|
||||
max_steps_i = int(max_steps)
|
||||
if max_steps_i <= 0:
|
||||
raise ValueError("simulation.max_steps must be > 0")
|
||||
|
||||
result = rollout_headless(
|
||||
env=env,
|
||||
policy=policy,
|
||||
seed=seed,
|
||||
max_steps=max_steps_i,
|
||||
action_low=action_low,
|
||||
action_high=action_high,
|
||||
action_mask=action_mask,
|
||||
)
|
||||
final_dist_str = "n/a" if result.final_xy_dist is None else f"{result.final_xy_dist:.3f}"
|
||||
print(
|
||||
"episode done: "
|
||||
f"return={result.return_:.6f}, len={result.length}, "
|
||||
f"target_reached={result.reached_target}, final_xy_dist={final_dist_str}"
|
||||
)
|
||||
else:
|
||||
max_steps_val = None
|
||||
if max_steps is not None:
|
||||
max_steps_i = int(max_steps)
|
||||
if max_steps_i <= 0:
|
||||
raise ValueError("simulation.max_steps must be > 0")
|
||||
max_steps_val = max_steps_i
|
||||
|
||||
model_dt = float(state0.mj_model.opt.timestep)
|
||||
control_dt = model_dt * float(training.environment.num_physics_steps_per_control_step)
|
||||
|
||||
rollout_viewer(
|
||||
env=env,
|
||||
policy=policy,
|
||||
seed=seed,
|
||||
state=state0,
|
||||
control_dt=control_dt,
|
||||
max_steps=max_steps_val,
|
||||
action_low=action_low,
|
||||
action_high=action_high,
|
||||
action_mask=action_mask,
|
||||
)
|
||||
|
||||
env.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
register_configs()
|
||||
main()
|
||||
9
scripts/simulate.sh
Normal file
9
scripts/simulate.sh
Normal file
|
|
@ -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
|
||||
141
scripts/tools/dump_mjcf.py
Normal file
141
scripts/tools/dump_mjcf.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Dump MJCF XML for a brittle-star morphology using the project's Hydra configs.
|
||||
|
||||
Usage examples:
|
||||
|
||||
# Use a named morphology config from configs/morphology (Hydra style)
|
||||
uv run python scripts/analysis/dump_mjcf.py morphology=3_arms
|
||||
|
||||
# Use a morphology override YAML (same key as simulation.morphology_override)
|
||||
uv run python scripts/analysis/dump_mjcf.py \
|
||||
simulation.morphology_override=configs/morphology/3_arms.yaml
|
||||
|
||||
Output path:
|
||||
Provide `dump_out=path/to/file.xml` on the command line, otherwise writes `morphology.xml` in
|
||||
current directory or `runs/morphologies/<name>.xml`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
import hydra
|
||||
import yaml
|
||||
from omegaconf import DictConfig, OmegaConf
|
||||
|
||||
from brittle_star_project.configs.register_configs import register_configs
|
||||
from brittle_star_project.environment.env_config import MorphologyConfig
|
||||
from brittle_star_project.environment.factory import BrittleStarEnvFactory
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def extract_xml_string(obj: Any) -> Optional[str]:
|
||||
"""
|
||||
Attempts to serialize the morphology object to an XML string by checking
|
||||
common dm_control and internal API methods.
|
||||
"""
|
||||
serialization_methods = [
|
||||
"to_xml_string",
|
||||
"to_xml",
|
||||
"to_string",
|
||||
"to_mjcf",
|
||||
"to_mjcf_string",
|
||||
"get_mjcf",
|
||||
"get_mjcf_str",
|
||||
"export_to_xml_string",
|
||||
]
|
||||
|
||||
# If the object itself has an 'mjcf' attribute, try to serialize that instead
|
||||
target_obj = getattr(obj, "mjcf", obj)
|
||||
|
||||
for method_name in serialization_methods:
|
||||
method = getattr(target_obj, method_name, None)
|
||||
if callable(method):
|
||||
try:
|
||||
xml_data = method()
|
||||
# Safely handle both string and byte responses
|
||||
if isinstance(xml_data, str):
|
||||
return xml_data
|
||||
elif isinstance(xml_data, bytes):
|
||||
return xml_data.decode("utf-8")
|
||||
except Exception as e:
|
||||
logger.debug(f"Method {method_name}() failed during serialization: {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def resolve_output_path(cfg: DictConfig) -> Path:
|
||||
"""Determines the appropriate output path for the MJCF XML."""
|
||||
dump_out = cfg.get("dump_out", None)
|
||||
if dump_out is not None:
|
||||
return Path(hydra.utils.to_absolute_path(str(dump_out)))
|
||||
|
||||
morph_name = "morphology"
|
||||
for arg in sys.argv[1:]:
|
||||
if arg.startswith("morphology="):
|
||||
morph_name = arg.split("=", 1)[1]
|
||||
break
|
||||
|
||||
default_out = (
|
||||
f"runs/morphologies/{morph_name}.xml" if morph_name != "morphology" else "morphology.xml"
|
||||
)
|
||||
return Path(hydra.utils.to_absolute_path(default_out))
|
||||
|
||||
|
||||
@hydra.main(config_path="../../configs", config_name="main_config", version_base="1.3")
|
||||
def main(cfg: DictConfig) -> None:
|
||||
"""Main entry point to construct the morphology and dump its XML."""
|
||||
logger.info("Initializing morphology construction...")
|
||||
|
||||
# Extract morphology config safely using dict `.get()` to avoid OmegaConf AttributeErrors
|
||||
simulation_cfg = cfg.get("simulation", cfg)
|
||||
override_path = simulation_cfg.get("morphology_override", None)
|
||||
|
||||
if override_path:
|
||||
logger.info(f"Using morphology override: {override_path}")
|
||||
with open(hydra.utils.to_absolute_path(override_path), "r") as f:
|
||||
data = yaml.safe_load(f) or {}
|
||||
morph_cfg = MorphologyConfig(**data)
|
||||
else:
|
||||
# Fallback to default simulation morphology, or an empty base config
|
||||
morph_node = simulation_cfg.get("morphology", cfg.get("morphology", None))
|
||||
|
||||
if morph_node is not None:
|
||||
# Convert OmegaConf node to dict and instantiate MorphologyConfig.
|
||||
# This ensures any missing keys gracefully fall back to the dataclass defaults.
|
||||
morph_dict = OmegaConf.to_container(morph_node, resolve=True)
|
||||
if isinstance(morph_dict, dict):
|
||||
# Filter to avoid unexpected kwargs if the dataclass is strictly defined
|
||||
if dataclasses.is_dataclass(MorphologyConfig):
|
||||
valid_keys = {f.name for f in dataclasses.fields(MorphologyConfig)}
|
||||
morph_dict = {k: v for k, v in morph_dict.items() if k in valid_keys}
|
||||
morph_cfg = MorphologyConfig(**morph_dict)
|
||||
else:
|
||||
morph_cfg = MorphologyConfig()
|
||||
else:
|
||||
morph_cfg = MorphologyConfig()
|
||||
|
||||
morphology = BrittleStarEnvFactory.create_morphology(morph_cfg)
|
||||
|
||||
xml_text = extract_xml_string(morphology)
|
||||
if not xml_text:
|
||||
raise RuntimeError("Failed to serialize morphology to MJCF/XML. ")
|
||||
|
||||
out_path = resolve_output_path(cfg)
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with out_path.open("w", encoding="utf-8") as f:
|
||||
f.write(xml_text)
|
||||
|
||||
logger.info(f"Successfully exported MJCF XML to: {out_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
register_configs()
|
||||
main()
|
||||
138
scripts/tools/extract_observation_bounds.py
Normal file
138
scripts/tools/extract_observation_bounds.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Empirically extract observation bounds (focused on joint velocities).
|
||||
|
||||
This script creates a MuJoCo environment using the project's factory and
|
||||
randomly samples actions to discover observed maxima for selected
|
||||
observation keys (joint_velocity, joint_position, joint_actuator_force).
|
||||
|
||||
Usage:
|
||||
python scripts/extract_observation_bounds.py \
|
||||
--morphology configs/morphology/3_arms.yaml --num-steps 5000 --seed 42
|
||||
|
||||
If `--morphology` is omitted the default `MorphologyConfig()` is used.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
import yaml
|
||||
import numpy as np
|
||||
|
||||
from brittle_star_project import BrittleStarEnvFactory, BrittleStarEnv, Backend
|
||||
from brittle_star_project.environment.env_config import (
|
||||
MorphologyConfig,
|
||||
ArenaConfig,
|
||||
EnvConfig,
|
||||
)
|
||||
|
||||
|
||||
def load_morphology(path: str | None) -> MorphologyConfig:
|
||||
if path is None:
|
||||
return MorphologyConfig()
|
||||
p = Path(path)
|
||||
if not p.exists():
|
||||
raise FileNotFoundError(f"Morphology file not found: {p}")
|
||||
with open(p, "r") as f:
|
||||
data = yaml.safe_load(f) or {}
|
||||
return MorphologyConfig(**data)
|
||||
|
||||
|
||||
def _extract_observations(state):
|
||||
# Under different backends the returned state may be a dict or an object
|
||||
obs = getattr(state, "observations", None)
|
||||
if obs is None and isinstance(state, dict):
|
||||
obs = state.get("observations", state)
|
||||
return obs
|
||||
|
||||
|
||||
def find_empirical_bounds(
|
||||
morph_cfg: MorphologyConfig,
|
||||
arena_cfg: ArenaConfig,
|
||||
env_cfg: EnvConfig,
|
||||
num_steps: int = 5000,
|
||||
seed: int = 42,
|
||||
) -> None:
|
||||
factory = BrittleStarEnvFactory()
|
||||
raw_env = factory.create_environment(Backend.MJC, morph_cfg, arena_cfg, env_cfg)
|
||||
env = BrittleStarEnv(raw_env, backend=Backend.MJC, config=env_cfg, morphology_config=morph_cfg)
|
||||
|
||||
# Initial reset
|
||||
state = env.reset(seed=seed)
|
||||
|
||||
# Determine action bounds
|
||||
action_space = getattr(raw_env, "action_space", None)
|
||||
if action_space is None:
|
||||
raise RuntimeError("Environment missing `action_space`; cannot sample actions.")
|
||||
|
||||
action_low = np.asarray(action_space.low, dtype=np.float32)
|
||||
action_high = np.asarray(action_space.high, dtype=np.float32)
|
||||
action_shape = action_low.shape
|
||||
|
||||
# Track maximum absolute observed values
|
||||
tracked_keys = ["joint_velocity", "joint_position", "joint_actuator_force"]
|
||||
max_observed = {k: 0.0 for k in tracked_keys}
|
||||
|
||||
# Include observation at reset
|
||||
obs0 = _extract_observations(state)
|
||||
if isinstance(obs0, dict):
|
||||
for k in tracked_keys:
|
||||
if k in obs0:
|
||||
max_observed[k] = max(max_observed[k], float(np.max(np.abs(np.asarray(obs0[k])))))
|
||||
|
||||
rng = np.random.RandomState(seed)
|
||||
for i in range(num_steps):
|
||||
u = rng.uniform(size=action_shape)
|
||||
action = action_low + (action_high - action_low) * u
|
||||
|
||||
# Provide a numpy RNG to the env step; wrapper will pass it if accepted.
|
||||
step_out = env.step(state=state, action=action, rng=env.make_rng(seed + i + 1))
|
||||
|
||||
# Unpack next state from common return conventions
|
||||
if hasattr(step_out, "state"):
|
||||
next_state = step_out.state
|
||||
elif isinstance(step_out, (tuple, list)) and len(step_out) >= 1:
|
||||
next_state = step_out[0]
|
||||
else:
|
||||
next_state = step_out
|
||||
|
||||
obs = _extract_observations(next_state)
|
||||
if isinstance(obs, dict):
|
||||
for k in tracked_keys:
|
||||
if k in obs:
|
||||
val = float(np.max(np.abs(np.asarray(obs[k]))))
|
||||
if val > max_observed[k]:
|
||||
max_observed[k] = val
|
||||
|
||||
state = next_state
|
||||
|
||||
# Print recommended bounds with a 20% safety margin
|
||||
print("\n--- Recommended Observation Bounds (20% margin) ---")
|
||||
for k, v in max_observed.items():
|
||||
if v == 0.0:
|
||||
print(f"{k}: observed max 0.0 (increase sampling or inspect env)")
|
||||
else:
|
||||
safe = v * 1.2
|
||||
print(f"{k}: [-{safe:.6f}, {safe:.6f}] (observed max: {v:.6f})")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--morphology", type=str, default=None, help="Path to morphology YAML (optional)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-steps", type=int, default=5000, help="Number of random steps to sample"
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=42, help="RNG seed")
|
||||
args = parser.parse_args()
|
||||
|
||||
morph_cfg = load_morphology(args.morphology)
|
||||
arena_cfg = ArenaConfig()
|
||||
env_cfg = EnvConfig()
|
||||
|
||||
find_empirical_bounds(morph_cfg, arena_cfg, env_cfg, num_steps=args.num_steps, seed=args.seed)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
58
scripts/train.py
Normal file
58
scripts/train.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
import os
|
||||
import torch
|
||||
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.trainers.PPOTrainer import PPOTrainer
|
||||
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
||||
from experiment_logger import init_logger, get_logger
|
||||
|
||||
|
||||
def make_env(cfg: BrittleStarConfig) -> BrittleStarJaxEnvWrapper:
|
||||
"""Create the environment using the structured configuration."""
|
||||
return BrittleStarJaxEnvWrapper(
|
||||
morphology=cfg.morphology,
|
||||
arena=cfg.arena,
|
||||
env_config=cfg.environment,
|
||||
num_envs=cfg.ppo.num_envs,
|
||||
)
|
||||
|
||||
|
||||
@hydra.main(config_path="../configs", config_name="main_config", version_base="1.3")
|
||||
def main(dict_cfg: DictConfig):
|
||||
# 1. Convert DictConfig to structured dataclass, ensuring the root schema is applied correctly.
|
||||
config: BrittleStarConfig = OmegaConf.to_object(
|
||||
OmegaConf.merge(OmegaConf.structured(BrittleStarConfig), dict_cfg)
|
||||
)
|
||||
|
||||
# 2. Setup run metadata
|
||||
# Hydra changes CWD to the output directory by default.
|
||||
run_dir = os.getcwd()
|
||||
run_name = os.path.basename(run_dir)
|
||||
|
||||
# 3. Initialize Logger
|
||||
cfg_dict = OmegaConf.to_container(dict_cfg, resolve=True, throw_on_missing=True)
|
||||
init_logger(
|
||||
run_name=run_name,
|
||||
full_config=cfg_dict,
|
||||
logging_cfg=config.logging,
|
||||
base_dir=os.path.dirname(run_dir),
|
||||
)
|
||||
logger = get_logger()
|
||||
logger.info(f"Hydra-initialized run: {run_name}")
|
||||
logger.info(f"Output directory: {run_dir}")
|
||||
|
||||
# 4. Setup Environment and Torch
|
||||
env = make_env(config)
|
||||
torch.backends.cudnn.deterministic = config.experiment.torch_deterministic
|
||||
|
||||
# 5. Train - pass structured config directly
|
||||
ppo_trainer = PPOTrainer(config, env, run_dir, run_name)
|
||||
ppo_trainer.train()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
register_configs()
|
||||
main()
|
||||
File diff suppressed because one or more lines are too long
BIN
sitemap.xml.gz
BIN
sitemap.xml.gz
Binary file not shown.
19
src/brittle_star_project/MLPs/__init__.py
Normal file
19
src/brittle_star_project/MLPs/__init__.py
Normal file
|
|
@ -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",
|
||||
]
|
||||
67
src/brittle_star_project/MLPs/adjancency_builder.py
Normal file
67
src/brittle_star_project/MLPs/adjancency_builder.py
Normal file
|
|
@ -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
|
||||
93
src/brittle_star_project/MLPs/mlps.py
Normal file
93
src/brittle_star_project/MLPs/mlps.py
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
from dataclasses import dataclass, fields, field
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
|
||||
class OneDenseLayerMLP(nn.Module):
|
||||
@nn.compact
|
||||
def __call__(self, x):
|
||||
return nn.Dense(1, kernel_init=orthogonal(1), bias_init=constant(0.0))(x)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
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: 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.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 = 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)
|
||||
return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs})
|
||||
22
src/brittle_star_project/MLPs/routing.py
Normal file
22
src/brittle_star_project/MLPs/routing.py
Normal file
|
|
@ -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)
|
||||
28
src/brittle_star_project/__init__.py
Normal file
28
src/brittle_star_project/__init__.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
from .environment.env_types import Backend, Task
|
||||
from .environment.env_config import ArenaConfig, EnvConfig, MorphologyConfig
|
||||
from .environment.factory import BrittleStarEnvFactory
|
||||
from .environment.env_wrapper import BrittleStarEnv
|
||||
from .evaluation import (
|
||||
PolicyAgent,
|
||||
ControlPolicy,
|
||||
load_metadata,
|
||||
rollout_headless,
|
||||
rollout_viewer,
|
||||
EpisodeResult,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ArenaConfig",
|
||||
"Backend",
|
||||
"BrittleStarEnv",
|
||||
"BrittleStarEnvFactory",
|
||||
"EnvConfig",
|
||||
"MorphologyConfig",
|
||||
"Task",
|
||||
"PolicyAgent",
|
||||
"ControlPolicy",
|
||||
"load_metadata",
|
||||
"rollout_headless",
|
||||
"rollout_viewer",
|
||||
"EpisodeResult",
|
||||
]
|
||||
66
src/brittle_star_project/configs/config_architecture.py
Normal file
66
src/brittle_star_project/configs/config_architecture.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class LayerConfig:
|
||||
hidden_dims: List[int] = field(default_factory=lambda: [64, 64])
|
||||
activation: str = "tanh"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ArchitectureConfig:
|
||||
"""Base class for actor-critic network configurations.
|
||||
|
||||
Both centralized and decentralized architectures share a centralized critic
|
||||
composed of a feature extractor followed by a shallow output layer.
|
||||
|
||||
See docs/design/actor-critic.md for the full design rationale.
|
||||
"""
|
||||
|
||||
name: str = "base"
|
||||
|
||||
# Actor pipeline
|
||||
sensor: Optional[LayerConfig] = None
|
||||
propagator: Optional[LayerConfig] = None
|
||||
motor: Optional[LayerConfig] = None
|
||||
|
||||
# Critic pipeline
|
||||
feature_extractor: Optional[LayerConfig] = None
|
||||
critic: Optional[LayerConfig] = None
|
||||
|
||||
# Decentralized
|
||||
message_passing_steps: Optional[int] = None
|
||||
topology_type: Optional[str] = None # Supported values: "ring", "fully_connected"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CentralizedConfig(ArchitectureConfig):
|
||||
"""Centralized actor-critic architecture (baseline).
|
||||
|
||||
The actor is a single global policy composed of a sensor (input network)
|
||||
and a motor (output network). The sensor receives the full concatenated
|
||||
global observation; the motor projects the hidden state to all joint actions.
|
||||
|
||||
See docs/design/actor-critic.md for the full design rationale.
|
||||
"""
|
||||
|
||||
name: str = "centralized"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DecentralizedConfig(ArchitectureConfig):
|
||||
"""Decentralized actor architecture (NerveNet-MLP variant).
|
||||
|
||||
Each node runs a local sensor, exchanges messages with neighbours via a
|
||||
propagator for a fixed number of steps, and then a local motor produces
|
||||
the joint offset for that node only.
|
||||
|
||||
The critic remains centralized (shared with the base class): it receives the
|
||||
full concatenated global observation and outputs a single scalar.
|
||||
|
||||
See docs/design/actor-critic.md and docs/design/communication.md for the
|
||||
full design rationale.
|
||||
"""
|
||||
|
||||
name: str = "decentralized"
|
||||
38
src/brittle_star_project/configs/config_evaluation.py
Normal file
38
src/brittle_star_project/configs/config_evaluation.py
Normal file
|
|
@ -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."
|
||||
)
|
||||
11
src/brittle_star_project/configs/config_experiment.py
Normal file
11
src/brittle_star_project/configs/config_experiment.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ExperimentConfig:
|
||||
exp_name: str = "brittle_star_ppo"
|
||||
seed: int = 1
|
||||
torch_deterministic: bool = True
|
||||
cuda: bool = True
|
||||
debug_sanity: bool = False
|
||||
base_run_dir: str = "runs"
|
||||
22
src/brittle_star_project/configs/config_ppo.py
Normal file
22
src/brittle_star_project/configs/config_ppo.py
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class PPOConfig:
|
||||
learning_rate: float = 2.5e-4
|
||||
total_timesteps: int = 10000000
|
||||
num_envs: int = 100
|
||||
num_steps: int = 128
|
||||
anneal_lr: bool = True
|
||||
gamma: float = 0.99
|
||||
gae_lambda: float = 0.95
|
||||
num_minibatches: int = 4
|
||||
update_epochs: int = 4
|
||||
norm_adv: bool = True
|
||||
clip_coef: float = 0.1
|
||||
clip_vloss: bool = True
|
||||
ent_coef: float = 0.01
|
||||
vf_coef: float = 0.5
|
||||
max_grad_norm: float = 0.5
|
||||
target_kl: Optional[float] = None
|
||||
32
src/brittle_star_project/configs/config_simulation.py
Normal file
32
src/brittle_star_project/configs/config_simulation.py
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class SimulationSettings:
|
||||
"""Settings for the simulation script."""
|
||||
|
||||
model_path: Optional[str] = None
|
||||
|
||||
# Script behavior
|
||||
headless: bool = False
|
||||
# If None, viewer mode runs until window closed or target reached.
|
||||
max_steps: Optional[int] = None
|
||||
|
||||
# Override morphology for amputation experiments.
|
||||
# When set, the environment uses this morphology instead of the trained one.
|
||||
# Points to a morphology config YAML file (e.g. configs/morphology/3_arms.yaml).
|
||||
# Observations are padded from the override morphology UP TO the training
|
||||
# morphology's shape via compute_padding_masks(override, reference=training).
|
||||
morphology_override: Optional[str] = None
|
||||
|
||||
# Video recording (requires [evaluation] extra)
|
||||
record_video: bool = False
|
||||
# When None, video is saved in a per-model evaluation folder alongside the model.
|
||||
video_output_path: Optional[str] = None
|
||||
# Camera ID to use for video recording (1 is usually the close-up camera)
|
||||
camera_id: int = 1
|
||||
|
||||
# Optional override for the sidecar metadata YAML file.
|
||||
# If None, it defaults to the model_path with a `_metadata.yaml` suffix.
|
||||
metadata_path: Optional[str] = None
|
||||
36
src/brittle_star_project/configs/main_config.py
Normal file
36
src/brittle_star_project/configs/main_config.py
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
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
|
||||
from brittle_star_project.environment.env_config import (
|
||||
MorphologyConfig,
|
||||
ArenaConfig,
|
||||
EnvConfig,
|
||||
ObservationBoundsConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class BrittleStarConfig:
|
||||
"""Root configuration for a brittle star training run.
|
||||
|
||||
Composed of strictly separated sub-configs. Each sub-config can be swapped
|
||||
independently via CLI or a different YAML file. See configs/README.md.
|
||||
"""
|
||||
|
||||
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.
|
||||
architecture: ArchitectureConfig = field(default_factory=ArchitectureConfig)
|
||||
morphology: MorphologyConfig = field(default_factory=MorphologyConfig)
|
||||
arena: ArenaConfig = field(default_factory=ArenaConfig)
|
||||
environment: EnvConfig = field(default_factory=EnvConfig)
|
||||
obs_bounds: ObservationBoundsConfig = field(default_factory=ObservationBoundsConfig)
|
||||
simulation: SimulationSettings = field(default_factory=SimulationSettings)
|
||||
48
src/brittle_star_project/configs/register_configs.py
Normal file
48
src/brittle_star_project/configs/register_configs.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
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,
|
||||
DecentralizedConfig,
|
||||
)
|
||||
from brittle_star_project.configs.config_simulation import SimulationSettings
|
||||
from brittle_star_project.environment.env_config import (
|
||||
MorphologyConfig,
|
||||
ArenaConfig,
|
||||
EnvConfig,
|
||||
ObservationBoundsConfig,
|
||||
)
|
||||
from brittle_star_project.configs.main_config import BrittleStarConfig
|
||||
|
||||
|
||||
def register_configs() -> None:
|
||||
"""Register all dataclasses with Hydra's ConfigStore.
|
||||
|
||||
This must be called before hydra.main() processes the config, ensuring
|
||||
every structured config is validated against its Python schema. Typos in
|
||||
YAML keys will raise ConfigAttributeError at startup.
|
||||
"""
|
||||
cs = ConfigStore.instance()
|
||||
|
||||
# Root schema
|
||||
cs.store(name="brittle_star_config", node=BrittleStarConfig)
|
||||
|
||||
# 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
|
||||
cs.store(group="architecture", name="centralized_schema", node=CentralizedConfig)
|
||||
cs.store(group="architecture", name="decentralized_schema", node=DecentralizedConfig)
|
||||
|
||||
# Environment configs
|
||||
cs.store(group="morphology", name="base_morphology", node=MorphologyConfig)
|
||||
cs.store(group="arena", name="base_arena", node=ArenaConfig)
|
||||
cs.store(group="environment", name="base_environment", node=EnvConfig)
|
||||
cs.store(group="obs_bounds", name="base_obs_bounds", node=ObservationBoundsConfig)
|
||||
cs.store(group="simulation", name="base_simulation", node=SimulationSettings)
|
||||
10
src/brittle_star_project/dataclasses/EpisodeStatistics.py
Normal file
10
src/brittle_star_project/dataclasses/EpisodeStatistics.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
import flax.struct
|
||||
import jax.numpy as jnp
|
||||
|
||||
|
||||
@flax.struct.dataclass
|
||||
class EpisodeStatistics:
|
||||
episode_returns: jnp.ndarray
|
||||
episode_lengths: jnp.ndarray
|
||||
returned_episode_returns: jnp.ndarray
|
||||
returned_episode_lengths: jnp.ndarray
|
||||
6
src/brittle_star_project/dataclasses/__init__.py
Normal file
6
src/brittle_star_project/dataclasses/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from .EpisodeStatistics import EpisodeStatistics
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EpisodeStatistics",
|
||||
]
|
||||
104
src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py
Normal file
104
src/brittle_star_project/environment/BrittleStarJaxEnvWrapper.py
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
import jax
|
||||
import jax.numpy as jnp
|
||||
|
||||
from experiment_logger import get_logger
|
||||
from .env_config import EnvConfig, MorphologyConfig, ArenaConfig
|
||||
from .env_types import Backend
|
||||
from .factory import BrittleStarEnvFactory
|
||||
from .padded_obs_wrapper import compute_padding_masks
|
||||
|
||||
|
||||
class BrittleStarJaxEnvWrapper:
|
||||
def __init__(
|
||||
self,
|
||||
morphology: MorphologyConfig,
|
||||
arena: ArenaConfig,
|
||||
env_config: EnvConfig,
|
||||
num_envs: int,
|
||||
backend: Backend = Backend.MJX,
|
||||
):
|
||||
self._morphology = morphology
|
||||
self._arena = arena
|
||||
self._env_config = env_config
|
||||
self._backend = backend
|
||||
self._num_envs = num_envs
|
||||
self._env = BrittleStarEnvFactory.create_environment(
|
||||
self._backend, self._morphology, self._arena, self._env_config
|
||||
)
|
||||
|
||||
# Pre-compute masks for observation padding
|
||||
self._padding_masks = compute_padding_masks(self._morphology.segments_per_arm)
|
||||
|
||||
self._vectorized_reset = jax.jit(jax.vmap(self._env.reset))
|
||||
self._vectorized_step = jax.jit(jax.vmap(self._env.step))
|
||||
self._vectorized_action_sample = jax.jit(jax.vmap(self._env.action_space.sample))
|
||||
|
||||
self._action_rng = None
|
||||
|
||||
self.logger = get_logger()
|
||||
self.logger.info(
|
||||
f"Initialized BrittleStarJaxEnvWrapper with {num_envs} envs on {backend.value}"
|
||||
)
|
||||
|
||||
@property
|
||||
def backend(self):
|
||||
return self._backend
|
||||
|
||||
@property
|
||||
def raw(self):
|
||||
return self._env
|
||||
|
||||
@property
|
||||
def padding_masks(self) -> dict:
|
||||
"""Pre-computed boolean masks for amputated limb padding.
|
||||
|
||||
Pass to create_obs_processor so the processor handles padding
|
||||
after normalization in the correct pipeline order.
|
||||
"""
|
||||
return self._padding_masks
|
||||
|
||||
@property
|
||||
def single_action_space(self):
|
||||
return self._env.action_space
|
||||
|
||||
@property
|
||||
def single_observation_space(self):
|
||||
return self._env.observation_space
|
||||
|
||||
def reset(self, seed: int = 0):
|
||||
self.logger.info(f"Resetting vectorized environment environments with seed {seed}")
|
||||
self._action_rng, env_rng = jax.random.split(jax.random.PRNGKey(seed), 2)
|
||||
env_rngs = jnp.array(jax.random.split(env_rng, self._num_envs))
|
||||
state = self._vectorized_reset(rng=env_rngs)
|
||||
return state
|
||||
|
||||
def sample_actions(self):
|
||||
assert self._action_rng is not None, "Call reset() before sample_actions()"
|
||||
self._action_rng, *sub_rngs = jnp.array(
|
||||
jax.random.split(self._action_rng, self._num_envs + 1)
|
||||
)
|
||||
return self._vectorized_action_sample(rng=jnp.array(sub_rngs))
|
||||
|
||||
def step(self, state, action):
|
||||
return self._vectorized_step(state=state, action=action)
|
||||
|
||||
def close(self):
|
||||
self._env.close()
|
||||
|
||||
@staticmethod
|
||||
def default(num_envs: int, backend: Backend = Backend.MJX) -> "BrittleStarJaxEnvWrapper":
|
||||
morphology = MorphologyConfig()
|
||||
arena = ArenaConfig()
|
||||
env_config = EnvConfig()
|
||||
return BrittleStarJaxEnvWrapper(
|
||||
morphology, arena, env_config, num_envs=num_envs, backend=backend
|
||||
)
|
||||
|
||||
def __str__(self):
|
||||
morphology_str = str(self._morphology)
|
||||
arena_str = str(self._arena)
|
||||
env_config_str = str(self._env_config)
|
||||
return (
|
||||
f"BrittleStarJaxEnvWrapper(backend={self._backend}, num_envs={self._num_envs}, "
|
||||
+ f"morphology={morphology_str}, arena={arena_str}, env_config={env_config_str})"
|
||||
)
|
||||
19
src/brittle_star_project/environment/__init__.py
Normal file
19
src/brittle_star_project/environment/__init__.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig, MorphMode
|
||||
from .env_types import Backend, Task
|
||||
from .env_wrapper import BrittleStarEnv
|
||||
from .factory import BrittleStarEnvFactory
|
||||
from .obs_processing import create_obs_processor
|
||||
from .padded_obs_wrapper import compute_padding_masks
|
||||
|
||||
__all__ = [
|
||||
"ArenaConfig",
|
||||
"EnvConfig",
|
||||
"MorphologyConfig",
|
||||
"Backend",
|
||||
"Task",
|
||||
"BrittleStarEnv",
|
||||
"BrittleStarEnvFactory",
|
||||
"MorphMode",
|
||||
"create_obs_processor",
|
||||
"compute_padding_masks",
|
||||
]
|
||||
99
src/brittle_star_project/environment/env_config.py
Normal file
99
src/brittle_star_project/environment/env_config.py
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
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.
|
||||
|
||||
segments_per_arm defines the number of segments for each arm. The length of
|
||||
this list implicitly sets the number of arms. Use 0 segments to represent
|
||||
a fully amputated arm (e.g., [4, 0, 4, 2, 4] for a 5-arm morphology with
|
||||
arm 1 removed and arm 3 shortened).
|
||||
|
||||
The upstream biorobot library natively supports per-arm segment counts.
|
||||
"""
|
||||
|
||||
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:
|
||||
return len(self.segments_per_arm)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ArenaConfig:
|
||||
size: list[float] = field(default_factory=lambda: [10.0, 5.0])
|
||||
sand_ground_color: bool = True
|
||||
attach_target: bool = True
|
||||
wall_height: float = 1.5
|
||||
wall_thickness: float = 0.1
|
||||
|
||||
|
||||
@dataclass
|
||||
class EnvConfig:
|
||||
"""Shared environment settings.
|
||||
|
||||
Note: Some tasks have additional parameters (see fields below).
|
||||
"""
|
||||
|
||||
task: Task = Task.DIRECTED_LOCOMOTION
|
||||
|
||||
simulation_time: float = 10000.0
|
||||
num_physics_steps_per_control_step: int = 10
|
||||
time_scale: int = 2
|
||||
|
||||
camera_ids: list[int] = field(default_factory=lambda: [0, 1])
|
||||
# (height, width)
|
||||
render_size: list[int] = field(default_factory=lambda: [480, 640])
|
||||
|
||||
joint_randomization_noise_scale: float = 0.0
|
||||
|
||||
# Directed locomotion
|
||||
target_distance: float = 3.0
|
||||
|
||||
# Light escape
|
||||
# Per docs in upstream env config: integer factors of 200.
|
||||
light_perlin_noise_scale: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class ObservationBoundsConfig:
|
||||
"""Physical observation bounds for deterministic min-max normalization."""
|
||||
|
||||
# Empirical testing based on the extract_observation_bounds.py script run for 1.000.000 steps
|
||||
|
||||
# Based on max. ctrlrange (0.78539816339744828) in XML, but empirical testing went slightly over
|
||||
joint_position: list[float] = field(default_factory=lambda: [-0.8, 0.8])
|
||||
# Empirical testing showed max. 3.22, adding buffer to be safe. Consider higher values "fast".
|
||||
joint_velocity: list[float] = field(default_factory=lambda: [-5.0, 5.0])
|
||||
# Based on max. forceRange in XML, verified with empirical testing
|
||||
joint_actuator_force: list[float] = field(default_factory=lambda: [-3.75, 3.75])
|
||||
# Based on intuition and reasoning
|
||||
segment_contact: list[float] = field(default_factory=lambda: [0.0, 1.0])
|
||||
robot_direction_to_target: list[float] = field(default_factory=lambda: [-1.0, 1.0])
|
||||
disk_z_tilt: list[float] = field(default_factory=lambda: [0.0, 3.141592653589793])
|
||||
|
||||
def to_bounds_dict(self) -> dict[str, tuple[float, float]]:
|
||||
return {
|
||||
"disk_z_tilt": tuple(self.disk_z_tilt),
|
||||
"joint_actuator_force": tuple(self.joint_actuator_force),
|
||||
"joint_position": tuple(self.joint_position),
|
||||
"joint_velocity": tuple(self.joint_velocity),
|
||||
"robot_direction_to_target": tuple(self.robot_direction_to_target),
|
||||
"segment_contact": tuple(self.segment_contact),
|
||||
}
|
||||
21
src/brittle_star_project/environment/env_types.py
Normal file
21
src/brittle_star_project/environment/env_types.py
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class Backend(str, Enum):
|
||||
"""Physics backend.
|
||||
|
||||
- MJC: MuJoCo C engine
|
||||
- MJX: MuJoCo XLA (JAX) engine
|
||||
"""
|
||||
|
||||
MJC = "MJC"
|
||||
MJX = "MJX"
|
||||
|
||||
|
||||
class Task(str, Enum):
|
||||
"""Which brittle-star task/environment to instantiate."""
|
||||
|
||||
DIRECTED_LOCOMOTION = "directed_locomotion"
|
||||
LIGHT_ESCAPE = "light_escape"
|
||||
99
src/brittle_star_project/environment/env_wrapper.py
Normal file
99
src/brittle_star_project/environment/env_wrapper.py
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .env_config import EnvConfig, MorphologyConfig
|
||||
from .env_types import Backend
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class StepResult:
|
||||
state: Any
|
||||
reward: float | None = None
|
||||
terminated: bool | None = None
|
||||
truncated: bool | None = None
|
||||
info: dict[str, Any] | None = None
|
||||
|
||||
|
||||
class BrittleStarEnv:
|
||||
"""Thin wrapper around the underlying DualMuJoCoEnvironment.
|
||||
|
||||
Goal: hide backend-specific RNG setup and provide a stable place to plug in RL.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
env: Any,
|
||||
*,
|
||||
backend: Backend,
|
||||
config: EnvConfig,
|
||||
morphology_config: MorphologyConfig | None = None,
|
||||
) -> None:
|
||||
self._env = env
|
||||
self._backend = backend
|
||||
self._config = config
|
||||
self._morphology_config = morphology_config
|
||||
|
||||
@property
|
||||
def raw(self) -> Any:
|
||||
return self._env
|
||||
|
||||
@property
|
||||
def backend(self) -> Backend:
|
||||
return self._backend
|
||||
|
||||
@property
|
||||
def config(self) -> EnvConfig:
|
||||
return self._config
|
||||
|
||||
@property
|
||||
def morphology_config(self) -> MorphologyConfig | None:
|
||||
return self._morphology_config
|
||||
|
||||
def make_rng(self, seed: int):
|
||||
if self._backend == Backend.MJC:
|
||||
return np.random.RandomState(seed)
|
||||
|
||||
import jax
|
||||
|
||||
return jax.random.PRNGKey(seed)
|
||||
|
||||
def reset(self, *, seed: int = 0):
|
||||
rng = self.make_rng(seed)
|
||||
state = self._env.reset(rng=rng)
|
||||
return state
|
||||
|
||||
def render(self, *, state: Any):
|
||||
return self._env.render(state=state)
|
||||
|
||||
def close(self) -> None:
|
||||
self._env.close()
|
||||
|
||||
def step(self, *, state: Any, action: Any, rng: Any | None = None) -> StepResult:
|
||||
"""Best-effort step wrapper.
|
||||
|
||||
Different env libraries return different tuples; we normalize common cases.
|
||||
"""
|
||||
|
||||
if not hasattr(self._env, "step"):
|
||||
raise AttributeError("Underlying env has no step() method")
|
||||
|
||||
step_fn = self._env.step
|
||||
sig = inspect.signature(step_fn)
|
||||
params = list(sig.parameters)
|
||||
|
||||
# Common patterns:
|
||||
# - step(state, action)
|
||||
# - step(state, action, rng)
|
||||
# - step(state, action, key)
|
||||
# We pass rng only if the callable accepts a 3rd arg.
|
||||
if len(params) >= 3 and rng is not None:
|
||||
out = step_fn(state, action, rng)
|
||||
else:
|
||||
out = step_fn(state, action)
|
||||
|
||||
return out
|
||||
112
src/brittle_star_project/environment/factory.py
Normal file
112
src/brittle_star_project/environment/factory.py
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import asdict
|
||||
|
||||
from moojoco.environment.dual import DualMuJoCoEnvironment
|
||||
|
||||
from .env_config import ArenaConfig, EnvConfig, MorphologyConfig
|
||||
from .env_types import Backend, Task
|
||||
|
||||
|
||||
class BrittleStarEnvFactory:
|
||||
"""Creates brittle-star morphology, arena, and task environment instances."""
|
||||
|
||||
@staticmethod
|
||||
def create_morphology(config: MorphologyConfig):
|
||||
from biorobot.brittle_star.mjcf.morphology.morphology import (
|
||||
MJCFBrittleStarMorphology,
|
||||
)
|
||||
from biorobot.brittle_star.mjcf.morphology.specification.default import (
|
||||
default_brittle_star_morphology_specification,
|
||||
)
|
||||
|
||||
spec = default_brittle_star_morphology_specification(
|
||||
num_arms=config.num_arms,
|
||||
num_segments_per_arm=list(config.segments_per_arm),
|
||||
use_p_control=config.use_p_control,
|
||||
use_torque_control=config.use_torque_control,
|
||||
)
|
||||
return MJCFBrittleStarMorphology(specification=spec)
|
||||
|
||||
@staticmethod
|
||||
def create_arena(config: ArenaConfig):
|
||||
from biorobot.brittle_star.mjcf.arena.aquarium import (
|
||||
AquariumArenaConfiguration,
|
||||
MJCFAquariumArena,
|
||||
)
|
||||
|
||||
arena_config = AquariumArenaConfiguration(**asdict(config))
|
||||
return MJCFAquariumArena(configuration=arena_config)
|
||||
|
||||
@staticmethod
|
||||
def create_environment_configuration(config: EnvConfig):
|
||||
# Import locally so the project can still be imported without these deps.
|
||||
from biorobot.brittle_star.environment.directed_locomotion.shared import (
|
||||
BrittleStarDirectedLocomotionEnvironmentConfiguration,
|
||||
)
|
||||
from biorobot.brittle_star.environment.light_escape.shared import (
|
||||
BrittleStarLightEscapeEnvironmentConfiguration,
|
||||
)
|
||||
|
||||
common = dict(
|
||||
joint_randomization_noise_scale=config.joint_randomization_noise_scale,
|
||||
render_mode="human",
|
||||
simulation_time=config.simulation_time,
|
||||
num_physics_steps_per_control_step=config.num_physics_steps_per_control_step,
|
||||
time_scale=config.time_scale,
|
||||
camera_ids=config.camera_ids,
|
||||
render_size=config.render_size,
|
||||
)
|
||||
|
||||
match config.task:
|
||||
case Task.DIRECTED_LOCOMOTION:
|
||||
return BrittleStarDirectedLocomotionEnvironmentConfiguration(
|
||||
target_distance=config.target_distance,
|
||||
**common,
|
||||
)
|
||||
case Task.LIGHT_ESCAPE:
|
||||
return BrittleStarLightEscapeEnvironmentConfiguration(
|
||||
light_perlin_noise_scale=config.light_perlin_noise_scale,
|
||||
**common,
|
||||
)
|
||||
case _:
|
||||
raise ValueError(f"Unsupported task: {config.task}")
|
||||
|
||||
@staticmethod
|
||||
def create_environment(
|
||||
backend: Backend,
|
||||
morphology_config: MorphologyConfig,
|
||||
arena_config: ArenaConfig,
|
||||
env_config: EnvConfig,
|
||||
) -> DualMuJoCoEnvironment:
|
||||
from biorobot.brittle_star.environment.directed_locomotion.dual import (
|
||||
BrittleStarDirectedLocomotionEnvironment,
|
||||
)
|
||||
from biorobot.brittle_star.environment.light_escape.dual import (
|
||||
BrittleStarLightEscapeEnvironment,
|
||||
)
|
||||
|
||||
morphology = BrittleStarEnvFactory.create_morphology(morphology_config)
|
||||
arena = BrittleStarEnvFactory.create_arena(arena_config)
|
||||
env_configuration = BrittleStarEnvFactory.create_environment_configuration(env_config)
|
||||
|
||||
match env_config.task:
|
||||
case Task.DIRECTED_LOCOMOTION:
|
||||
env_class = BrittleStarDirectedLocomotionEnvironment
|
||||
case Task.LIGHT_ESCAPE:
|
||||
env_class = BrittleStarLightEscapeEnvironment
|
||||
case _:
|
||||
raise ValueError(f"Unsupported task: {env_config.task}")
|
||||
|
||||
env = env_class.from_morphology_and_arena(
|
||||
morphology=morphology,
|
||||
arena=arena,
|
||||
configuration=env_configuration,
|
||||
backend=backend.value,
|
||||
)
|
||||
|
||||
from experiment_logger import get_logger
|
||||
|
||||
get_logger().info(f"Created {env_config.task.value} env on backend {backend.value}")
|
||||
|
||||
return env
|
||||
192
src/brittle_star_project/environment/obs_processing.py
Normal file
192
src/brittle_star_project/environment/obs_processing.py
Normal file
|
|
@ -0,0 +1,192 @@
|
|||
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",
|
||||
"joint_velocity",
|
||||
"joint_actuator_force",
|
||||
"actuator_force",
|
||||
}
|
||||
)
|
||||
|
||||
_SEGMENT_SCALED_KEYS = frozenset(
|
||||
{
|
||||
"segment_contact",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
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]],
|
||||
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:
|
||||
rot = new_obs["disk_rotation"]
|
||||
new_obs["disk_z_tilt"] = jnp.sqrt(jnp.pow(rot[0], 2) + jnp.pow(rot[1], 2))
|
||||
|
||||
if "unit_xy_direction_to_target" in new_obs:
|
||||
yaw = rot[2]
|
||||
unit_x, unit_y = new_obs["unit_xy_direction_to_target"]
|
||||
cos_yaw, sin_yaw = jnp.cos(yaw), jnp.sin(yaw)
|
||||
new_x = unit_x * cos_yaw + unit_y * sin_yaw
|
||||
new_y = -unit_x * sin_yaw + unit_y * cos_yaw
|
||||
new_obs["robot_direction_to_target"] = jnp.stack([new_x, new_y])
|
||||
|
||||
return new_obs
|
||||
|
||||
def _normalize_features(obs: dict) -> dict:
|
||||
normalized = {}
|
||||
for key, arr in obs.items():
|
||||
if key in bounds_dict:
|
||||
low, high = bounds_dict[key]
|
||||
if low == -1.0 and high == 1.0:
|
||||
normalized[key] = jnp.clip(arr, -1.0, 1.0)
|
||||
else:
|
||||
arr_clipped = jnp.clip(arr, low, high)
|
||||
normalized[key] = 2.0 * (arr_clipped - low) / (high - low) - 1.0
|
||||
else:
|
||||
normalized[key] = arr
|
||||
return normalized
|
||||
|
||||
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():
|
||||
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:
|
||||
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:
|
||||
"""
|
||||
Input:
|
||||
key -> (num_arms, feat_per_key)
|
||||
|
||||
Output:
|
||||
(num_arms, total_features)
|
||||
"""
|
||||
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)
|
||||
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))
|
||||
54
src/brittle_star_project/environment/padded_obs_wrapper.py
Normal file
54
src/brittle_star_project/environment/padded_obs_wrapper.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
"""Observation padding masks for amputated brittle star morphologies."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Sequence
|
||||
|
||||
import jax.numpy as jnp
|
||||
|
||||
|
||||
def compute_padding_masks(
|
||||
segments_per_arm: Sequence[int],
|
||||
reference_segments_per_arm: Sequence[int] = (4, 4, 4, 4, 4),
|
||||
) -> dict[str, Any]:
|
||||
"""Pre-compute boolean masks for spatial insertion of observations.
|
||||
|
||||
Args:
|
||||
segments_per_arm: The current (possibly amputated) morphology.
|
||||
reference_segments_per_arm: The full morphology that defines the expected size.
|
||||
|
||||
Returns:
|
||||
A dict containing 1D boolean masks and target sizes.
|
||||
"""
|
||||
if len(segments_per_arm) != len(reference_segments_per_arm):
|
||||
raise ValueError(
|
||||
f"Morphology mismatch: current has {len(segments_per_arm)} arms, "
|
||||
f"but reference requires {len(reference_segments_per_arm)} arms."
|
||||
)
|
||||
|
||||
mask_1x = []
|
||||
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}: "
|
||||
f"actual segments ({actual}) must be between 0 and reference ({ref})."
|
||||
)
|
||||
# 1x scaling (e.g., contacts: 1 value per segment)
|
||||
mask_1x.extend([True] * actual + [False] * (ref - actual))
|
||||
# 2x scaling (e.g., joints: 2 values per segment)
|
||||
mask_2x.extend([True] * (actual * 2) + [False] * ((ref - actual) * 2))
|
||||
|
||||
return {
|
||||
"mask_1x": jnp.array(mask_1x, dtype=bool),
|
||||
"mask_2x": jnp.array(mask_2x, dtype=bool),
|
||||
"target_size_1x": sum(reference_segments_per_arm),
|
||||
"target_size_2x": sum(reference_segments_per_arm) * 2,
|
||||
}
|
||||
43
src/brittle_star_project/evaluation/__init__.py
Normal file
43
src/brittle_star_project/evaluation/__init__.py
Normal file
|
|
@ -0,0 +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",
|
||||
]
|
||||
114
src/brittle_star_project/evaluation/checkpoint.py
Normal file
114
src/brittle_star_project/evaluation/checkpoint.py
Normal file
|
|
@ -0,0 +1,114 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import yaml
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from collections.abc import Mapping
|
||||
|
||||
import flax
|
||||
from omegaconf import OmegaConf
|
||||
|
||||
from brittle_star_project.environment.env_config import (
|
||||
MorphologyConfig,
|
||||
ArenaConfig,
|
||||
EnvConfig,
|
||||
ObservationBoundsConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainingConfig:
|
||||
"""Holds typed configurations extracted from a training run's metadata."""
|
||||
|
||||
morphology: MorphologyConfig
|
||||
arena: ArenaConfig
|
||||
environment: EnvConfig
|
||||
obs_bounds: ObservationBoundsConfig
|
||||
|
||||
|
||||
def load_params(path: Path) -> dict:
|
||||
"""Load model parameters from a .flax checkpoint file."""
|
||||
payload = path.read_bytes()
|
||||
restored = flax.serialization.msgpack_restore(payload)
|
||||
|
||||
sensor_params = None
|
||||
actor_params = None
|
||||
message_passer_params = None
|
||||
|
||||
# Extract params from restored checkpoint
|
||||
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, 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:
|
||||
sensor_params = params_part[0]
|
||||
actor_params = params_part[1]
|
||||
|
||||
if sensor_params is None or actor_params is None:
|
||||
raise ValueError(f"Could not extract sensor and actor params from checkpoint: {path}")
|
||||
|
||||
return {
|
||||
"sensor_params": sensor_params,
|
||||
"actor_params": actor_params,
|
||||
"message_passer_params": message_passer_params,
|
||||
}
|
||||
|
||||
|
||||
def load_metadata(model_path: Path, metadata_override_path: Path | None = None) -> dict:
|
||||
"""Discover and load the sidecar metadata YAML file."""
|
||||
if metadata_override_path is not None:
|
||||
metadata_path = metadata_override_path
|
||||
else:
|
||||
metadata_path = model_path.with_name(model_path.stem + "_metadata.yaml")
|
||||
|
||||
if not metadata_path.exists():
|
||||
raise FileNotFoundError(f"Could not find metadata YAML at {metadata_path}")
|
||||
with open(metadata_path, "r") as f:
|
||||
return yaml.safe_load(f)
|
||||
|
||||
|
||||
def metadata_to_configs(metadata: dict) -> TrainingConfig:
|
||||
"""Reconstruct typed configuration objects from a metadata dictionary."""
|
||||
trained_morphology = OmegaConf.to_object(
|
||||
OmegaConf.merge(OmegaConf.structured(MorphologyConfig), metadata.get("morphology", {}))
|
||||
)
|
||||
trained_arena = OmegaConf.to_object(
|
||||
OmegaConf.merge(OmegaConf.structured(ArenaConfig), metadata.get("arena", {}))
|
||||
)
|
||||
|
||||
env_dict = metadata.get("environment", {})
|
||||
if isinstance(env_dict.get("task"), str):
|
||||
from brittle_star_project.environment.env_types import Task
|
||||
|
||||
try:
|
||||
env_dict["task"] = Task[env_dict["task"]].name
|
||||
except Exception:
|
||||
try:
|
||||
env_dict["task"] = Task(env_dict["task"]).name
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
trained_environment = OmegaConf.to_object(
|
||||
OmegaConf.merge(OmegaConf.structured(EnvConfig), env_dict)
|
||||
)
|
||||
trained_obs_bounds = OmegaConf.to_object(
|
||||
OmegaConf.merge(
|
||||
OmegaConf.structured(ObservationBoundsConfig), metadata.get("obs_bounds", {})
|
||||
)
|
||||
)
|
||||
|
||||
return TrainingConfig(
|
||||
morphology=trained_morphology,
|
||||
arena=trained_arena,
|
||||
environment=trained_environment,
|
||||
obs_bounds=trained_obs_bounds,
|
||||
)
|
||||
176
src/brittle_star_project/evaluation/eval_env_builder.py
Normal file
176
src/brittle_star_project/evaluation/eval_env_builder.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
58
src/brittle_star_project/evaluation/evaluate.py
Normal file
58
src/brittle_star_project/evaluation/evaluate.py
Normal file
|
|
@ -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,
|
||||
)
|
||||
258
src/brittle_star_project/evaluation/evaluate_mjx.py
Normal file
258
src/brittle_star_project/evaluation/evaluate_mjx.py
Normal file
|
|
@ -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 `<run_dir>/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
|
||||
168
src/brittle_star_project/evaluation/policy.py
Normal file
168
src/brittle_star_project/evaluation/policy.py
Normal file
|
|
@ -0,0 +1,168 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any, Protocol
|
||||
|
||||
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
|
||||
|
||||
|
||||
class ControlPolicy(Protocol):
|
||||
"""Protocol for any policy that can produce actions from observations."""
|
||||
|
||||
def act(self, *, observations: dict[str, Any]) -> np.ndarray: ...
|
||||
|
||||
|
||||
class PolicyAgent:
|
||||
"""Wraps a trained Flax actor for deterministic inference."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
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,
|
||||
MessagePasser,
|
||||
)
|
||||
|
||||
# Infer layer sizes from params
|
||||
try:
|
||||
dense_params = (
|
||||
sensor_params.get("params", {})
|
||||
if isinstance(sensor_params, dict)
|
||||
else sensor_params["params"]
|
||||
)
|
||||
except Exception:
|
||||
dense_params = sensor_params
|
||||
|
||||
layer_sizes = []
|
||||
idx = 0
|
||||
while True:
|
||||
key = f"Dense_{idx}"
|
||||
if key not in dense_params:
|
||||
break
|
||||
|
||||
layer_sizes.append(int(np.asarray(dense_params[key]["kernel"]).shape[-1]))
|
||||
idx += 1
|
||||
|
||||
if not layer_sizes:
|
||||
raise ValueError("Could not infer Dense_* layers from sensor params")
|
||||
|
||||
self._sensor = GenericDenseLayersWithActivation(layer_sizes=layer_sizes)
|
||||
self._actor = Actor(action_dim=action_dim)
|
||||
|
||||
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,
|
||||
model_path: Path,
|
||||
*,
|
||||
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)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
def act(self, *, observations: dict[str, Any]) -> np.ndarray:
|
||||
"""Return deterministic action (actor mean, no exploration noise)."""
|
||||
batched_obs = jax.tree.map(lambda x: jnp.asarray(x)[None, ...], observations)
|
||||
obs = self._obs_processor(batched_obs)
|
||||
|
||||
hidden = 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()
|
||||
168
src/brittle_star_project/evaluation/rollout.py
Normal file
168
src/brittle_star_project/evaluation/rollout.py
Normal file
|
|
@ -0,0 +1,168 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from brittle_star_project import BrittleStarEnv
|
||||
from brittle_star_project.evaluation.policy import ControlPolicy
|
||||
|
||||
|
||||
@dataclass
|
||||
class EpisodeResult:
|
||||
return_: float
|
||||
length: int
|
||||
reached_target: bool
|
||||
final_xy_dist: float | None
|
||||
initial_target_distance: float | None
|
||||
|
||||
|
||||
def _get_observations(state: Any) -> dict[str, Any] | None:
|
||||
return getattr(state, "observations", None)
|
||||
|
||||
|
||||
def _get_xy_distance_to_target(observations: dict[str, Any]) -> float | None:
|
||||
return float(np.asarray(observations["xy_distance_to_target"]).reshape(-1)[0])
|
||||
|
||||
|
||||
def _target_reached(*, state: Any) -> bool:
|
||||
return bool(getattr(state, "terminated", False) or getattr(state, "truncated", False))
|
||||
|
||||
|
||||
def _maybe_clip_action(
|
||||
action: np.ndarray,
|
||||
low: np.ndarray | None,
|
||||
high: np.ndarray | None,
|
||||
) -> np.ndarray:
|
||||
if low is None or high is None:
|
||||
return action
|
||||
low = np.asarray(low, dtype=np.float32).ravel()
|
||||
high = np.asarray(high, dtype=np.float32).ravel()
|
||||
if low.shape != action.shape or high.shape != action.shape:
|
||||
return action
|
||||
return np.clip(action, low, high)
|
||||
|
||||
|
||||
def rollout_headless(
|
||||
*,
|
||||
env: BrittleStarEnv,
|
||||
policy: ControlPolicy,
|
||||
seed: int,
|
||||
max_steps: int,
|
||||
action_low: np.ndarray | None,
|
||||
action_high: np.ndarray | None,
|
||||
action_mask: np.ndarray | None = None,
|
||||
) -> EpisodeResult:
|
||||
"""Run an episode headlessly and return the result."""
|
||||
state = env.reset(seed=seed)
|
||||
|
||||
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
|
||||
for _ in range(int(max_steps)):
|
||||
obs_dict = observations or {}
|
||||
|
||||
action = policy.act(observations=obs_dict)
|
||||
if action_mask is not None:
|
||||
action = action[action_mask]
|
||||
action = _maybe_clip_action(action, action_low, action_high)
|
||||
|
||||
state = env.step(state=state, action=action)
|
||||
steps += 1
|
||||
|
||||
observations = _get_observations(state)
|
||||
cur_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
if prev_dist is not None and cur_dist is not None:
|
||||
ep_return += prev_dist - cur_dist
|
||||
prev_dist = cur_dist
|
||||
|
||||
reached_target = _target_reached(state=state)
|
||||
if reached_target:
|
||||
break
|
||||
|
||||
final_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
return EpisodeResult(
|
||||
return_=ep_return,
|
||||
length=steps,
|
||||
reached_target=reached_target,
|
||||
final_xy_dist=final_dist,
|
||||
initial_target_distance=initial_target_distance,
|
||||
)
|
||||
|
||||
|
||||
def rollout_viewer(
|
||||
*,
|
||||
env: BrittleStarEnv,
|
||||
policy: ControlPolicy,
|
||||
seed: int,
|
||||
state: Any,
|
||||
control_dt: float,
|
||||
max_steps: int | None,
|
||||
action_low: np.ndarray | None,
|
||||
action_high: np.ndarray | None,
|
||||
action_mask: np.ndarray | None = None,
|
||||
) -> None:
|
||||
"""Run an episode using the interactive MuJoCo viewer."""
|
||||
import mujoco.viewer
|
||||
|
||||
model = state.mj_model
|
||||
data = state.mj_data
|
||||
|
||||
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 _ 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]
|
||||
action = _maybe_clip_action(action, action_low, action_high)
|
||||
|
||||
with viewer.lock():
|
||||
state = env.step(state=state, action=action)
|
||||
|
||||
if not viewer.is_running():
|
||||
break
|
||||
viewer.sync()
|
||||
|
||||
steps += 1
|
||||
|
||||
observations = _get_observations(state)
|
||||
cur_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
if prev_dist is not None and cur_dist is not None:
|
||||
episode_return += prev_dist - cur_dist
|
||||
prev_dist = cur_dist
|
||||
|
||||
reached_target = _target_reached(state=state)
|
||||
if reached_target:
|
||||
break
|
||||
|
||||
remaining = control_dt - (time.time() - step_start)
|
||||
if remaining > 0:
|
||||
time.sleep(remaining)
|
||||
|
||||
dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
dist_str = "n/a" if dist is None else f"{dist:.3f}"
|
||||
print(
|
||||
"episode done: "
|
||||
f"return={episode_return:.6f}, len={steps}, "
|
||||
f"target_reached={reached_target}, final_xy_dist={dist_str}"
|
||||
)
|
||||
149
src/brittle_star_project/evaluation/video.py
Normal file
149
src/brittle_star_project/evaluation/video.py
Normal file
|
|
@ -0,0 +1,149 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import yaml
|
||||
|
||||
from brittle_star_project import BrittleStarEnv
|
||||
from brittle_star_project.evaluation.policy import ControlPolicy
|
||||
from brittle_star_project.evaluation.rollout import (
|
||||
EpisodeResult,
|
||||
_get_observations,
|
||||
_get_xy_distance_to_target,
|
||||
_target_reached,
|
||||
_maybe_clip_action,
|
||||
)
|
||||
|
||||
|
||||
def create_evaluation_dir(model_path: Path) -> Path:
|
||||
"""Create a unique timestamped directory for saving evaluation results."""
|
||||
timestamp = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S")
|
||||
eval_dir = model_path.parent / f"{model_path.stem}_evaluations" / f"eval_{timestamp}"
|
||||
eval_dir.mkdir(parents=True, exist_ok=True)
|
||||
return eval_dir
|
||||
|
||||
|
||||
def save_evaluation_metadata(
|
||||
eval_dir: Path,
|
||||
*,
|
||||
morphology_override_path: str | None,
|
||||
seed: int,
|
||||
max_steps: int | None,
|
||||
result: EpisodeResult,
|
||||
) -> None:
|
||||
"""Save metadata about the evaluation run."""
|
||||
metadata = {
|
||||
"timestamp": datetime.datetime.now().isoformat(),
|
||||
"morphology_override": morphology_override_path,
|
||||
"seed": seed,
|
||||
"max_steps": max_steps,
|
||||
"result": {
|
||||
"return": float(result.return_),
|
||||
"length": int(result.length),
|
||||
"reached_target": bool(result.reached_target),
|
||||
"final_xy_dist": float(result.final_xy_dist)
|
||||
if result.final_xy_dist is not None
|
||||
else None,
|
||||
},
|
||||
}
|
||||
with open(eval_dir / "evaluation_metadata.yaml", "w") as f:
|
||||
yaml.safe_dump(metadata, f, sort_keys=False)
|
||||
|
||||
|
||||
def record_episode(
|
||||
*,
|
||||
env: BrittleStarEnv,
|
||||
policy: ControlPolicy,
|
||||
seed: int,
|
||||
max_steps: int,
|
||||
action_low: np.ndarray | None,
|
||||
action_high: np.ndarray | None,
|
||||
action_mask: np.ndarray | None = None,
|
||||
output_path: Path,
|
||||
camera_id: int = 1,
|
||||
fps: int = 60,
|
||||
width: int = 640,
|
||||
height: int = 480,
|
||||
) -> EpisodeResult:
|
||||
"""Run an episode headlessly and record a video using MuJoCo's Renderer and imageio.
|
||||
|
||||
Args:
|
||||
env: The environment.
|
||||
policy: The policy agent.
|
||||
seed: Random seed.
|
||||
max_steps: Maximum number of steps.
|
||||
action_low: Minimum action values.
|
||||
action_high: Maximum action values.
|
||||
action_mask: Boolean mask for the actions.
|
||||
output_path: Where to save the .mp4 file.
|
||||
camera_id: Camera index to use for rendering (1 is usually close-up).
|
||||
fps: Frames per second for the video.
|
||||
width: Video width.
|
||||
height: Video height.
|
||||
"""
|
||||
try:
|
||||
import imageio
|
||||
import mujoco
|
||||
except ImportError as e:
|
||||
raise ImportError(
|
||||
"Video recording requires 'imageio' and 'mujoco'. "
|
||||
"Please install the evaluation dependencies: `uv pip install .[evaluation]`"
|
||||
) from e
|
||||
|
||||
state = env.reset(seed=seed)
|
||||
model = state.mj_model
|
||||
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 = []
|
||||
steps = 0
|
||||
|
||||
for _ in range(int(max_steps)):
|
||||
# Capture frame
|
||||
renderer.update_scene(data, camera=camera_id)
|
||||
frames.append(renderer.render())
|
||||
|
||||
# Step environment
|
||||
obs_dict = observations or {}
|
||||
action = policy.act(observations=obs_dict)
|
||||
if action_mask is not None:
|
||||
action = action[action_mask]
|
||||
action = _maybe_clip_action(action, action_low, action_high)
|
||||
|
||||
state = env.step(state=state, action=action)
|
||||
steps += 1
|
||||
|
||||
observations = _get_observations(state)
|
||||
cur_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
if prev_dist is not None and cur_dist is not None:
|
||||
ep_return += prev_dist - cur_dist
|
||||
prev_dist = cur_dist
|
||||
|
||||
reached_target = _target_reached(state=state)
|
||||
if reached_target:
|
||||
break
|
||||
|
||||
# Capture final frame
|
||||
renderer.update_scene(data, camera=camera_id)
|
||||
frames.append(renderer.render())
|
||||
renderer.close()
|
||||
|
||||
# Save video
|
||||
imageio.mimsave(str(output_path), frames, fps=fps)
|
||||
|
||||
final_dist = _get_xy_distance_to_target(observations) if observations else None
|
||||
return EpisodeResult(
|
||||
return_=ep_return,
|
||||
length=steps,
|
||||
reached_target=reached_target,
|
||||
final_xy_dist=final_dist,
|
||||
initial_target_distance=initial_dist,
|
||||
)
|
||||
201
src/brittle_star_project/ppo.py
Normal file
201
src/brittle_star_project/ppo.py
Normal file
|
|
@ -0,0 +1,201 @@
|
|||
from functools import partial
|
||||
|
||||
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_apply,
|
||||
actor_apply,
|
||||
critic_apply,
|
||||
feature_extractor_apply,
|
||||
message_passer=None,
|
||||
):
|
||||
self.args = args
|
||||
|
||||
if not message_passer:
|
||||
message_passer = identity
|
||||
|
||||
self.ppo_loss_grad_fn = jax.value_and_grad(
|
||||
partial(
|
||||
ppo_loss,
|
||||
args=args,
|
||||
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,
|
||||
)
|
||||
|
||||
# This PPO class should be initialized only once,
|
||||
# or this function will need to recompile
|
||||
@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
|
||||
|
||||
def update_epoch(carry, _):
|
||||
agent_state, key = carry
|
||||
key, subkey = jax.random.split(key)
|
||||
|
||||
def flatten(x):
|
||||
return x.reshape((-1,) + x.shape[2:])
|
||||
|
||||
def convert_data(x):
|
||||
x = jax.random.permutation(subkey, x)
|
||||
return jnp.reshape(x, (args.num_minibatches, -1) + x.shape[1:])
|
||||
|
||||
flatten_storage = jax.tree.map(flatten, storage)
|
||||
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,
|
||||
minibatch.actions,
|
||||
minibatch.logprobs,
|
||||
minibatch.advantages,
|
||||
minibatch.returns,
|
||||
)
|
||||
agent_state = agent_state.apply_gradients(grads=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) = 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
|
||||
|
||||
|
||||
"""
|
||||
Should be ok to use partial here, since the references to network,
|
||||
actor and critic should not change at runtime
|
||||
The cost of seperating concerns is to somehow pass these values
|
||||
that are now not in the same scope
|
||||
"""
|
||||
|
||||
|
||||
@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: 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)
|
||||
|
||||
# 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))
|
||||
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
|
||||
|
||||
|
||||
def ppo_loss(
|
||||
params,
|
||||
x,
|
||||
a,
|
||||
logp,
|
||||
mb_advantages,
|
||||
mb_returns,
|
||||
args,
|
||||
sensor_apply,
|
||||
actor_apply,
|
||||
message_passer,
|
||||
critic_apply,
|
||||
feature_extractor_apply,
|
||||
):
|
||||
newlogprob, entropy, newvalue = get_action_and_value(
|
||||
sensor_apply,
|
||||
actor_apply,
|
||||
message_passer,
|
||||
critic_apply,
|
||||
feature_extractor_apply,
|
||||
params,
|
||||
x,
|
||||
a,
|
||||
)
|
||||
logratio = newlogprob - logp
|
||||
ratio = jnp.exp(logratio)
|
||||
approx_kl = ((ratio - 1) - logratio).mean()
|
||||
|
||||
if args.norm_adv:
|
||||
mb_advantages = (mb_advantages - mb_advantages.mean()) / (mb_advantages.std() + 1e-8)
|
||||
|
||||
pg_loss1 = -mb_advantages * ratio
|
||||
pg_loss2 = -mb_advantages * jnp.clip(ratio, 1 - args.clip_coef, 1 + args.clip_coef)
|
||||
pg_loss = jnp.maximum(pg_loss1, pg_loss2).mean()
|
||||
|
||||
v_loss = 0.5 * ((newvalue - mb_returns) ** 2).mean()
|
||||
entropy_loss = entropy.mean()
|
||||
loss = pg_loss - args.ent_coef * entropy_loss + v_loss * args.vf_coef
|
||||
return loss, (pg_loss, v_loss, entropy_loss, jax.lax.stop_gradient(approx_kl))
|
||||
|
||||
|
||||
def identity(_, hidden):
|
||||
"""
|
||||
Used for seamless jax integration,
|
||||
avoids having branching inside jitted function,
|
||||
used as message_passer in case it is not given,
|
||||
(in case of centralized lvl)
|
||||
"""
|
||||
|
||||
return hidden
|
||||
988
src/brittle_star_project/trainers/PPOTrainer.py
Normal file
988
src/brittle_star_project/trainers/PPOTrainer.py
Normal file
|
|
@ -0,0 +1,988 @@
|
|||
import datetime
|
||||
import random
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
from functools import partial
|
||||
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
|
||||
|
||||
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?
|
||||
|
||||
|
||||
@logged_jit
|
||||
def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray:
|
||||
return jnp.clip(action, low, high)
|
||||
|
||||
|
||||
def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float:
|
||||
var_returns = jnp.var(returns)
|
||||
explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8)
|
||||
return float(explained_var)
|
||||
|
||||
|
||||
@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: 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,
|
||||
action_low,
|
||||
action_high,
|
||||
):
|
||||
# (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 = 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
|
||||
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,
|
||||
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, 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,
|
||||
)
|
||||
|
||||
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,
|
||||
actions=raw_action,
|
||||
raw_actions=raw_action,
|
||||
logprobs=logprob,
|
||||
dones=done,
|
||||
values=value,
|
||||
rewards=reward,
|
||||
means=mean,
|
||||
stds=std,
|
||||
returns=jnp.zeros_like(reward),
|
||||
advantages=jnp.zeros_like(reward),
|
||||
)
|
||||
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):
|
||||
"""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"]
|
||||
).squeeze(-1)
|
||||
|
||||
env_reward = next_env_state.reward
|
||||
clipped_env_reward = jnp.clip(100 * env_reward, -10, 10)
|
||||
|
||||
time_penalty = 0.1
|
||||
distance_penalty = jnp.clip(0.5 * delta_distance, -0.5, 0.5)
|
||||
penalty = time_penalty + distance_penalty
|
||||
|
||||
return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty)
|
||||
|
||||
|
||||
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_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
|
||||
new_episode_length = episode_stats.episode_lengths + 1
|
||||
|
||||
episode_stats = episode_stats.replace(
|
||||
episode_returns=new_episode_return * (1 - done),
|
||||
episode_lengths=new_episode_length * (1 - done),
|
||||
returned_episode_returns=jnp.where(
|
||||
done, new_episode_return, episode_stats.returned_episode_returns
|
||||
),
|
||||
returned_episode_lengths=jnp.where(
|
||||
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, 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,
|
||||
env_state,
|
||||
next_obs,
|
||||
next_done,
|
||||
key,
|
||||
max_steps,
|
||||
step_env_fn,
|
||||
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,
|
||||
):
|
||||
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,
|
||||
terminated_any0,
|
||||
truncated_any0,
|
||||
),
|
||||
(),
|
||||
max_steps,
|
||||
)
|
||||
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):
|
||||
advantages = carry
|
||||
nextdone, nextvalues, curvalues, reward = inp
|
||||
nextnonterminal = 1.0 - nextdone
|
||||
delta = reward + gamma * nextvalues * nextnonterminal - curvalues
|
||||
advantages = delta + gamma * gae_lambda * nextnonterminal * advantages
|
||||
return advantages, advantages
|
||||
|
||||
|
||||
def _compute_gae_jit(
|
||||
agent_state,
|
||||
storage,
|
||||
next_obs,
|
||||
next_done,
|
||||
gamma,
|
||||
gae_lambda,
|
||||
num_envs,
|
||||
feature_extractor,
|
||||
critic,
|
||||
):
|
||||
next_value = apply_shared(
|
||||
critic,
|
||||
agent_state.params["critic_params"],
|
||||
apply_shared(feature_extractor, agent_state.params["feature_extractor_params"], next_obs),
|
||||
).squeeze(-1)
|
||||
|
||||
advantages = jnp.zeros((num_envs,))
|
||||
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
|
||||
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
|
||||
_, advantages = jax.lax.scan(
|
||||
partial(_compute_gae_once, gamma=gamma, gae_lambda=gae_lambda),
|
||||
advantages,
|
||||
(dones[1:], values[1:], values[:-1], storage.rewards),
|
||||
reverse=True,
|
||||
)
|
||||
returns = advantages + storage.values
|
||||
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
|
||||
return storage.replace(advantages=advantages, returns=returns)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainingMeasurements:
|
||||
loss: jnp.ndarray
|
||||
pg_loss: jnp.ndarray
|
||||
v_loss: jnp.ndarray
|
||||
entropy_loss: jnp.ndarray
|
||||
approx_kl: jnp.ndarray
|
||||
avg_episodic_return: float
|
||||
explained_variance: float
|
||||
num_terminated: int
|
||||
num_truncated: int
|
||||
avg_terminated_length: Any
|
||||
avg_truncated_length: Any
|
||||
|
||||
|
||||
class PPOTrainer:
|
||||
def __init__(
|
||||
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
|
||||
self.logger = get_logger()
|
||||
|
||||
# Derived runtime fields
|
||||
self.batch_size = self.ppo.num_envs * self.ppo.num_steps
|
||||
self.num_iterations = self.ppo.total_timesteps // self.batch_size
|
||||
|
||||
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.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 = 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 = logged_jit(
|
||||
partial(
|
||||
_compute_gae_jit,
|
||||
num_envs=self.ppo.num_envs,
|
||||
gamma=self.ppo.gamma,
|
||||
gae_lambda=self.ppo.gae_lambda,
|
||||
feature_extractor=self.feature_extractor,
|
||||
critic=self.critic,
|
||||
)
|
||||
)
|
||||
|
||||
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}")
|
||||
|
||||
random.seed(self.experiment.seed)
|
||||
np.random.seed(self.experiment.seed)
|
||||
|
||||
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])
|
||||
critic = OneDenseLayerMLP()
|
||||
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, 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
|
||||
|
||||
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,
|
||||
message_passer_params,
|
||||
)
|
||||
),
|
||||
tx=optax.chain(
|
||||
optax.clip_by_global_norm(self.ppo.max_grad_norm),
|
||||
optax.inject_hyperparams(optax.adam)(
|
||||
learning_rate=partial(
|
||||
_linear_schedule,
|
||||
minibatch_count=self.ppo.num_minibatches,
|
||||
update_epochs=self.ppo.update_epochs,
|
||||
num_iterations=self.num_iterations,
|
||||
learning_rate=self.ppo.learning_rate,
|
||||
)
|
||||
if self.ppo.anneal_lr
|
||||
else self.ppo.learning_rate,
|
||||
eps=1e-5,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def _init_episode_stats(self) -> EpisodeStatistics:
|
||||
self.logger.info("[EPISODE STATS]: Initializing episode stats...")
|
||||
|
||||
return EpisodeStatistics(
|
||||
episode_returns=jnp.zeros(self.ppo.num_envs, dtype=jnp.float32),
|
||||
episode_lengths=jnp.zeros(self.ppo.num_envs, dtype=jnp.int32),
|
||||
returned_episode_returns=jnp.zeros(self.ppo.num_envs, jnp.float32),
|
||||
returned_episode_lengths=jnp.zeros(self.ppo.num_envs, dtype=jnp.int32),
|
||||
)
|
||||
|
||||
def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]:
|
||||
return self._rollout_jit(
|
||||
self.agent_state,
|
||||
self.episode_stats,
|
||||
env_state,
|
||||
next_obs,
|
||||
next_done,
|
||||
self.key,
|
||||
)
|
||||
|
||||
def _compute_gae(self, storage, next_obs, next_done) -> Storage:
|
||||
return self._compute_gae_jit(
|
||||
self.agent_state,
|
||||
storage,
|
||||
next_obs,
|
||||
next_done,
|
||||
)
|
||||
|
||||
def _log(
|
||||
self,
|
||||
global_step,
|
||||
episode_stats,
|
||||
start_time,
|
||||
iteration_time_start,
|
||||
training_measurements,
|
||||
storage,
|
||||
):
|
||||
data = jax.device_get(
|
||||
{
|
||||
"rewards": storage.rewards,
|
||||
"values": storage.values,
|
||||
"returns": storage.returns,
|
||||
"advantages": storage.advantages,
|
||||
}
|
||||
)
|
||||
|
||||
rollout_metrics = {
|
||||
"rollout/reward_mean": float(np.mean(data["rewards"])),
|
||||
"rollout/return_mean": float(np.mean(data["returns"])),
|
||||
"rollout/value_mean": float(np.mean(data["values"])),
|
||||
"rollout/advantage_mean": float(np.mean(data["advantages"])),
|
||||
"rollout/advantage_std": float(np.std(data["advantages"])),
|
||||
"rollout/value_vs_return_mse": float(np.mean((data["values"] - data["returns"]) ** 2)),
|
||||
}
|
||||
|
||||
metrics = {
|
||||
"charts/episodic_return": training_measurements.avg_episodic_return,
|
||||
"charts/episodic_length": float(
|
||||
np.mean(jax.device_get(episode_stats.returned_episode_lengths))
|
||||
),
|
||||
"charts/explained_variance": training_measurements.explained_variance,
|
||||
"losses/value_loss": training_measurements.v_loss[-1, -1].item(),
|
||||
"losses/policy_loss": training_measurements.pg_loss[-1, -1].item(),
|
||||
"losses/entropy": training_measurements.entropy_loss[-1, -1].item(),
|
||||
"losses/approx_kl": training_measurements.approx_kl[-1, -1].item(),
|
||||
"charts/learning_rate": self.agent_state.opt_state[1]
|
||||
.hyperparams["learning_rate"]
|
||||
.item(),
|
||||
"charts/SPS": int(global_step / (time.time() - start_time)),
|
||||
"charts/SPS_update": int(
|
||||
self.ppo.num_envs * self.ppo.num_steps / (time.time() - iteration_time_start)
|
||||
),
|
||||
"termi_trunci/num_terminated": training_measurements.num_terminated,
|
||||
"termi_trunci/num_truncated": training_measurements.num_truncated,
|
||||
"termi_trunci/avg_terminated_ep_length": training_measurements.avg_terminated_length,
|
||||
"termi_trunci/avg_truncated_ep_length": training_measurements.avg_truncated_length,
|
||||
**rollout_metrics,
|
||||
}
|
||||
|
||||
self.logger.log(metrics, step=global_step)
|
||||
|
||||
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,
|
||||
next_obs,
|
||||
next_done,
|
||||
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()}")
|
||||
|
||||
self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = (
|
||||
self._ppo.update_ppo(self.agent_state, storage, self.key)
|
||||
)
|
||||
|
||||
if iteration == 1:
|
||||
self.logger.log_non_interactive(f"First PPO update completed: {time.ctime()}")
|
||||
|
||||
avg_episodic_return = float(
|
||||
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
|
||||
)
|
||||
|
||||
explained_var = _compute_explained_variance(storage.values, storage.returns)
|
||||
|
||||
terminated = terminated_any
|
||||
truncated = truncated_any
|
||||
episode_lengths = self.episode_stats.returned_episode_lengths
|
||||
|
||||
num_terminated = int(jnp.sum(terminated).item())
|
||||
num_truncated = int(jnp.sum(truncated).item())
|
||||
|
||||
avg_terminated_length = jnp.sum(episode_lengths * terminated) / jnp.maximum(
|
||||
jnp.sum(terminated), 1
|
||||
)
|
||||
|
||||
avg_truncated_length = jnp.sum(episode_lengths * truncated) / jnp.maximum(
|
||||
jnp.sum(truncated), 1
|
||||
)
|
||||
|
||||
return (
|
||||
next_env_state,
|
||||
next_obs,
|
||||
next_done,
|
||||
TrainingMeasurements(
|
||||
loss=loss,
|
||||
pg_loss=pg_loss,
|
||||
v_loss=v_loss,
|
||||
entropy_loss=entropy_loss,
|
||||
approx_kl=approx_kl,
|
||||
avg_episodic_return=avg_episodic_return,
|
||||
explained_variance=explained_var,
|
||||
num_terminated=num_terminated,
|
||||
num_truncated=num_truncated,
|
||||
avg_terminated_length=avg_terminated_length,
|
||||
avg_truncated_length=avg_truncated_length,
|
||||
),
|
||||
storage,
|
||||
)
|
||||
|
||||
def _close(self):
|
||||
self.env.close()
|
||||
|
||||
def _save_model(self, model_path: str):
|
||||
self.logger.info("[SAVE]: Saving the final model...")
|
||||
self.logger.save_final_model(params=self.agent_state.params, metadata=asdict(self.cfg))
|
||||
|
||||
def _save_checkpoint(self, iteration: int):
|
||||
self.logger.info(f"[SAVE]: Saving checkpoint at iteration {iteration}...")
|
||||
self.logger.save_checkpoint(
|
||||
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.
|
||||
Closes the environment at the end of training.
|
||||
"""
|
||||
self.logger.info(f"running name: {self.run_name}")
|
||||
|
||||
self.logger.info("[TRAIN]: Resetting environment...")
|
||||
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()}")
|
||||
|
||||
global_step = 0
|
||||
start_time = time.time()
|
||||
|
||||
iter_bar = self.logger.progress_bar(range(1, self.num_iterations + 1))
|
||||
for iteration in iter_bar:
|
||||
iteration_time_start = time.time()
|
||||
|
||||
env_state, next_obs, next_done, training_measurements, storage = self._step(
|
||||
env_state, next_obs, next_done, iteration=iteration
|
||||
)
|
||||
|
||||
global_step += self.ppo.num_steps * self.ppo.num_envs
|
||||
self._log(
|
||||
global_step,
|
||||
self.episode_stats,
|
||||
start_time,
|
||||
iteration_time_start,
|
||||
training_measurements,
|
||||
storage,
|
||||
)
|
||||
|
||||
sps = int(global_step / (time.time() - start_time))
|
||||
remaining_steps = self.ppo.total_timesteps - global_step
|
||||
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
|
||||
eta_str = str(datetime.timedelta(seconds=eta_seconds))
|
||||
|
||||
self.logger.log_non_interactive(
|
||||
f"Iteration {iteration}/{self.num_iterations} | "
|
||||
f"Step {global_step}/{self.ppo.total_timesteps} | "
|
||||
f"SPS {sps} | "
|
||||
f"Return {training_measurements.avg_episodic_return:.4f} | "
|
||||
f"ETA {eta_str}"
|
||||
)
|
||||
|
||||
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")
|
||||
break
|
||||
|
||||
if self.logging_cfg.save_model:
|
||||
model_path = f"{self.run_dir}/{self.experiment.exp_name}.cleanrl_model"
|
||||
self._save_model(model_path=model_path)
|
||||
|
||||
self._close()
|
||||
0
src/brittle_star_project/trainers/__init__.py
Normal file
0
src/brittle_star_project/trainers/__init__.py
Normal file
3
src/brittle_star_project/utils/__init__.py
Normal file
3
src/brittle_star_project/utils/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .logged_jit import logged_jit
|
||||
|
||||
__all__ = ["logged_jit"]
|
||||
17
src/brittle_star_project/utils/logged_jit.py
Normal file
17
src/brittle_star_project/utils/logged_jit.py
Normal file
|
|
@ -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)
|
||||
19
src/experiment_logger/__init__.py
Normal file
19
src/experiment_logger/__init__.py
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
"""Unified logging framework for machine learning experiments.
|
||||
|
||||
This package provides a unified interface for logging to multiple backends
|
||||
(WandB, disk, stdout) simultaneously, ensuring no data loss.
|
||||
"""
|
||||
|
||||
from experiment_logger.unified_logger import UnifiedLogger, get_logger, init_logger
|
||||
from experiment_logger.simple_logger import SimpleLogger
|
||||
from experiment_logger.wandb_utils import finish_wandb, init_wandb
|
||||
|
||||
__all__ = [
|
||||
"UnifiedLogger",
|
||||
"SimpleLogger",
|
||||
"get_logger",
|
||||
"init_logger",
|
||||
"init_wandb",
|
||||
"finish_wandb",
|
||||
]
|
||||
__version__ = "0.1.0"
|
||||
36
src/experiment_logger/config_logger.py
Normal file
36
src/experiment_logger/config_logger.py
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoggingConfig:
|
||||
track: bool = False
|
||||
wandb_project_name: str = "default-project"
|
||||
wandb_entity: Optional[str] = "SEL3-2026-Groep-4"
|
||||
capture_video: bool = False
|
||||
|
||||
# Local Saving
|
||||
save_model: bool = True # Final model
|
||||
save_checkpoints: bool = True # Intermediate checkpoints
|
||||
checkpoint_frequency: int = 100
|
||||
|
||||
# Remote Uploading (WandB Artifacts)
|
||||
upload_final_model: bool = False
|
||||
upload_checkpoints: bool = False
|
||||
|
||||
hf_entity: str = ""
|
||||
|
||||
def __post_init__(self):
|
||||
if self.upload_final_model and not (self.track and self.save_model):
|
||||
raise ValueError(
|
||||
"Configuration Error: 'upload_final_model' is True, but it requires "
|
||||
"both 'track' and 'save_model' to also be True."
|
||||
)
|
||||
if self.upload_checkpoints and not (self.track and self.save_checkpoints):
|
||||
raise ValueError(
|
||||
"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).
|
||||
1417
src/experiment_logger/index.html
Normal file
1417
src/experiment_logger/index.html
Normal file
File diff suppressed because it is too large
Load diff
85
src/experiment_logger/simple_logger.py
Normal file
85
src/experiment_logger/simple_logger.py
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
"""Simple terminal logger for running without external backends.
|
||||
|
||||
This is used for standalone package usage where WandB or TensorBoard are not desired.
|
||||
It preserves the same API as UnifiedLogger but simply prints to stdout.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
class SimpleLogger:
|
||||
"""Simple logger that implements the UnifiedLogger interface via print statements."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
run_name: str = "simple_run",
|
||||
full_config: Optional[Dict[str, Any]] = None,
|
||||
logging_cfg: Optional[Any] = None,
|
||||
base_dir: str = "runs",
|
||||
save_code: bool = False,
|
||||
log_level: int = logging.INFO,
|
||||
_set_as_global: bool = False,
|
||||
):
|
||||
self.is_interactive = True
|
||||
self.run_name = run_name
|
||||
self.full_config = full_config or {}
|
||||
print(f"[INIT] SimpleLogger initialized for run: {run_name}")
|
||||
|
||||
def set_level(self, level: int):
|
||||
pass
|
||||
|
||||
def log_non_interactive(self, msg: str, *args, **kwargs):
|
||||
"""In SimpleLogger, we just print everything as we assume interactive use."""
|
||||
self.info(msg, *args, **kwargs)
|
||||
|
||||
def progress_bar(self, iterable=None, *args, **kwargs):
|
||||
"""Standard tqdm wrapper that falls back to range if tqdm is missing."""
|
||||
try:
|
||||
import tqdm
|
||||
|
||||
return tqdm.tqdm(iterable, *args, **kwargs)
|
||||
except ImportError:
|
||||
return iterable
|
||||
|
||||
def info(self, msg: str, *args, **kwargs):
|
||||
print(f"[INFO] {msg}")
|
||||
|
||||
def warning(self, msg: str, *args, **kwargs):
|
||||
print(f"[WARNING] {msg}")
|
||||
|
||||
def error(self, msg: str, *args, **kwargs):
|
||||
print(f"[ERROR] {msg}")
|
||||
|
||||
def debug(self, msg: str, *args, **kwargs):
|
||||
print(f"[DEBUG] {msg}")
|
||||
|
||||
def log(self, metrics: Dict[str, Any], step: Optional[int] = None, commit: bool = True):
|
||||
step_str = f"Step {step}" if step is not None else "Log"
|
||||
metric_str = ", ".join(f"{k}: {v}" for k, v in metrics.items())
|
||||
print(f"[{step_str}] {metric_str}")
|
||||
|
||||
def save_checkpoint(
|
||||
self,
|
||||
params: Any,
|
||||
step: int,
|
||||
prefix: str = "checkpoint",
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
print(f"[SAVE] Checkpoint '{prefix}' would be saved at step {step} (SimpleLogger: No-Op)")
|
||||
|
||||
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}")
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
self.finish()
|
||||
470
src/experiment_logger/unified_logger.py
Normal file
470
src/experiment_logger/unified_logger.py
Normal file
|
|
@ -0,0 +1,470 @@
|
|||
"""Unified logger that writes to multiple backends simultaneously.
|
||||
|
||||
This logger ensures all experimental data is preserved by writing to:
|
||||
1. Weights & Biases (when available)
|
||||
2. Local disk (JSON files, model checkpoints, run.log)
|
||||
3. stdout (for real-time monitoring)
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
import logging
|
||||
import yaml
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import flax
|
||||
import jax.numpy as jnp
|
||||
import numpy as np
|
||||
|
||||
from experiment_logger.wandb_utils import finish_wandb, init_wandb
|
||||
from experiment_logger.config_logger import LoggingConfig
|
||||
|
||||
# Global storage for the active logger and the proxy singleton
|
||||
_active_logger: Optional[Any] = None
|
||||
_proxy_instance: Optional["LoggerProxy"] = None
|
||||
|
||||
|
||||
def _sanitize_for_yaml(obj: Any) -> Any:
|
||||
"""Convert non-primitive values into YAML-safe structures.
|
||||
|
||||
In particular, avoids PyYAML serializing Enums as
|
||||
``!!python/object/apply:...`` which OmegaConf will not load.
|
||||
"""
|
||||
|
||||
if isinstance(obj, Enum):
|
||||
return obj.name
|
||||
if isinstance(obj, Path):
|
||||
return str(obj)
|
||||
if isinstance(obj, (np.generic, jnp.ndarray)):
|
||||
try:
|
||||
return obj.item()
|
||||
except Exception:
|
||||
pass
|
||||
if isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
if isinstance(obj, dict):
|
||||
return {str(k): _sanitize_for_yaml(v) for k, v in obj.items()}
|
||||
if isinstance(obj, list):
|
||||
return [_sanitize_for_yaml(v) for v in obj]
|
||||
if isinstance(obj, tuple):
|
||||
return [_sanitize_for_yaml(v) for v in obj]
|
||||
return obj
|
||||
|
||||
|
||||
def get_logger() -> "LoggerProxy":
|
||||
"""Retrieve the global LoggerProxy.
|
||||
|
||||
This should be used for all logging calls. It returns a proxy that
|
||||
delegates to the active logger (defaulting to a SimpleLogger until
|
||||
init_logger is called).
|
||||
"""
|
||||
global _proxy_instance, _active_logger
|
||||
if _proxy_instance is None:
|
||||
if _active_logger is None:
|
||||
# Fallback to SimpleLogger to avoid premature directory creation
|
||||
from experiment_logger.simple_logger import SimpleLogger
|
||||
|
||||
_active_logger = SimpleLogger(run_name="pre_init")
|
||||
|
||||
_proxy_instance = LoggerProxy()
|
||||
|
||||
return _proxy_instance
|
||||
|
||||
|
||||
def init_logger(**kwargs) -> "UnifiedLogger":
|
||||
"""Initialize the full UnifiedLogger and set it as the active logger.
|
||||
|
||||
This should be called once the configuration is ready. It will create
|
||||
the output directories and set up all logging backends.
|
||||
"""
|
||||
global _active_logger
|
||||
logger = UnifiedLogger(**kwargs)
|
||||
_active_logger = logger
|
||||
return logger
|
||||
|
||||
|
||||
class LoggerProxy:
|
||||
"""Proxy that delegates all method calls to the active logger instance.
|
||||
|
||||
This allows the logger to be swapped out (e.g., from a SimpleLogger to
|
||||
a UnifiedLogger) without any clients needing to update their references.
|
||||
"""
|
||||
|
||||
def _get_logger(self) -> Any:
|
||||
global _active_logger
|
||||
if _active_logger is None:
|
||||
# This shouldn't normally happen since get_logger handles it
|
||||
from experiment_logger.simple_logger import SimpleLogger
|
||||
|
||||
_active_logger = SimpleLogger(run_name="pre_init_fallback")
|
||||
return _active_logger
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._get_logger(), name)
|
||||
|
||||
def __enter__(self):
|
||||
return self._get_logger().__enter__()
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
return self._get_logger().__exit__(exc_type, exc_val, exc_tb)
|
||||
|
||||
|
||||
class UnifiedLogger:
|
||||
"""Unified logger for scientific experiments with redundant backup."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
run_name: str,
|
||||
full_config: Dict[str, Any],
|
||||
logging_cfg: LoggingConfig,
|
||||
base_dir: str = "runs",
|
||||
save_code: bool = True,
|
||||
log_level: int = logging.INFO,
|
||||
):
|
||||
"""Initialize the unified logger.
|
||||
|
||||
Args:
|
||||
run_name: Unique name for this run
|
||||
full_config: Full configuration dictionary with hyperparameters to be saved
|
||||
logging_cfg: Structured logging configuration dataclass
|
||||
base_dir: Base directory for local storage
|
||||
save_code: Whether to save code to WandB
|
||||
"""
|
||||
self.run_name = run_name
|
||||
self.full_config = full_config
|
||||
self.use_wandb = logging_cfg.track
|
||||
self.upload_final_model = logging_cfg.upload_final_model
|
||||
self.upload_checkpoints = logging_cfg.upload_checkpoints
|
||||
self.wandb_available = False
|
||||
self.wandb_run = None
|
||||
self.is_interactive = sys.stdout.isatty()
|
||||
|
||||
# Setup local storage
|
||||
self.run_dir = Path(base_dir) / run_name
|
||||
self.run_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self.checkpoints_dir = self.run_dir / "checkpoints"
|
||||
self.checkpoints_dir.mkdir(exist_ok=True)
|
||||
|
||||
self.metrics_dir = self.run_dir / "metrics"
|
||||
self.metrics_dir.mkdir(exist_ok=True)
|
||||
|
||||
self.config_file = self.run_dir / "config.yaml"
|
||||
|
||||
# Setup standard Python logging mirror
|
||||
self.text_log_file = self.run_dir / "run.log"
|
||||
self._text_logger = logging.getLogger(f"UnifiedLogger_{self.run_name}")
|
||||
self._text_logger.setLevel(log_level)
|
||||
self._text_logger.propagate = False
|
||||
|
||||
# Avoid duplicate handlers if re-instantiated
|
||||
if not self._text_logger.handlers:
|
||||
fh = logging.FileHandler(self.text_log_file)
|
||||
ch = logging.StreamHandler()
|
||||
|
||||
formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
|
||||
fh.setFormatter(formatter)
|
||||
ch.setFormatter(formatter)
|
||||
|
||||
self._text_logger.addHandler(fh)
|
||||
self._text_logger.addHandler(ch)
|
||||
|
||||
# Save config to disk
|
||||
self._save_config()
|
||||
|
||||
# Setup TensorBoard
|
||||
self.writer = None
|
||||
try:
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
|
||||
self.writer = SummaryWriter(self.run_dir)
|
||||
self.info("TensorBoard SummaryWriter initialized.")
|
||||
except ImportError:
|
||||
self.warning("tensorboard not installed. Skipping SummaryWriter.")
|
||||
|
||||
# Initialize WandB if requested
|
||||
if self.use_wandb:
|
||||
self._init_wandb(logging_cfg.wandb_project_name, logging_cfg.wandb_entity, save_code)
|
||||
|
||||
# Initialize metrics storage
|
||||
self.metrics_buffer: List[Dict[str, Any]] = []
|
||||
self.step_counter = 0
|
||||
|
||||
self.info(f"Initialized UnifiedLogger for run: {run_name}")
|
||||
self.info(f"Local storage: {self.run_dir.absolute()}")
|
||||
self.info(f"WandB logging: {self.wandb_available}")
|
||||
|
||||
def set_level(self, level: int):
|
||||
"""Dynamically update the verbosity of the stdout/text logger."""
|
||||
self._text_logger.setLevel(level)
|
||||
|
||||
def log_non_interactive(self, msg: str, *args, **kwargs):
|
||||
"""Log an info message only if running in a non-interactive environment."""
|
||||
if not self.is_interactive:
|
||||
self.info(msg, *args, **kwargs)
|
||||
|
||||
def progress_bar(self, iterable=None, *args, **kwargs):
|
||||
"""Wrapper around tqdm that automatically disables in non-interactive environments."""
|
||||
import tqdm
|
||||
|
||||
kwargs.setdefault("disable", not self.is_interactive)
|
||||
return tqdm.tqdm(iterable, *args, **kwargs)
|
||||
|
||||
def info(self, msg: str, *args, **kwargs):
|
||||
"""Log an info message to stdout and disk."""
|
||||
self._text_logger.info(msg, *args, **kwargs)
|
||||
|
||||
def warning(self, msg: str, *args, **kwargs):
|
||||
"""Log a warning message to stdout and disk."""
|
||||
self._text_logger.warning(msg, *args, **kwargs)
|
||||
|
||||
def error(self, msg: str, *args, **kwargs):
|
||||
"""Log an error message to stdout and disk."""
|
||||
self._text_logger.error(msg, *args, **kwargs)
|
||||
|
||||
def debug(self, msg: str, *args, **kwargs):
|
||||
"""Log a debug message to stdout and disk."""
|
||||
self._text_logger.debug(msg, *args, **kwargs)
|
||||
|
||||
def _init_wandb(self, project_name: str, entity: Optional[str], save_code: bool):
|
||||
"""Initialize Weights & Biases logging."""
|
||||
self.wandb_run = init_wandb(
|
||||
project=project_name,
|
||||
entity=entity,
|
||||
name=self.run_name,
|
||||
config=self.full_config,
|
||||
save_code=save_code,
|
||||
resume="allow",
|
||||
)
|
||||
self.wandb_available = self.wandb_run is not None
|
||||
|
||||
def _save_config(self):
|
||||
"""Save configuration to disk."""
|
||||
try:
|
||||
with open(self.config_file, "w") as f:
|
||||
yaml.safe_dump(
|
||||
_sanitize_for_yaml(self.full_config),
|
||||
f,
|
||||
default_flow_style=False,
|
||||
indent=2,
|
||||
sort_keys=False,
|
||||
)
|
||||
self.info(f"Config saved to {self.config_file}")
|
||||
except Exception as e:
|
||||
self.error(f"Error saving config: {e}")
|
||||
|
||||
def log(self, metrics: Dict[str, Any], step: Optional[int] = None, commit: bool = True):
|
||||
"""Log metrics to all backends.
|
||||
|
||||
Args:
|
||||
metrics: Dictionary of metric name -> value
|
||||
step: Global step counter (auto-incremented if None)
|
||||
commit: Whether to commit to WandB immediately
|
||||
"""
|
||||
if step is None:
|
||||
step = self.step_counter
|
||||
self.step_counter += 1
|
||||
|
||||
# Add timestamp
|
||||
metrics_with_metadata = {
|
||||
"step": step,
|
||||
"timestamp": time.time(),
|
||||
**metrics,
|
||||
}
|
||||
|
||||
# Log to stdout
|
||||
self._log_to_stdout(metrics_with_metadata)
|
||||
|
||||
# Log to WandB
|
||||
if self.wandb_run is not None:
|
||||
try:
|
||||
self.wandb_run.log(metrics, step=step, commit=commit)
|
||||
except Exception as e:
|
||||
self.warning(f"WandB logging failed: {e}")
|
||||
|
||||
# Log to TensorBoard
|
||||
if self.writer is not None:
|
||||
for k, v in metrics.items():
|
||||
if isinstance(v, (int, float, np.floating, np.integer)):
|
||||
self.writer.add_scalar(k, v, step)
|
||||
elif hasattr(v, "item"):
|
||||
self.writer.add_scalar(k, v.item(), step)
|
||||
elif isinstance(v, (np.ndarray, jnp.ndarray)) and v.size == 1:
|
||||
self.writer.add_scalar(k, v.item(), step)
|
||||
|
||||
# Buffer for disk storage
|
||||
self.metrics_buffer.append(metrics_with_metadata)
|
||||
|
||||
# Periodically flush to disk
|
||||
if len(self.metrics_buffer) >= 100:
|
||||
self._flush_metrics()
|
||||
|
||||
def _log_to_stdout(self, metrics: Dict[str, Any]):
|
||||
"""Log metrics to stdout for real-time monitoring."""
|
||||
step = metrics.get("step", "?")
|
||||
metric_str = ", ".join(
|
||||
f"{k}={v:.6f}" if isinstance(v, (float, np.floating)) else f"{k}={v}"
|
||||
for k, v in metrics.items()
|
||||
if k not in ["step", "timestamp"]
|
||||
)
|
||||
self.info(f"[Step {step}] {metric_str}")
|
||||
|
||||
def _flush_metrics(self):
|
||||
"""Flush buffered metrics to disk."""
|
||||
if not self.metrics_buffer:
|
||||
return
|
||||
|
||||
try:
|
||||
metrics_file = self.metrics_dir / "metrics.yaml"
|
||||
with open(metrics_file, "a") as f:
|
||||
for metric in self.metrics_buffer:
|
||||
# Convert numpy/jax types to native Python types for YAML serialization
|
||||
serializable_metric = {}
|
||||
for k, v in metric.items():
|
||||
if hasattr(v, "item"): # numpy/jax scalar
|
||||
serializable_metric[k] = v.item()
|
||||
elif isinstance(v, (np.ndarray, jnp.ndarray)):
|
||||
serializable_metric[k] = v.tolist()
|
||||
else:
|
||||
serializable_metric[k] = v
|
||||
f.write("---\n")
|
||||
yaml.safe_dump(
|
||||
_sanitize_for_yaml(serializable_metric),
|
||||
f,
|
||||
default_flow_style=False,
|
||||
sort_keys=False,
|
||||
)
|
||||
self.metrics_buffer.clear()
|
||||
except Exception as e:
|
||||
self.error(f"Error flushing metrics: {e}")
|
||||
|
||||
def save_checkpoint(
|
||||
self,
|
||||
params: Any,
|
||||
step: int,
|
||||
prefix: str = "checkpoint",
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
"""Save model checkpoint to disk and optionally to WandB."""
|
||||
checkpoint_name = f"{prefix}_step_{step}.flax"
|
||||
checkpoint_path = self.checkpoints_dir / checkpoint_name
|
||||
|
||||
try:
|
||||
# Save to disk using Flax serialization
|
||||
with open(checkpoint_path, "wb") as f:
|
||||
f.write(flax.serialization.to_bytes(params))
|
||||
|
||||
# Save metadata if provided
|
||||
if metadata:
|
||||
metadata_path = self.checkpoints_dir / f"{prefix}_step_{step}_metadata.yaml"
|
||||
with open(metadata_path, "w") as f:
|
||||
yaml.safe_dump(
|
||||
_sanitize_for_yaml(metadata),
|
||||
f,
|
||||
default_flow_style=False,
|
||||
indent=2,
|
||||
sort_keys=False,
|
||||
)
|
||||
|
||||
self.info(f"Checkpoint saved: {checkpoint_path}")
|
||||
|
||||
# Log to WandB as artifact
|
||||
if self.wandb_run is not None and self.upload_checkpoints:
|
||||
try:
|
||||
import wandb
|
||||
|
||||
artifact = wandb.Artifact(
|
||||
name=f"{self.run_name}_{prefix}",
|
||||
type="model",
|
||||
metadata=metadata or {},
|
||||
)
|
||||
artifact.add_file(str(checkpoint_path))
|
||||
if metadata:
|
||||
artifact.add_file(str(metadata_path))
|
||||
self.wandb_run.log_artifact(artifact)
|
||||
self.info("Checkpoint uploaded to WandB")
|
||||
except Exception as e:
|
||||
self.warning(f"Could not upload checkpoint to WandB: {e}")
|
||||
|
||||
except Exception as e:
|
||||
self.error(f"Error saving checkpoint: {e}")
|
||||
|
||||
def save_final_model(self, params: Any, metadata: Optional[Dict[str, Any]] = None):
|
||||
"""Save the final trained model."""
|
||||
final_model_path = self.run_dir / "final_model.flax"
|
||||
|
||||
try:
|
||||
with open(final_model_path, "wb") as f:
|
||||
f.write(flax.serialization.to_bytes(params))
|
||||
|
||||
if metadata:
|
||||
metadata_path = self.run_dir / "final_model_metadata.yaml"
|
||||
with open(metadata_path, "w") as f:
|
||||
yaml.safe_dump(
|
||||
_sanitize_for_yaml(metadata),
|
||||
f,
|
||||
default_flow_style=False,
|
||||
indent=2,
|
||||
sort_keys=False,
|
||||
)
|
||||
|
||||
self.info(f"Final model saved: {final_model_path}")
|
||||
|
||||
# Log to WandB
|
||||
if self.wandb_run is not None and self.upload_final_model:
|
||||
try:
|
||||
import wandb
|
||||
|
||||
artifact = wandb.Artifact(
|
||||
name=f"{self.run_name}_final_model",
|
||||
type="model",
|
||||
metadata=metadata or {},
|
||||
)
|
||||
artifact.add_file(str(final_model_path))
|
||||
if metadata:
|
||||
artifact.add_file(str(metadata_path))
|
||||
self.wandb_run.log_artifact(artifact)
|
||||
except Exception as e:
|
||||
self.warning(f"Could not upload final model to WandB: {e}")
|
||||
|
||||
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
|
||||
self._flush_metrics()
|
||||
|
||||
if self.writer is not None:
|
||||
self.writer.close()
|
||||
|
||||
self.info(f"Run complete. Results saved to: {self.run_dir.absolute()}")
|
||||
|
||||
# Finish WandB run
|
||||
if self.wandb_available:
|
||||
finish_wandb()
|
||||
|
||||
def __enter__(self):
|
||||
"""Context manager entry."""
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
"""Context manager exit."""
|
||||
self.finish()
|
||||
91
src/experiment_logger/wandb_utils.py
Normal file
91
src/experiment_logger/wandb_utils.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Centralized WandB initialization utilities."""
|
||||
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def init_wandb(
|
||||
project: str,
|
||||
config: Dict[str, Any],
|
||||
name: Optional[str] = None,
|
||||
entity: Optional[str] = None,
|
||||
sync_tensorboard: bool = False,
|
||||
save_code: bool = True,
|
||||
resume: str = "allow",
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize WandB with standardized settings.
|
||||
|
||||
This function provides a centralized way to initialize WandB across different
|
||||
scripts, ensuring consistent configuration and error handling.
|
||||
|
||||
Args:
|
||||
project: WandB project name
|
||||
config: Configuration dictionary to log
|
||||
name: Run name (auto-generated if None)
|
||||
entity: WandB entity (team/user name)
|
||||
sync_tensorboard: Whether to sync tensorboard logs
|
||||
save_code: Whether to save code snapshots
|
||||
resume: Resume strategy ("allow", "must", "never", "auto")
|
||||
**kwargs: Additional arguments to pass to wandb.init()
|
||||
|
||||
Returns:
|
||||
wandb.Run object if successful, None otherwise
|
||||
"""
|
||||
try:
|
||||
import wandb
|
||||
import os
|
||||
import sys
|
||||
|
||||
# Robust HPC checking: check for API key
|
||||
has_key = os.environ.get("WANDB_API_KEY") is not None
|
||||
if not has_key:
|
||||
try:
|
||||
# Check if logged in locally via settings/netrc
|
||||
has_key = wandb.setup().settings.api_key is not None
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
is_interactive = sys.stdout.isatty()
|
||||
|
||||
if not has_key and not is_interactive and os.environ.get("WANDB_MODE") != "offline":
|
||||
logger.warning(
|
||||
"WANDB_API_KEY not found and environment is non-interactive. "
|
||||
"Switching to offline mode."
|
||||
)
|
||||
sync_path = f"runs/{name}" if name else "runs"
|
||||
logger.warning(f"WandB is offline. Use 'wandb sync {sync_path}' to upload logs later.")
|
||||
os.environ["WANDB_MODE"] = "offline"
|
||||
|
||||
run = wandb.init(
|
||||
project=project,
|
||||
entity=entity,
|
||||
name=name,
|
||||
config=config,
|
||||
sync_tensorboard=sync_tensorboard,
|
||||
save_code=save_code,
|
||||
resume=resume,
|
||||
**kwargs,
|
||||
)
|
||||
logger.info(f"WandB initialized successfully for project '{project}', run '{run.name}'")
|
||||
return run
|
||||
except ImportError:
|
||||
logger.warning("WandB not installed. Skipping WandB initialization.")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to initialize WandB: {e}")
|
||||
return None
|
||||
|
||||
|
||||
def finish_wandb():
|
||||
"""Safely finish the current WandB run."""
|
||||
try:
|
||||
import wandb
|
||||
|
||||
if wandb.run is not None:
|
||||
wandb.finish()
|
||||
logger.info("WandB run finished successfully")
|
||||
except Exception as e:
|
||||
logger.warning(f"Error finishing WandB run: {e}")
|
||||
Reference in a new issue