create-rl-app 0.0.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,63 @@
1
+ Metadata-Version: 2.3
2
+ Name: create-rl-app
3
+ Version: 0.0.1
4
+ Summary: Add your description here
5
+ Author: Koen Ponse
6
+ Author-email: Koen Ponse <k.ponse@liacs.leidenuniv.nl>
7
+ Requires-Python: >=3.11
8
+ Description-Content-Type: text/markdown
9
+
10
+ # create-rl-app
11
+
12
+ [![Python Version](https://img.shields.io/badge/python-3.11%2B-blue.svg)](https://python.org)
13
+ [![Jaxnasium Version](https://badge.fury.io/py/jaxnasium.svg)](https://github.com/ponseko/jaxnasium)
14
+ [![Ruff](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json)](https://github.com/astral-sh/ruff)
15
+
16
+ A CLI application to bootstrap reinforcement learning applications within the Jaxnasium ecosystem. Quickly scaffold new RL projects for either developing environments with baseline algorithms or for altering existing baselines.
17
+
18
+ ## What it does
19
+
20
+ `create-rl-app` is a command-line tool that helps you quickly set up new reinforcement learning projects using the Jaxnasium framework. It creates a well-structured project template with:
21
+
22
+ - 🚀 **Quick Setup**: Get a new RL project running in seconds
23
+ - 🏗️ **Helpful Templates**: Templates for environments and algorithms for you to start with.
24
+ - ⚡ **Performance Optimized**: Sets you up with PureJaxRL compatible agents and environments for performance and GPU scalability.
25
+
26
+ ## Useage
27
+
28
+ ### uvx (Recommended)
29
+
30
+ ```bash
31
+ uvx create-rl-app <project_name>
32
+ cd <project_name>
33
+ uv run train_example.py
34
+ ```
35
+
36
+ ### pipx
37
+
38
+ ```bash
39
+ pipx run create-rl-app <project_name>
40
+ cd <project_name>
41
+ # Create a new environment (e.g. conda, venv, etc.)
42
+ # source .../bin/activate
43
+ python train_example.py
44
+ ```
45
+
46
+ ### Or Install Globally
47
+
48
+ ```bash
49
+ uv tool install create-rl-app
50
+ ```
51
+
52
+ ```bash
53
+ pip install create-rl-app
54
+ ```
55
+
56
+ ## Dependencies
57
+
58
+ None.
59
+
60
+ ## Contributing
61
+
62
+ Contributions are welcome! Please feel free to submit a Pull Request.
63
+
@@ -0,0 +1,54 @@
1
+ # create-rl-app
2
+
3
+ [![Python Version](https://img.shields.io/badge/python-3.11%2B-blue.svg)](https://python.org)
4
+ [![Jaxnasium Version](https://badge.fury.io/py/jaxnasium.svg)](https://github.com/ponseko/jaxnasium)
5
+ [![Ruff](https://img.shields.io/endpoint?url=https://raw.githubusercontent.com/astral-sh/ruff/main/assets/badge/v2.json)](https://github.com/astral-sh/ruff)
6
+
7
+ A CLI application to bootstrap reinforcement learning applications within the Jaxnasium ecosystem. Quickly scaffold new RL projects for either developing environments with baseline algorithms or for altering existing baselines.
8
+
9
+ ## What it does
10
+
11
+ `create-rl-app` is a command-line tool that helps you quickly set up new reinforcement learning projects using the Jaxnasium framework. It creates a well-structured project template with:
12
+
13
+ - 🚀 **Quick Setup**: Get a new RL project running in seconds
14
+ - 🏗️ **Helpful Templates**: Templates for environments and algorithms for you to start with.
15
+ - ⚡ **Performance Optimized**: Sets you up with PureJaxRL compatible agents and environments for performance and GPU scalability.
16
+
17
+ ## Useage
18
+
19
+ ### uvx (Recommended)
20
+
21
+ ```bash
22
+ uvx create-rl-app <project_name>
23
+ cd <project_name>
24
+ uv run train_example.py
25
+ ```
26
+
27
+ ### pipx
28
+
29
+ ```bash
30
+ pipx run create-rl-app <project_name>
31
+ cd <project_name>
32
+ # Create a new environment (e.g. conda, venv, etc.)
33
+ # source .../bin/activate
34
+ python train_example.py
35
+ ```
36
+
37
+ ### Or Install Globally
38
+
39
+ ```bash
40
+ uv tool install create-rl-app
41
+ ```
42
+
43
+ ```bash
44
+ pip install create-rl-app
45
+ ```
46
+
47
+ ## Dependencies
48
+
49
+ None.
50
+
51
+ ## Contributing
52
+
53
+ Contributions are welcome! Please feel free to submit a Pull Request.
54
+
@@ -0,0 +1,17 @@
1
+ [project]
2
+ name = "create-rl-app"
3
+ version = "0.0.1"
4
+ description = "Add your description here"
5
+ readme = "README.md"
6
+ authors = [
7
+ { name = "Koen Ponse", email = "k.ponse@liacs.leidenuniv.nl" }
8
+ ]
9
+ requires-python = ">=3.11"
10
+ dependencies = []
11
+
12
+ [project.scripts]
13
+ create-rl-app = "create_rl_app.cli:main"
14
+
15
+ [build-system]
16
+ requires = ["uv_build>=0.8.4,<0.9.0"]
17
+ build-backend = "uv_build"
File without changes
@@ -0,0 +1,2 @@
1
+ # This file ensures the vendored directory is tracked by git
2
+ # even when empty. It will be removed when files are vendored.
@@ -0,0 +1,317 @@
1
+ import argparse
2
+ import importlib.resources as pkg_resources
3
+ import json
4
+ import re
5
+ import shutil
6
+ import subprocess
7
+ from pathlib import Path
8
+
9
+ # ANSI color codes
10
+ CYAN = "\033[96m"
11
+ GREEN = "\033[92m"
12
+ YELLOW = "\033[93m"
13
+ BOLD = "\033[1m"
14
+ RESET = "\033[0m"
15
+ LIGHT_GRAY = "\033[37m"
16
+
17
+
18
+ def get_vendored_jaxnasium_version():
19
+ """Get the vendored jaxnasium version from .vendor_info"""
20
+ try:
21
+ files = pkg_resources.files("create_rl_app._vendored")
22
+ vendor_info = files / "jaxnasium" / ".vendor_info"
23
+ with vendor_info.open("r") as f:
24
+ data = json.load(f)
25
+ version = data["version"]
26
+ return version.lstrip("v") if version.startswith("v") else version
27
+ except (FileNotFoundError, ImportError, KeyError):
28
+ raise RuntimeError("Vendored jaxnasium version not found")
29
+
30
+
31
+ def colored_input(prompt, default=""):
32
+ """Display a colored input prompt with default value."""
33
+ default_display = f" ({LIGHT_GRAY}{default}{RESET})" if default else ""
34
+ user_input = input(
35
+ f"{CYAN}{BOLD}?{RESET} {prompt}{YELLOW}{default_display}{RESET}: "
36
+ ).strip()
37
+ return user_input if user_input else default
38
+
39
+
40
+ def yes_no_prompt(question, default="y"):
41
+ """Ask a yes/no question with colored output."""
42
+ options = "(Y/n)" if default.lower() == "y" else "(y/N)"
43
+ while True:
44
+ response = colored_input(f"{question} {options}")
45
+ if not response:
46
+ return default.lower() == "y"
47
+ if response.lower() in ["y", "yes"]:
48
+ return True
49
+ elif response.lower() in ["n", "no"]:
50
+ return False
51
+ print(f"{YELLOW}Please answer with 'y' or 'n'.{RESET}")
52
+
53
+
54
+ def update_pyproject_toml_file(project_path):
55
+ toml_path = project_path / "pyproject.toml"
56
+ if toml_path.exists():
57
+ with toml_path.open("r") as file:
58
+ content = file.read()
59
+
60
+ # Update requires-python version
61
+ content = re.sub(
62
+ r"^requires-python\s*=\s*.*$",
63
+ 'requires-python = ">=3.11"',
64
+ content,
65
+ flags=re.MULTILINE,
66
+ )
67
+
68
+ # Update dependencies
69
+ def update_dependencies(match):
70
+ deps_str = match.group(1).strip()
71
+ if deps_str == "[]":
72
+ dependencies = []
73
+ else:
74
+ dependencies = re.findall(r'"([^"]+)"', deps_str)
75
+
76
+ # Add jaxnasium with the required version
77
+ dependencies.append(f"jaxnasium[algs]>={get_vendored_jaxnasium_version()}")
78
+ return f"dependencies = {dependencies}"
79
+
80
+ content = re.sub(
81
+ r"^dependencies\s*=\s*(\[.*?\])",
82
+ update_dependencies,
83
+ content,
84
+ flags=re.MULTILINE | re.DOTALL,
85
+ )
86
+
87
+ # UV-build expects a src folder by default. Remove this expectation here.
88
+ if re.search(r'build-backend\s*=\s*["\']uv_build["\']', content):
89
+ if not re.search(r"\[tool\.uv\.build-backend\]", content):
90
+ content += "\n\n[tool.uv.build-backend]\nmodule-root = ''\n"
91
+
92
+ with toml_path.open("w") as file:
93
+ file.write(content)
94
+
95
+
96
+ def replace_init_file(src_path, import_env):
97
+ """Clear the project __init__.py file (or create one if it does not exist)
98
+ Fill it with "from projectname import ExampleEnv".
99
+ """
100
+ init_file_path = src_path / "__init__.py"
101
+ with init_file_path.open("w") as file:
102
+ file.write("# This file is auto-generated by jaxnasium\n")
103
+ if import_env:
104
+ file.write(f"from .{src_path.name}Env import ExampleEnv\n")
105
+
106
+
107
+ def copy_env_template(target_path: Path):
108
+ target_path.parent.mkdir(parents=True, exist_ok=True)
109
+ with pkg_resources.path(
110
+ "create_rl_app.resources", "env_template.py"
111
+ ) as template_path:
112
+ shutil.copy(template_path, target_path)
113
+
114
+
115
+ def copy_algorithms_source(target_path: Path):
116
+ target_path.mkdir(parents=True, exist_ok=True)
117
+ with pkg_resources.path(
118
+ "create_rl_app._vendored.jaxnasium", "algorithms"
119
+ ) as template_path:
120
+ # ignore the "utils" folder in "algorithms"
121
+ ignored = ["utils", "_algorithm.py"]
122
+ shutil.copytree(
123
+ template_path,
124
+ target_path,
125
+ dirs_exist_ok=True,
126
+ ignore=shutil.ignore_patterns(*ignored),
127
+ )
128
+
129
+
130
+ def main():
131
+ # Parse the project name argument
132
+ parser = argparse.ArgumentParser(description="Initialize a new jaxnasium project.")
133
+ parser.add_argument("projectname", help="The path to the new project directory.")
134
+ parser.add_argument(
135
+ "-y",
136
+ "--yes",
137
+ action="store_true",
138
+ help="Automatically answer 'yes' to all prompts.",
139
+ )
140
+ parser.add_argument(
141
+ "--env-template",
142
+ type=str,
143
+ choices=["true", "false"],
144
+ help="Include environment template in the project (true/false).",
145
+ )
146
+ parser.add_argument(
147
+ "--algorithm-source",
148
+ type=str,
149
+ choices=["true", "false"],
150
+ help="Copy algorithm source code into the project instead of importing (true/false).",
151
+ )
152
+ parser.add_argument(
153
+ "--init-algorithms",
154
+ action="store_true",
155
+ help="Skip project creation and only copy algorithm source code at the <projectname> location",
156
+ )
157
+ parser.add_argument(
158
+ "--init-env",
159
+ action="store_true",
160
+ help="Skip project creation and only copy environment source code at the <projectname> location",
161
+ )
162
+ args = parser.parse_args()
163
+ projectname = args.projectname
164
+
165
+ if args.init_algorithms:
166
+ copy_algorithms_source(Path(projectname))
167
+
168
+ if args.init_env:
169
+ copy_env_template(Path(projectname))
170
+
171
+ if args.init_algorithms or args.init_env:
172
+ return
173
+
174
+ print(f"{CYAN}{BOLD}")
175
+ print(" ██████╗██████╗ ███████╗ █████╗ ████████╗███████╗")
176
+ print(" ██╔════╝██╔══██╗██╔════╝██╔══██╗╚══██╔══╝██╔════╝")
177
+ print(" ██║ ██████╔╝█████╗ ███████║ ██║ █████╗")
178
+ print(" ██║ ██╔══██╗██╔══╝ ██╔══██║ ██║ ██╔══╝")
179
+ print(" ╚██████╗██║ ██║███████╗██║ ██║ ██║ ███████╗")
180
+ print(" ╚═════╝╚═╝ ╚═╝╚══════╝╚═╝ ╚═╝ ╚═╝ ╚══════╝")
181
+ print("")
182
+ print(" ██████╗ ██╗ █████╗ ██████╗ ██████╗")
183
+ print(" ██╔══██╗██║ ██╔══██╗██╔══██╗██╔══██╗")
184
+ print(" ██████╔╝██║ ███████║██████╔╝██████╔╝")
185
+ print(" ██╔══██╗██║ ██╔══██║██╔═══╝ ██╔═══╝")
186
+ print(" ██║ ██║███████╗ ██║ ██║██║ ██║")
187
+ print(" ╚═╝ ╚═╝╚══════╝ ╚═╝ ╚═╝╚═╝ ╚═╝")
188
+ print(f"{RESET}")
189
+
190
+ print(
191
+ f"{LIGHT_GRAY}Setting up a new Jaxnasium project (v{get_vendored_jaxnasium_version()}){RESET}"
192
+ )
193
+ print(f"{LIGHT_GRAY}{'─' * 60}{RESET}\n")
194
+
195
+ if any(char.isupper() for char in projectname):
196
+ print(f"{YELLOW} Project name has been altered to lowercase{RESET}")
197
+ projectname = projectname.lower()
198
+ print(f"{LIGHT_GRAY}Project name: {projectname}{RESET}")
199
+
200
+ ########## Questions ##########
201
+
202
+ # Use CLI arguments if provided, otherwise prompt the user
203
+ if args.env_template is not None:
204
+ build_environment_template = args.env_template == "true"
205
+ else:
206
+ build_environment_template = args.yes or yes_no_prompt(
207
+ "Would you like to include a environment template?", default="y"
208
+ )
209
+
210
+ if args.algorithm_source is not None:
211
+ include_algorithm_source = args.algorithm_source == "true"
212
+ else:
213
+ include_algorithm_source = args.yes or yes_no_prompt(
214
+ "Instead of importing, would you like to copy the algorithm source code into your project?",
215
+ default="n",
216
+ )
217
+
218
+ # Display summary of choices
219
+ print(f"\n{BOLD}📋 Project configuration summary:{RESET}")
220
+ print(f" • Project name: {projectname}")
221
+ print(
222
+ f" • Include environment template: {'Yes' if build_environment_template else 'No'}"
223
+ )
224
+ print(
225
+ f" • Copy algorithm source code: {'Yes' if include_algorithm_source else 'No'}"
226
+ )
227
+
228
+ # Confirm setup
229
+ if not args.yes:
230
+ if not yes_no_prompt("\nDo you want to proceed with this configuration?"):
231
+ print("Setup cancelled.")
232
+ return
233
+
234
+ #### SETUP PROJECT ####
235
+
236
+ # Determine the command to run
237
+ command = (
238
+ ["uv", "init", "--package", projectname]
239
+ if shutil.which("uv")
240
+ else ["pipx", "run", "uv", "init", "--package", projectname]
241
+ )
242
+
243
+ # Run UV init command
244
+ try:
245
+ subprocess.run(command, check=True)
246
+ except subprocess.CalledProcessError as e:
247
+ raise RuntimeError(
248
+ f"Failed to initialize the project: {e}. \n Most likely neither uv nor pipx are installed."
249
+ )
250
+
251
+ # Set project paths
252
+ project_path = Path(projectname).resolve()
253
+ src_path = project_path / "src" / project_path.name
254
+
255
+ # UV defaults to src structure. Create a flat structure instead:
256
+ if src_path.exists():
257
+ shutil.move(src_path, project_path / project_path.name)
258
+ if (project_path / "src").exists():
259
+ shutil.rmtree(project_path / "src")
260
+ src_path = project_path / project_path.name
261
+ assert src_path.exists(), "Src path does not exist"
262
+
263
+ # Update the dependencies in pyproject.toml and set the minimum python version
264
+ update_pyproject_toml_file(project_path)
265
+
266
+ # Update default __init__.py file generated by uv
267
+ replace_init_file(src_path, build_environment_template)
268
+
269
+ # Copy over the required files.
270
+ if build_environment_template:
271
+ file_path = src_path / f"{project_path.name}Env.py"
272
+ copy_env_template(file_path)
273
+
274
+ if include_algorithm_source:
275
+ file_path = src_path / "algorithms"
276
+ copy_algorithms_source(file_path)
277
+
278
+ with pkg_resources.path(
279
+ "create_rl_app.resources", "train_template.py"
280
+ ) as template_path:
281
+ shutil.copy(template_path, project_path / "train_example.py")
282
+
283
+ # Read template content and apply regex replacements
284
+ with open(template_path, "r") as file:
285
+ content = file.read()
286
+
287
+ # Replace import statements and environment creation based on configuration
288
+ if build_environment_template:
289
+ # Replace the jym.make call to use ExampleEnv
290
+ content = re.sub(
291
+ r'jym\.make\("CartPole-v1"\)',
292
+ 'jym.make("ExampleEnv")',
293
+ content,
294
+ flags=re.MULTILINE,
295
+ )
296
+ # Add import statement of example environment
297
+ import_line = (
298
+ f"from {project_path.name}.{project_path.name}Env import ExampleEnv\n"
299
+ )
300
+ content = import_line + content
301
+
302
+ # Replace algorithm import if copying source code
303
+ if include_algorithm_source:
304
+ content = re.sub(
305
+ r"^from jaxnasium\.algorithms.*$",
306
+ f"from {project_path.name}.algorithms import PPO",
307
+ content,
308
+ flags=re.MULTILINE,
309
+ )
310
+
311
+ # Write the modified content
312
+ with open(project_path / "train_example.py", "w") as file:
313
+ file.write(content)
314
+
315
+
316
+ if __name__ == "__main__":
317
+ main()
@@ -0,0 +1,89 @@
1
+ from typing import Tuple
2
+
3
+ import equinox as eqx
4
+ import jax.numpy as jnp
5
+ import jaxnasium as jym
6
+ from jaxtyping import Array, Float, PRNGKeyArray
7
+
8
+
9
+ class EnvState(eqx.Module):
10
+ x: int
11
+ y: int
12
+ time: int = 0
13
+
14
+ @property
15
+ def location(self):
16
+ return (self.x, self.y)
17
+
18
+
19
+ @jym.registry.register("ExampleEnv")
20
+ class ExampleEnv(jym.Environment):
21
+ max_episode_steps: int = 100
22
+
23
+ def step_env(
24
+ self, key: PRNGKeyArray, state: EnvState, action: int
25
+ ) -> Tuple[jym.TimeStep, EnvState]:
26
+ """
27
+ Update the environment state based on the action taken.
28
+ """
29
+ # action 0 -> move up
30
+ # action 1 -> move right
31
+ # action 2 -> move down
32
+ # action 3 -> move left
33
+ new_x = state.x + (action == 0) - (action == 2)
34
+ new_y = state.y + (action == 1) - (action == 3)
35
+
36
+ state = EnvState(x=new_x, y=new_y, time=state.time + 1)
37
+
38
+ timestep = jym.TimeStep(
39
+ observation=self.get_observation(state),
40
+ reward=self.get_reward(),
41
+ terminated=self.get_terminated(state),
42
+ truncated=state.time >= self.max_episode_steps,
43
+ info={},
44
+ )
45
+ return timestep, state
46
+
47
+ def reset_env(self, key: PRNGKeyArray) -> Tuple[Array, EnvState]:
48
+ """
49
+ Reset the environment to its initial state.
50
+ """
51
+ state = EnvState(x=5, y=5) # Start in the center
52
+ observation = self.get_observation(state)
53
+ return observation, state
54
+
55
+ def get_observation(self, state: EnvState) -> Array:
56
+ """
57
+ Get the observation from the environment state.
58
+ """
59
+ return jnp.array(state.location)
60
+
61
+ def get_reward(self) -> float:
62
+ """
63
+ Get the reward from the environment state.
64
+ """
65
+ # Example reward function: 1 for each step taken
66
+ return 1.0
67
+
68
+ def get_terminated(self, state: EnvState) -> Float[Array, "1"]:
69
+ """
70
+ Check if the episode has terminated.
71
+ """
72
+ # Example termination condition: if the agent moves out of bounds
73
+ out_of_bounds_x = jnp.logical_or(state.x < 0, state.x >= 10)
74
+ out_of_bounds_y = jnp.logical_or(state.y < 0, state.y >= 10)
75
+ return jnp.logical_or(out_of_bounds_x, out_of_bounds_y)
76
+
77
+ @property
78
+ def observation_space(self) -> jym.Space:
79
+ """
80
+ Define the observation space of the environment.
81
+ """
82
+ return jym.Box(low=0, high=10, shape=(2,))
83
+
84
+ @property
85
+ def action_space(self) -> jym.Space:
86
+ """
87
+ Define the action space of the environment.
88
+ """
89
+ return jym.Discrete(n=4)
@@ -0,0 +1,43 @@
1
+ import jax
2
+ import jaxnasium as jym
3
+ from jaxnasium.algorithms import PPO
4
+ from jaxtyping import PRNGKeyArray
5
+
6
+
7
+ def do_random_evaluation(
8
+ key: PRNGKeyArray, env: jym.Environment, num_repitions: int = 10
9
+ ):
10
+ """Perform some random steps to set a baseline for the environment."""
11
+ rewards = 0.0
12
+ for _ in range(num_repitions):
13
+ obs, env_state = env.reset(key)
14
+ while True:
15
+ key, key = jax.random.split(key)
16
+ action = env.action_space.sample(key)
17
+ (obs, reward, terminated, truncated, info), env_state = env.step(
18
+ key, env_state, action
19
+ )
20
+ rewards += reward
21
+ if terminated or truncated:
22
+ break
23
+ return rewards / num_repitions
24
+
25
+
26
+ if __name__ == "__main__":
27
+ env = jym.make("CartPole-v1")
28
+ env = jym.LogWrapper(env)
29
+ rng = jax.random.PRNGKey(0)
30
+
31
+ random_rewards = do_random_evaluation(rng, env)
32
+ print(f"Random Agent average reward: {random_rewards}")
33
+
34
+ # RL Training with PPO
35
+ agent = PPO(
36
+ total_timesteps=50000,
37
+ num_steps=64,
38
+ learning_rate=2.5e-3,
39
+ ent_coef=0.0,
40
+ num_epochs=1,
41
+ )
42
+ agent = agent.train(rng, env)
43
+ print(f"Agent average reward: {agent.evaluate(rng, env, num_eval_episodes=5)}")