JAXtari 0.1.0__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.
- jaxtari-0.1.0/.gitignore +12 -0
- jaxtari-0.1.0/LICENSE +21 -0
- jaxtari-0.1.0/PKG-INFO +409 -0
- jaxtari-0.1.0/README.md +341 -0
- jaxtari-0.1.0/pyproject.toml +79 -0
- jaxtari-0.1.0/src/jaxatari/__init__.py +32 -0
- jaxtari-0.1.0/src/jaxatari/core.py +207 -0
- jaxtari-0.1.0/src/jaxatari/environment.py +289 -0
- jaxtari-0.1.0/src/jaxatari/games/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/amidar_mazes.py +23 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_airraid.py +1168 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_alien.py +3108 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_amidar.py +1303 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_asterix.py +1114 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_asteroids.py +1420 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_atlantis.py +1592 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_bankheist.py +1644 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_beamrider.py +4909 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_berzerk.py +2258 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_blackjack.py +1021 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_breakout.py +1065 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_casino.py +355 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_casino_blackjack.py +1174 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_casino_five_stud_poker.py +499 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_casino_poker_solitaire.py +421 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_centipede.py +2569 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_choppercommand.py +2211 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_donkeykong.py +2532 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_enduro.py +1769 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_fishingderby.py +1879 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_flagcapture.py +718 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_freeway.py +723 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_frostbite.py +3742 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_galaxian.py +1779 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_gravitar.py +4124 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_hangman.py +767 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_hauntedhouse.py +1524 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_humancannonball.py +1022 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_kangaroo.py +2367 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_kingkong.py +2631 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_klax.py +1161 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_lasergates.py +3511 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_montezumarevenge.py +1016 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_mspacman.py +1687 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_namethisgame.py +1649 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_pacman.py +1301 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_phoenix.py +2883 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_pong.py +590 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_qbert.py +1581 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_riverraid.py +2209 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_seaquest.py +2856 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_sirlancelot.py +2911 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_skiing.py +1242 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_slotmachine.py +1474 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_spaceinvaders.py +1427 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_spacewar.py +1128 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_surround.py +765 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_tennis.py +1737 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_tetris.py +787 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_timepilot.py +1934 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_tron.py +2927 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_turmoil.py +2145 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_venture.py +2136 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_videocheckers.py +1461 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_videocube.py +1396 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_videopinball.py +4381 -0
- jaxtari-0.1.0/src/jaxatari/games/jax_wordzapper.py +2115 -0
- jaxtari-0.1.0/src/jaxatari/games/kangaroo_levels.py +249 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/alien/alien_mod_plugins.py +180 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/alien_mods.py +39 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/asteroids/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/asteroids/asteroids_mod_plugins.py +345 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/asteroids_mods.py +30 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/atlantis/atlantis_mod_plugins.py +114 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/atlantis_mods.py +39 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/bankheist/bankheist_mod_plugins.py +668 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/bankheist_mods.py +61 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/beamrider/beamrider_mod_plugins.py +3623 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/beamrider_mods.py +41 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/breakout/breakout_mod_plugins.py +141 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/breakout_mods.py +44 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/enduro/enduro_mod_plugins.py +216 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/enduro_mods.py +36 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/fishingderby/fishingderby_mod_plugins.py +82 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/fishingderby_mods.py +39 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/freeway/freeway_mod_plugins.py +234 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/freeway_mods.py +40 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/frostbite/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/frostbite/frostbite_mod_plugins.py +190 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/frostbite_mods.py +48 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/gravitar/gravitar_mod_plugins.py +166 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/gravitar_mods.py +53 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/kangaroo_mod_plugins.py +1189 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/cactus.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/cactus_tall.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/chicken.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/danger_sign.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/dragon.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/fireball.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/flame_0.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/flame_1.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/honey_bee.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/kangaroo_rope_climb.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/polarbear.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/snake.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/tank.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/tank_15x8.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/wasp.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo_mods.py +86 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/montezuma_revenge/montezuma_revenge_mod_plugins.py +285 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/montezuma_revenge_mods.py +50 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/mspacman/mspacman_mod_plugins.py +364 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/mspacman_mods.py +53 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/pacman/pacman_mod_plugins.py +269 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/pacman_mods.py +40 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/phoenix/phoenix_mod_plugins.py +138 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/phoenix_mods.py +48 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/pong/pong_mod_plugins.py +165 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/pong_mods.py +36 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/qbert/qbert_mod_plugins.py +673 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/qbert_mods.py +55 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/seaquest_mod_plugins.py +123 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/sprites/fireball.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/sprites/mine.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/seaquest_mods.py +115 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/skiing_mod_plugins.py +390 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skier_fallen.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_0.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_1.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_2.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_3.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_4.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_5.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_6.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_7.npy +0 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/skiing_mods.py +53 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/spaceinvaders/spaceinvaders_mod_plugins.py +108 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/spaceinvaders_mods.py +37 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/tennis/tennis_mod_plugins.py +376 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/tennis_mods.py +54 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/venture/venture_mod_plugins.py +148 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/venture_mods.py +43 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/videopinball/videopinball_mod_plugins.py +75 -0
- jaxtari-0.1.0/src/jaxatari/games/mods/videopinball_mods.py +31 -0
- jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/core.py +224 -0
- jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/renderer.py +931 -0
- jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/rooms.py +1094 -0
- jaxtari-0.1.0/src/jaxatari/games/mspacman_mazes.py +285 -0
- jaxtari-0.1.0/src/jaxatari/games/timepilot_levels.py +177 -0
- jaxtari-0.1.0/src/jaxatari/games/videopinball_constants.py +1738 -0
- jaxtari-0.1.0/src/jaxatari/gym_wrapper.py +369 -0
- jaxtari-0.1.0/src/jaxatari/install_sprites.py +155 -0
- jaxtari-0.1.0/src/jaxatari/modification.py +1024 -0
- jaxtari-0.1.0/src/jaxatari/py.typed +0 -0
- jaxtari-0.1.0/src/jaxatari/renderers.py +15 -0
- jaxtari-0.1.0/src/jaxatari/rendering/__init__.py +0 -0
- jaxtari-0.1.0/src/jaxatari/rendering/jax_rendering_utils.py +1355 -0
- jaxtari-0.1.0/src/jaxatari/spaces.py +386 -0
- jaxtari-0.1.0/src/jaxatari/wrappers.py +1038 -0
jaxtari-0.1.0/.gitignore
ADDED
jaxtari-0.1.0/LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2025 Quentin Delfosse, Jannis Blüml
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
jaxtari-0.1.0/PKG-INFO
ADDED
|
@@ -0,0 +1,409 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: JAXtari
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: GPU-accelerated, object-centric Atari environments for reinforcement learning with JAX.
|
|
5
|
+
Project-URL: Homepage, https://github.com/k4ntz/JAXAtari
|
|
6
|
+
Project-URL: Repository, https://github.com/k4ntz/JAXAtari
|
|
7
|
+
License: MIT License
|
|
8
|
+
|
|
9
|
+
Copyright (c) 2025 Quentin Delfosse, Jannis Blüml
|
|
10
|
+
|
|
11
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
12
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
13
|
+
in the Software without restriction, including without limitation the rights
|
|
14
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
15
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
16
|
+
furnished to do so, subject to the following conditions:
|
|
17
|
+
|
|
18
|
+
The above copyright notice and this permission notice shall be included in all
|
|
19
|
+
copies or substantial portions of the Software.
|
|
20
|
+
|
|
21
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
22
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
23
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
24
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
25
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
26
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
27
|
+
SOFTWARE.
|
|
28
|
+
License-File: LICENSE
|
|
29
|
+
Requires-Python: >=3.10
|
|
30
|
+
Requires-Dist: absl-py>=2.3
|
|
31
|
+
Requires-Dist: ale-py>=0.11.1
|
|
32
|
+
Requires-Dist: chex>=0.1.87
|
|
33
|
+
Requires-Dist: flax
|
|
34
|
+
Requires-Dist: gymnasium>=1.2.0
|
|
35
|
+
Requires-Dist: gymnax>=0.0.8
|
|
36
|
+
Requires-Dist: jax
|
|
37
|
+
Requires-Dist: ml-dtypes
|
|
38
|
+
Requires-Dist: numpy
|
|
39
|
+
Requires-Dist: opt-einsum>=3.4.0
|
|
40
|
+
Requires-Dist: platformdirs>=4.5.1
|
|
41
|
+
Requires-Dist: requests>=2.32.5
|
|
42
|
+
Requires-Dist: scipy>=1.15.3
|
|
43
|
+
Requires-Dist: toolz>=1.0.0
|
|
44
|
+
Requires-Dist: typing-extensions>=4.14.0
|
|
45
|
+
Provides-Extra: dev
|
|
46
|
+
Requires-Dist: gymnasium[other]==1.2.0; extra == 'dev'
|
|
47
|
+
Requires-Dist: pygame==2.5.0; extra == 'dev'
|
|
48
|
+
Requires-Dist: pytest; extra == 'dev'
|
|
49
|
+
Requires-Dist: syrupy==4.9.1; extra == 'dev'
|
|
50
|
+
Provides-Extra: gh-ci
|
|
51
|
+
Requires-Dist: opencv-python-headless; extra == 'gh-ci'
|
|
52
|
+
Requires-Dist: pygame==2.5.0; extra == 'gh-ci'
|
|
53
|
+
Requires-Dist: pytest; extra == 'gh-ci'
|
|
54
|
+
Requires-Dist: pytest-github-actions-annotate-failures; extra == 'gh-ci'
|
|
55
|
+
Requires-Dist: pytest-sugar; extra == 'gh-ci'
|
|
56
|
+
Requires-Dist: pytest-xdist; extra == 'gh-ci'
|
|
57
|
+
Requires-Dist: syrupy==4.9.1; extra == 'gh-ci'
|
|
58
|
+
Provides-Extra: training
|
|
59
|
+
Requires-Dist: hydra-core>=1.3.2; extra == 'training'
|
|
60
|
+
Requires-Dist: omegaconf>=2.3.0; extra == 'training'
|
|
61
|
+
Requires-Dist: rtpt>=0.0.4; extra == 'training'
|
|
62
|
+
Requires-Dist: safetensors>=0.7.0; extra == 'training'
|
|
63
|
+
Requires-Dist: tensorboard; extra == 'training'
|
|
64
|
+
Requires-Dist: torch; extra == 'training'
|
|
65
|
+
Requires-Dist: tyro; extra == 'training'
|
|
66
|
+
Requires-Dist: wandb[media]>=0.24.0; extra == 'training'
|
|
67
|
+
Description-Content-Type: text/markdown
|
|
68
|
+
|
|
69
|
+
# JAXtari: High-Throughput and Easy-to-Modify Arcade Learning Environment
|
|
70
|
+
|
|
71
|
+
Quentin Delfosse*, Raban Emunds*, Paul Seitz*, Sebastian Wette*, Jannis Blüml*, Daniel Kirn, Dominik Mandok, Kristian Kersting —
|
|
72
|
+
[AI/ML Lab, TU Darmstadt](https://www.aiml.informatik.tu-darmstadt.de/)
|
|
73
|
+
|
|
74
|
+
[Citation](#citation) • [Features](#features) • [Installation](#installation) • [Quick Start](#quick-start) • [Wrappers](#wrapper-reference) • [Environments](#available-environments) • [Contributing](#contributing) • [License](LICENSE)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
**JAXtari** is a GPU-accelerated, object-centric Atari environment framework powered by [JAX](https://github.com/google/jax). Inspired by [OCAtari](https://github.com/k4ntz/OC_Atari), it enables training agents with 100M steps in under 1 hour (pixel-based observations) or under 15 minutes (object-centric observation) through JIT compilation, vectorization, and full GPU parallelization — while exposing structured, object-centric observations alongside standard pixel inputs. Similar to [HackAtari](https://github.com/k4ntz/HackAtari), it also supports game modifications for testing agent generalization.
|
|
78
|
+
|
|
79
|
+
---
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
<div class="collage">
|
|
83
|
+
<div class="row" align="center">
|
|
84
|
+
<img src="./docs/source/_static/gifs/pong.gif" alt="Pong" width="24%">
|
|
85
|
+
<img src="./docs/source/_static/gifs/beamrider.gif" alt="Beamrider" width="24%">
|
|
86
|
+
<img src="./docs/source/_static/gifs/phoenix.gif" alt="Phoenix" width="24%">
|
|
87
|
+
<img src="./docs/source/_static/gifs/tennis.gif" alt="Tennis" width="24%">
|
|
88
|
+
</div>
|
|
89
|
+
<div class="row" align="center">
|
|
90
|
+
<img src="./docs/source/_static/gifs/skiing.gif" alt="Skiing" width="24%">
|
|
91
|
+
<img src="./docs/source/_static/gifs/montezumarevenge.gif" alt="Montezuma" width="24%">
|
|
92
|
+
<img src="./docs/source/_static/gifs/seaquest.gif" alt="Seaquest" width="24%">
|
|
93
|
+
<img src="./docs/source/_static/gifs/kangaroo.gif" alt="Kangaroo" width="24%">
|
|
94
|
+
</div>
|
|
95
|
+
<div class="row" align="center">
|
|
96
|
+
<img src="./docs/source/_static/gifs/freeway.gif" alt="Freeway" width="24%">
|
|
97
|
+
<img src="./docs/source/_static/gifs/venture.gif" alt="Venture" width="24%">
|
|
98
|
+
<img src="./docs/source/_static/gifs/qbert.gif" alt="Qbert" width="24%">
|
|
99
|
+
<img src="./docs/source/_static/gifs/frostbite.gif" alt="Frostbite" width="24%">
|
|
100
|
+
</div>
|
|
101
|
+
<div class="row" align="center">
|
|
102
|
+
<img src="./docs/source/_static/gifs/bankheist.gif" alt="Bankheist" width="24%">
|
|
103
|
+
<img src="./docs/source/_static/gifs/mspacman.gif" alt="Ms. PacMan" width="24%">
|
|
104
|
+
<img src="./docs/source/_static/gifs/gravitar.gif" alt="Gravitar" width="24%">
|
|
105
|
+
<img src="./docs/source/_static/gifs/enduro.gif" alt="Enduro" width="24%">
|
|
106
|
+
</div>
|
|
107
|
+
</div>
|
|
108
|
+
|
|
109
|
+
---
|
|
110
|
+
|
|
111
|
+
## Features
|
|
112
|
+
|
|
113
|
+
- **Object-centric observations** — structured game state with per-object positions, types, and attributes
|
|
114
|
+
- **Full GPU pipeline** — end-to-end JAX with JIT compilation, `vmap`, and `lax.scan`; no CPU/GPU transfer bottlenecks
|
|
115
|
+
- **Comprehensive wrapper system** — pixel, object-centric, combined, normalized, flattened — all composable
|
|
116
|
+
- **Game modifications** — pre-built mods and a clean API for custom distribution shifts
|
|
117
|
+
|
|
118
|
+
---
|
|
119
|
+
|
|
120
|
+
## Installation
|
|
121
|
+
|
|
122
|
+
### Basic
|
|
123
|
+
|
|
124
|
+
```bash
|
|
125
|
+
python3 -m venv .venv
|
|
126
|
+
source .venv/bin/activate
|
|
127
|
+
pip install -U pip
|
|
128
|
+
pip install -e .
|
|
129
|
+
```
|
|
130
|
+
|
|
131
|
+
### With development tools (tests + manual play)
|
|
132
|
+
|
|
133
|
+
Includes `pytest`, `pygame`, and testing extras:
|
|
134
|
+
|
|
135
|
+
```bash
|
|
136
|
+
pip install -e ".[dev]"
|
|
137
|
+
```
|
|
138
|
+
|
|
139
|
+
### With training scripts
|
|
140
|
+
|
|
141
|
+
Includes `wandb`, `tensorboard`, `hydra`, and other training dependencies:
|
|
142
|
+
|
|
143
|
+
```bash
|
|
144
|
+
pip install -e ".[training]"
|
|
145
|
+
```
|
|
146
|
+
|
|
147
|
+
### GPU acceleration (CUDA)
|
|
148
|
+
|
|
149
|
+
```bash
|
|
150
|
+
pip install -U "jax[cuda12]"
|
|
151
|
+
```
|
|
152
|
+
|
|
153
|
+
For other accelerators see the [JAX installation guide](https://docs.jax.dev/en/latest/installation.html).
|
|
154
|
+
|
|
155
|
+
### Download sprites
|
|
156
|
+
|
|
157
|
+
Before running any environment for the first time you will be asked to confirm ROM ownership of the original Atari ROMs. This is necessary to download sprites that look similar to the original ALE sprites.
|
|
158
|
+
|
|
159
|
+
If you do not have ownership of the original Atari ROMs, you can continue with replacement/custom sprites. In that case, please decline the ownership and the installer will download the alternative sprites package.
|
|
160
|
+
You can also use your own sprites by placing them in the ~/.local/share/jaxatari/sprites directory.
|
|
161
|
+
|
|
162
|
+
```bash
|
|
163
|
+
python3 src/jaxatari/install_sprites.py
|
|
164
|
+
```
|
|
165
|
+
|
|
166
|
+
---
|
|
167
|
+
|
|
168
|
+
## Quick Start
|
|
169
|
+
|
|
170
|
+
### Basic environment creation
|
|
171
|
+
|
|
172
|
+
```python
|
|
173
|
+
import jax
|
|
174
|
+
import jaxatari
|
|
175
|
+
|
|
176
|
+
env = jaxatari.make("pong")
|
|
177
|
+
|
|
178
|
+
# List all available games
|
|
179
|
+
print(jaxatari.list_available_games())
|
|
180
|
+
```
|
|
181
|
+
|
|
182
|
+
### Game modifications
|
|
183
|
+
|
|
184
|
+
JAXtari ships with pre-built modifications for testing generalization:
|
|
185
|
+
|
|
186
|
+
```python
|
|
187
|
+
import jaxatari
|
|
188
|
+
|
|
189
|
+
# Single mod
|
|
190
|
+
env = jaxatari.make("pong", mods=["lazy_enemy"])
|
|
191
|
+
|
|
192
|
+
# Multiple mods simultaneously
|
|
193
|
+
env = jaxatari.make("pong", mods=["lazy_enemy", "shift_enemy"])
|
|
194
|
+
```
|
|
195
|
+
|
|
196
|
+
### Applying wrappers
|
|
197
|
+
|
|
198
|
+
Wrappers must be applied in order: `AtariWrapper` first, then an observation wrapper, then optional utility wrappers.
|
|
199
|
+
|
|
200
|
+
```python
|
|
201
|
+
import jaxatari
|
|
202
|
+
from jaxatari.wrappers import (
|
|
203
|
+
AtariWrapper,
|
|
204
|
+
ObjectCentricWrapper,
|
|
205
|
+
PixelObsWrapper,
|
|
206
|
+
PixelAndObjectCentricWrapper,
|
|
207
|
+
FlattenObservationWrapper,
|
|
208
|
+
NormalizeObservationWrapper,
|
|
209
|
+
LogWrapper,
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
base_env = jaxatari.make("pong")
|
|
213
|
+
atari_env = AtariWrapper(base_env)
|
|
214
|
+
|
|
215
|
+
# Choose one observation type:
|
|
216
|
+
env = ObjectCentricWrapper(atari_env, frame_stack_size=4, frame_skip=4) # shape: (frame_stack, features)
|
|
217
|
+
# env = PixelObsWrapper(atari_env) # shape: (frame_stack, H, W, C)
|
|
218
|
+
# env = PixelAndObjectCentricWrapper(atari_env) # both
|
|
219
|
+
|
|
220
|
+
# Optional: flatten to 1D
|
|
221
|
+
env = FlattenObservationWrapper(env)
|
|
222
|
+
|
|
223
|
+
# Optional: normalize observations to [0, 1]
|
|
224
|
+
env = NormalizeObservationWrapper(env)
|
|
225
|
+
|
|
226
|
+
# Optional: track episode returns and lengths
|
|
227
|
+
env = LogWrapper(env)
|
|
228
|
+
```
|
|
229
|
+
|
|
230
|
+
### Vectorized stepping
|
|
231
|
+
|
|
232
|
+
```python
|
|
233
|
+
import jax
|
|
234
|
+
import jaxatari
|
|
235
|
+
from jaxatari.wrappers import AtariWrapper, ObjectCentricWrapper, FlattenObservationWrapper
|
|
236
|
+
|
|
237
|
+
env = FlattenObservationWrapper(ObjectCentricWrapper(AtariWrapper(jaxatari.make("pong"))))
|
|
238
|
+
|
|
239
|
+
n_envs = 1024
|
|
240
|
+
rng = jax.random.PRNGKey(0)
|
|
241
|
+
reset_keys = jax.random.split(rng, n_envs)
|
|
242
|
+
|
|
243
|
+
# Initialise n_envs parallel environments
|
|
244
|
+
obs, env_state = jax.vmap(env.reset)(reset_keys)
|
|
245
|
+
|
|
246
|
+
# Single parallel step
|
|
247
|
+
action = jax.random.randint(rng, (n_envs,), 0, env.action_space().n)
|
|
248
|
+
obs, env_state, reward, terminated, truncated, info = jax.vmap(env.step)(env_state, action)
|
|
249
|
+
|
|
250
|
+
# 100 steps with scan
|
|
251
|
+
def step_fn(carry, _):
|
|
252
|
+
obs, state = carry
|
|
253
|
+
new_obs, new_state, reward, terminated, truncated, info = jax.vmap(env.step)(state, action)
|
|
254
|
+
return (new_obs, new_state), (reward, terminated, truncated, info)
|
|
255
|
+
|
|
256
|
+
_, (rewards, terminations, truncations, infos) = jax.lax.scan(
|
|
257
|
+
step_fn, (obs, env_state), None, length=100
|
|
258
|
+
)
|
|
259
|
+
```
|
|
260
|
+
|
|
261
|
+
### Gymnasium compatibility *(WIP)*
|
|
262
|
+
|
|
263
|
+
> **Note:** This wrapper is currently work in progress and supports interoperability with CPU-based Gymnasium pipelines (e.g. stable-baselines3). It currently only exposes pixel observations and does not accept JAXtari wrappers. For JAX-native training use the wrapper stack above instead.
|
|
264
|
+
|
|
265
|
+
```python
|
|
266
|
+
from jaxatari.gym_wrapper import GymnasiumJaxAtariWrapper
|
|
267
|
+
import jaxatari
|
|
268
|
+
|
|
269
|
+
base_env = jaxatari.make("pong")
|
|
270
|
+
gym_env = GymnasiumJaxAtariWrapper(base_env)
|
|
271
|
+
|
|
272
|
+
obs, info = gym_env.reset()
|
|
273
|
+
obs, reward, terminated, truncated, info = gym_env.step(gym_env.action_space.sample())
|
|
274
|
+
```
|
|
275
|
+
|
|
276
|
+
### Multiple reward functions
|
|
277
|
+
|
|
278
|
+
Use `MultiRewardWrapper` to compute several reward signals in parallel (apply it directly after the base environment, before any other wrapper):
|
|
279
|
+
|
|
280
|
+
```python
|
|
281
|
+
import jaxatari
|
|
282
|
+
from jaxatari.wrappers import MultiRewardWrapper, AtariWrapper, ObjectCentricWrapper, MultiRewardLogWrapper
|
|
283
|
+
|
|
284
|
+
def survival_reward(prev_state, state):
|
|
285
|
+
return 1.0 # reward every surviving step
|
|
286
|
+
|
|
287
|
+
def score_delta(prev_state, state):
|
|
288
|
+
return state.score - prev_state.score
|
|
289
|
+
|
|
290
|
+
base_env = jaxatari.make("pong")
|
|
291
|
+
env = MultiRewardWrapper(base_env, reward_funcs=[survival_reward, score_delta])
|
|
292
|
+
env = ObjectCentricWrapper(AtariWrapper(env))
|
|
293
|
+
env = MultiRewardLogWrapper(env)
|
|
294
|
+
```
|
|
295
|
+
|
|
296
|
+
### Manual play
|
|
297
|
+
|
|
298
|
+
```bash
|
|
299
|
+
# requires the [dev] extra (pygame)
|
|
300
|
+
python3 scripts/play.py -g Pong --mods lazy_enemy
|
|
301
|
+
```
|
|
302
|
+
|
|
303
|
+
---
|
|
304
|
+
|
|
305
|
+
## Wrapper Reference
|
|
306
|
+
|
|
307
|
+
All wrappers live in `src/jaxatari/wrappers.py`. The standard stack is:
|
|
308
|
+
|
|
309
|
+
```
|
|
310
|
+
base env → [MultiRewardWrapper] → AtariWrapper → <obs wrapper> → [utility wrappers]
|
|
311
|
+
```
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
| Wrapper | Description |
|
|
315
|
+
| ------------------------------ | -------------------------------------------------------------------------------------------------------------------------------------- |
|
|
316
|
+
| `AtariWrapper` | Atari-specific pre-processing: sticky actions, episodic life, noop reset, frame-skip config. Must come before any observation wrapper. |
|
|
317
|
+
| `ObjectCentricWrapper` | Stacked object-centric features. Output shape: `(frame_stack, features)`. |
|
|
318
|
+
| `PixelObsWrapper` | Stacked pixel frames with max-pooling. Output shape: `(frame_stack, H, W, C)`. |
|
|
319
|
+
| `PixelAndObjectCentricWrapper` | Both pixel and object-centric observations as a tuple. |
|
|
320
|
+
| `PixelAndObjectObsWrapper` | Same as above but returns structured (non-flattened) OC observations. |
|
|
321
|
+
| `FlattenObservationWrapper` | Flattens any observation pytree to a single 1D array. |
|
|
322
|
+
| `NormalizeObservationWrapper` | Normalizes observations to `[0, 1]` (or `[-1, 1]` with `to_neg_one=True`). Compatible with any pytree structure. |
|
|
323
|
+
| `LogWrapper` | Tracks episode returns and lengths. |
|
|
324
|
+
| `MultiRewardWrapper` | Computes multiple reward functions at every step. Apply before `AtariWrapper`. |
|
|
325
|
+
| `MultiRewardLogWrapper` | Tracks multiple reward components separately. Use with `MultiRewardWrapper`. |
|
|
326
|
+
|
|
327
|
+
---
|
|
328
|
+
|
|
329
|
+
## Project Structure
|
|
330
|
+
|
|
331
|
+
```
|
|
332
|
+
JAXAtari/
|
|
333
|
+
├── src/jaxatari/
|
|
334
|
+
│ ├── core.py # make() factory, game and mod registries
|
|
335
|
+
│ ├── environment.py # JaxEnvironment base class
|
|
336
|
+
│ ├── wrappers.py # all wrappers
|
|
337
|
+
│ ├── modification.py # mod system (plugins, conflict detection)
|
|
338
|
+
│ ├── gym_wrapper.py # Gymnasium compatibility adapter
|
|
339
|
+
│ ├── renderers.py # JAX rendering utilities
|
|
340
|
+
│ ├── spaces.py # action/observation space definitions
|
|
341
|
+
│ ├── install_sprites.py # sprite download script
|
|
342
|
+
│ └── games/
|
|
343
|
+
│ ├── jax_<game>.py # one file per environment
|
|
344
|
+
│ └── mods/
|
|
345
|
+
│ ├── <game>_mods.py # mod controller per game
|
|
346
|
+
│ └── <game>/
|
|
347
|
+
│ └── <game>_mod_plugins.py # individual mod plugin classes
|
|
348
|
+
├── scripts/ # see scripts/README.md for a full description
|
|
349
|
+
│ ├── play.py # interactive human play
|
|
350
|
+
│ └── benchmarks/ # PPO/PQN training and evaluation scripts
|
|
351
|
+
├── tests/ # pytest test suite
|
|
352
|
+
├── docs/ # Sphinx documentation source
|
|
353
|
+
├── games_covered.md # full environment status table
|
|
354
|
+
└── pyproject.toml
|
|
355
|
+
```
|
|
356
|
+
|
|
357
|
+
---
|
|
358
|
+
|
|
359
|
+
## Contributing
|
|
360
|
+
|
|
361
|
+
Contributions are welcome! See [CONTRIBUTING.md](CONTRIBUTING.md) for detailed guides on adding mods, environments, and wrappers. Quick overview below.
|
|
362
|
+
|
|
363
|
+
### Adding a new environment
|
|
364
|
+
|
|
365
|
+
1. Create `src/jaxatari/games/jax_<game>.py` implementing `JaxEnvironment`
|
|
366
|
+
2. Register it in `GAME_MODULES` in `src/jaxatari/core.py`
|
|
367
|
+
3. Add a test in `tests/games/`
|
|
368
|
+
4. Update your game's status in [games_covered.md](games_covered.md)
|
|
369
|
+
|
|
370
|
+
### Adding a mod
|
|
371
|
+
|
|
372
|
+
1. Create `src/jaxatari/games/mods/<game>/<game>_mod_plugins.py` with your plugin class(es) extending `JaxAtariInternalModPlugin` or `JaxAtariPostStepModPlugin`
|
|
373
|
+
2. Create or update `src/jaxatari/games/mods/<game>_mods.py` — add your mod key to the `REGISTRY` dict
|
|
374
|
+
3. Register the controller in `MOD_MODULES` in `src/jaxatari/core.py` (if not already present)
|
|
375
|
+
|
|
376
|
+
### Adding a wrapper
|
|
377
|
+
|
|
378
|
+
1. Subclass `JaxatariWrapper` in `src/jaxatari/wrappers.py`
|
|
379
|
+
2. Implement `reset()`, `step()`, and `observation_space()` / `action_space()`
|
|
380
|
+
3. Export it from `src/jaxatari/__init__.py`
|
|
381
|
+
|
|
382
|
+
### General
|
|
383
|
+
|
|
384
|
+
1. Fork this repository
|
|
385
|
+
2. Create a feature branch: `git checkout -b feature/my-feature`
|
|
386
|
+
3. Commit your changes and open a pull request
|
|
387
|
+
|
|
388
|
+
Feel free to share new mods or environments by opening a PR!
|
|
389
|
+
|
|
390
|
+
---
|
|
391
|
+
|
|
392
|
+
## Citation
|
|
393
|
+
|
|
394
|
+
```bibtex
|
|
395
|
+
@misc{jaxatari2026,
|
|
396
|
+
author = {Delfosse, Quentin and Emunds, Raban and Seitz, Paul and Wette, Sebastian and Kirn, Daniel and Mandok, Dominik and Bl{\"u}ml, Jannis and Kersting, Kristian},
|
|
397
|
+
title = {JAXtari: High-Throughput and Easy-to-Modify Arcade Learning Environment},
|
|
398
|
+
year = {2026},
|
|
399
|
+
publisher = {GitHub},
|
|
400
|
+
journal = {GitHub repository},
|
|
401
|
+
howpublished = {https://github.com/k4ntz/JAXAtari/},
|
|
402
|
+
}
|
|
403
|
+
```
|
|
404
|
+
|
|
405
|
+
---
|
|
406
|
+
|
|
407
|
+
## License
|
|
408
|
+
|
|
409
|
+
This project is licensed under the MIT License — see [LICENSE](LICENSE) for details.
|