create-rl-app 0.0.1__py3-none-any.whl
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.
- create_rl_app/__init__.py +0 -0
- create_rl_app/_vendored/.gitkeep +2 -0
- create_rl_app/_vendored/__init__.py +0 -0
- create_rl_app/cli.py +317 -0
- create_rl_app/resources/__init__.py +0 -0
- create_rl_app/resources/env_template.py +89 -0
- create_rl_app/resources/train_template.py +43 -0
- create_rl_app-0.0.1.dist-info/METADATA +63 -0
- create_rl_app-0.0.1.dist-info/RECORD +11 -0
- create_rl_app-0.0.1.dist-info/WHEEL +4 -0
- create_rl_app-0.0.1.dist-info/entry_points.txt +3 -0
|
File without changes
|
|
File without changes
|
create_rl_app/cli.py
ADDED
|
@@ -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()
|
|
File without changes
|
|
@@ -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)}")
|
|
@@ -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
|
+
[](https://python.org)
|
|
13
|
+
[](https://github.com/ponseko/jaxnasium)
|
|
14
|
+
[](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,11 @@
|
|
|
1
|
+
create_rl_app/__init__.py,sha256=e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855,0
|
|
2
|
+
create_rl_app/_vendored/.gitkeep,sha256=7db74051a8e96297141b78b840705b8bffbcb06612faeb73be5fbd5c220d5a40,124
|
|
3
|
+
create_rl_app/_vendored/__init__.py,sha256=e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855,0
|
|
4
|
+
create_rl_app/cli.py,sha256=b6f7116d449a47e38326857b140c5721c9d0b55cc201761a0b4c384c05fa88c6,11949
|
|
5
|
+
create_rl_app/resources/__init__.py,sha256=e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855,0
|
|
6
|
+
create_rl_app/resources/env_template.py,sha256=159f22c4eddfef96648364da1a86d4a0f333babcb353045f27e2790bef4f1010,2631
|
|
7
|
+
create_rl_app/resources/train_template.py,sha256=536bffe2ec33b766fd2900088985f75949bbe9a50a38be516ff90be36ddc362a,1282
|
|
8
|
+
create_rl_app-0.0.1.dist-info/WHEEL,sha256=b6dc288e80aa2d1b1518ddb3502fd5b53e8fd6cb507ed2a4f932e9e6088b264a,78
|
|
9
|
+
create_rl_app-0.0.1.dist-info/entry_points.txt,sha256=d014ab52c6461ade392347d61bc5a7b6f3a219c1074267ce11b418eba0942d3e,58
|
|
10
|
+
create_rl_app-0.0.1.dist-info/METADATA,sha256=65ededfa469afc6ef3283fc36e5bee0174b802ffc1bf5ae69056f2b7dc689e5c,1853
|
|
11
|
+
create_rl_app-0.0.1.dist-info/RECORD,,
|