restructure
This commit is contained in:
parent
b892e3777e
commit
431ecdf4b4
10 changed files with 8 additions and 27 deletions
|
|
@ -28,7 +28,7 @@ jobs:
|
||||||
uses: astral-sh/setup-uv@v5
|
uses: astral-sh/setup-uv@v5
|
||||||
|
|
||||||
- name: Regenerate env/hpc/requirements.txt
|
- name: Regenerate env/hpc/requirements.txt
|
||||||
run: uv run scripts/export_hpc_requirements.py
|
run: uv run scripts/hpc/export_requirements.py
|
||||||
|
|
||||||
- name: Commit updated requirements if changed
|
- name: Commit updated requirements if changed
|
||||||
uses: stefanzweifel/git-auto-commit-action@v5
|
uses: stefanzweifel/git-auto-commit-action@v5
|
||||||
|
|
|
||||||
|
|
@ -79,7 +79,7 @@ After installation, run these commands to ensure your environment is set up corr
|
||||||
`env/hpc/requirements.txt` is auto-generated from `pyproject.toml`. To regenerate:
|
`env/hpc/requirements.txt` is auto-generated from `pyproject.toml`. To regenerate:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv run scripts/export_hpc_requirements.py
|
uv run scripts/hpc/export_requirements.py
|
||||||
```
|
```
|
||||||
|
|
||||||
Modules listed in `env/hpc/modules.txt` are automatically excluded from the pip requirements to save space and use HPC-optimized binaries.
|
Modules listed in `env/hpc/modules.txt` are automatically excluded from the pip requirements to save space and use HPC-optimized binaries.
|
||||||
|
|
|
||||||
|
|
@ -1 +0,0 @@
|
||||||
# TODO
|
|
||||||
|
|
@ -1,3 +0,0 @@
|
||||||
from .plot import simple_plot
|
|
||||||
|
|
||||||
__all__ = ["simple_plot"]
|
|
||||||
|
|
@ -1,12 +0,0 @@
|
||||||
import matplotlib.pyplot as plt
|
|
||||||
|
|
||||||
|
|
||||||
def simple_plot(x: list, y: list, show_window: bool = False, filename: str = "plot.png") -> None:
|
|
||||||
plt.plot(x, y)
|
|
||||||
plt.savefig(filename)
|
|
||||||
|
|
||||||
if show_window:
|
|
||||||
# blocks until window is closed
|
|
||||||
plt.show()
|
|
||||||
|
|
||||||
plt.close()
|
|
||||||
|
|
@ -5,9 +5,6 @@ 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
|
code whenever pyproject.toml dependencies change. It reads the modules from
|
||||||
env/hpc/modules.txt and the full dependency list from pyproject.toml, then
|
env/hpc/modules.txt and the full dependency list from pyproject.toml, then
|
||||||
writes the remainder to env/hpc/requirements.txt.
|
writes the remainder to env/hpc/requirements.txt.
|
||||||
|
|
||||||
Usage:
|
|
||||||
uv run scripts/export_hpc_requirements.py
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
@ -7,7 +7,7 @@ import yaml
|
||||||
import os
|
import os
|
||||||
|
|
||||||
from brittle_star_project.dataclasses import PPOArgs
|
from brittle_star_project.dataclasses import PPOArgs
|
||||||
from PPOTrainer import PPOTrainer
|
from brittle_star_project.trainers.PPOTrainer import PPOTrainer
|
||||||
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -28,13 +28,13 @@ from ppo import PPO
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
|
def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
|
||||||
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
|
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
|
||||||
return learning_rate * frac
|
return learning_rate * frac
|
||||||
|
|
||||||
|
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
|
def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
|
||||||
return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))(
|
return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))(
|
||||||
obs_dict
|
obs_dict
|
||||||
)
|
)
|
||||||
|
|
@ -124,7 +124,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn):
|
||||||
return (
|
return (
|
||||||
episode_stats,
|
episode_stats,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
(convert_obs_dict_to_array(next_env_state.observations), reward, done),
|
(_convert_obs_dict_to_array(next_env_state.observations), reward, done),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -300,7 +300,7 @@ class PPOTrainer:
|
||||||
optax.clip_by_global_norm(self.args.max_grad_norm),
|
optax.clip_by_global_norm(self.args.max_grad_norm),
|
||||||
optax.inject_hyperparams(optax.adam)(
|
optax.inject_hyperparams(optax.adam)(
|
||||||
learning_rate=partial(
|
learning_rate=partial(
|
||||||
linear_schedule,
|
_linear_schedule,
|
||||||
minibatch_count=self.args.num_minibatches,
|
minibatch_count=self.args.num_minibatches,
|
||||||
update_epochs=self.args.update_epochs,
|
update_epochs=self.args.update_epochs,
|
||||||
num_iterations=self.args.num_iterations,
|
num_iterations=self.args.num_iterations,
|
||||||
|
|
@ -467,7 +467,7 @@ class PPOTrainer:
|
||||||
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
|
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
env_state = self.env.reset(seed=self.args.seed)
|
env_state = self.env.reset(seed=self.args.seed)
|
||||||
next_obs = convert_obs_dict_to_array(env_state.observations)
|
next_obs = _convert_obs_dict_to_array(env_state.observations)
|
||||||
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
|
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
|
||||||
|
|
||||||
if log and not is_tty:
|
if log and not is_tty:
|
||||||
0
src/brittle_star_project/trainers/__init__.py
Normal file
0
src/brittle_star_project/trainers/__init__.py
Normal file
Reference in a new issue