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.
Files changed (162) hide show
  1. jaxtari-0.1.0/.gitignore +12 -0
  2. jaxtari-0.1.0/LICENSE +21 -0
  3. jaxtari-0.1.0/PKG-INFO +409 -0
  4. jaxtari-0.1.0/README.md +341 -0
  5. jaxtari-0.1.0/pyproject.toml +79 -0
  6. jaxtari-0.1.0/src/jaxatari/__init__.py +32 -0
  7. jaxtari-0.1.0/src/jaxatari/core.py +207 -0
  8. jaxtari-0.1.0/src/jaxatari/environment.py +289 -0
  9. jaxtari-0.1.0/src/jaxatari/games/__init__.py +0 -0
  10. jaxtari-0.1.0/src/jaxatari/games/amidar_mazes.py +23 -0
  11. jaxtari-0.1.0/src/jaxatari/games/jax_airraid.py +1168 -0
  12. jaxtari-0.1.0/src/jaxatari/games/jax_alien.py +3108 -0
  13. jaxtari-0.1.0/src/jaxatari/games/jax_amidar.py +1303 -0
  14. jaxtari-0.1.0/src/jaxatari/games/jax_asterix.py +1114 -0
  15. jaxtari-0.1.0/src/jaxatari/games/jax_asteroids.py +1420 -0
  16. jaxtari-0.1.0/src/jaxatari/games/jax_atlantis.py +1592 -0
  17. jaxtari-0.1.0/src/jaxatari/games/jax_bankheist.py +1644 -0
  18. jaxtari-0.1.0/src/jaxatari/games/jax_beamrider.py +4909 -0
  19. jaxtari-0.1.0/src/jaxatari/games/jax_berzerk.py +2258 -0
  20. jaxtari-0.1.0/src/jaxatari/games/jax_blackjack.py +1021 -0
  21. jaxtari-0.1.0/src/jaxatari/games/jax_breakout.py +1065 -0
  22. jaxtari-0.1.0/src/jaxatari/games/jax_casino.py +355 -0
  23. jaxtari-0.1.0/src/jaxatari/games/jax_casino_blackjack.py +1174 -0
  24. jaxtari-0.1.0/src/jaxatari/games/jax_casino_five_stud_poker.py +499 -0
  25. jaxtari-0.1.0/src/jaxatari/games/jax_casino_poker_solitaire.py +421 -0
  26. jaxtari-0.1.0/src/jaxatari/games/jax_centipede.py +2569 -0
  27. jaxtari-0.1.0/src/jaxatari/games/jax_choppercommand.py +2211 -0
  28. jaxtari-0.1.0/src/jaxatari/games/jax_donkeykong.py +2532 -0
  29. jaxtari-0.1.0/src/jaxatari/games/jax_enduro.py +1769 -0
  30. jaxtari-0.1.0/src/jaxatari/games/jax_fishingderby.py +1879 -0
  31. jaxtari-0.1.0/src/jaxatari/games/jax_flagcapture.py +718 -0
  32. jaxtari-0.1.0/src/jaxatari/games/jax_freeway.py +723 -0
  33. jaxtari-0.1.0/src/jaxatari/games/jax_frostbite.py +3742 -0
  34. jaxtari-0.1.0/src/jaxatari/games/jax_galaxian.py +1779 -0
  35. jaxtari-0.1.0/src/jaxatari/games/jax_gravitar.py +4124 -0
  36. jaxtari-0.1.0/src/jaxatari/games/jax_hangman.py +767 -0
  37. jaxtari-0.1.0/src/jaxatari/games/jax_hauntedhouse.py +1524 -0
  38. jaxtari-0.1.0/src/jaxatari/games/jax_humancannonball.py +1022 -0
  39. jaxtari-0.1.0/src/jaxatari/games/jax_kangaroo.py +2367 -0
  40. jaxtari-0.1.0/src/jaxatari/games/jax_kingkong.py +2631 -0
  41. jaxtari-0.1.0/src/jaxatari/games/jax_klax.py +1161 -0
  42. jaxtari-0.1.0/src/jaxatari/games/jax_lasergates.py +3511 -0
  43. jaxtari-0.1.0/src/jaxatari/games/jax_montezumarevenge.py +1016 -0
  44. jaxtari-0.1.0/src/jaxatari/games/jax_mspacman.py +1687 -0
  45. jaxtari-0.1.0/src/jaxatari/games/jax_namethisgame.py +1649 -0
  46. jaxtari-0.1.0/src/jaxatari/games/jax_pacman.py +1301 -0
  47. jaxtari-0.1.0/src/jaxatari/games/jax_phoenix.py +2883 -0
  48. jaxtari-0.1.0/src/jaxatari/games/jax_pong.py +590 -0
  49. jaxtari-0.1.0/src/jaxatari/games/jax_qbert.py +1581 -0
  50. jaxtari-0.1.0/src/jaxatari/games/jax_riverraid.py +2209 -0
  51. jaxtari-0.1.0/src/jaxatari/games/jax_seaquest.py +2856 -0
  52. jaxtari-0.1.0/src/jaxatari/games/jax_sirlancelot.py +2911 -0
  53. jaxtari-0.1.0/src/jaxatari/games/jax_skiing.py +1242 -0
  54. jaxtari-0.1.0/src/jaxatari/games/jax_slotmachine.py +1474 -0
  55. jaxtari-0.1.0/src/jaxatari/games/jax_spaceinvaders.py +1427 -0
  56. jaxtari-0.1.0/src/jaxatari/games/jax_spacewar.py +1128 -0
  57. jaxtari-0.1.0/src/jaxatari/games/jax_surround.py +765 -0
  58. jaxtari-0.1.0/src/jaxatari/games/jax_tennis.py +1737 -0
  59. jaxtari-0.1.0/src/jaxatari/games/jax_tetris.py +787 -0
  60. jaxtari-0.1.0/src/jaxatari/games/jax_timepilot.py +1934 -0
  61. jaxtari-0.1.0/src/jaxatari/games/jax_tron.py +2927 -0
  62. jaxtari-0.1.0/src/jaxatari/games/jax_turmoil.py +2145 -0
  63. jaxtari-0.1.0/src/jaxatari/games/jax_venture.py +2136 -0
  64. jaxtari-0.1.0/src/jaxatari/games/jax_videocheckers.py +1461 -0
  65. jaxtari-0.1.0/src/jaxatari/games/jax_videocube.py +1396 -0
  66. jaxtari-0.1.0/src/jaxatari/games/jax_videopinball.py +4381 -0
  67. jaxtari-0.1.0/src/jaxatari/games/jax_wordzapper.py +2115 -0
  68. jaxtari-0.1.0/src/jaxatari/games/kangaroo_levels.py +249 -0
  69. jaxtari-0.1.0/src/jaxatari/games/mods/__init__.py +0 -0
  70. jaxtari-0.1.0/src/jaxatari/games/mods/alien/alien_mod_plugins.py +180 -0
  71. jaxtari-0.1.0/src/jaxatari/games/mods/alien_mods.py +39 -0
  72. jaxtari-0.1.0/src/jaxatari/games/mods/asteroids/__init__.py +0 -0
  73. jaxtari-0.1.0/src/jaxatari/games/mods/asteroids/asteroids_mod_plugins.py +345 -0
  74. jaxtari-0.1.0/src/jaxatari/games/mods/asteroids_mods.py +30 -0
  75. jaxtari-0.1.0/src/jaxatari/games/mods/atlantis/atlantis_mod_plugins.py +114 -0
  76. jaxtari-0.1.0/src/jaxatari/games/mods/atlantis_mods.py +39 -0
  77. jaxtari-0.1.0/src/jaxatari/games/mods/bankheist/bankheist_mod_plugins.py +668 -0
  78. jaxtari-0.1.0/src/jaxatari/games/mods/bankheist_mods.py +61 -0
  79. jaxtari-0.1.0/src/jaxatari/games/mods/beamrider/beamrider_mod_plugins.py +3623 -0
  80. jaxtari-0.1.0/src/jaxatari/games/mods/beamrider_mods.py +41 -0
  81. jaxtari-0.1.0/src/jaxatari/games/mods/breakout/breakout_mod_plugins.py +141 -0
  82. jaxtari-0.1.0/src/jaxatari/games/mods/breakout_mods.py +44 -0
  83. jaxtari-0.1.0/src/jaxatari/games/mods/enduro/enduro_mod_plugins.py +216 -0
  84. jaxtari-0.1.0/src/jaxatari/games/mods/enduro_mods.py +36 -0
  85. jaxtari-0.1.0/src/jaxatari/games/mods/fishingderby/fishingderby_mod_plugins.py +82 -0
  86. jaxtari-0.1.0/src/jaxatari/games/mods/fishingderby_mods.py +39 -0
  87. jaxtari-0.1.0/src/jaxatari/games/mods/freeway/freeway_mod_plugins.py +234 -0
  88. jaxtari-0.1.0/src/jaxatari/games/mods/freeway_mods.py +40 -0
  89. jaxtari-0.1.0/src/jaxatari/games/mods/frostbite/__init__.py +0 -0
  90. jaxtari-0.1.0/src/jaxatari/games/mods/frostbite/frostbite_mod_plugins.py +190 -0
  91. jaxtari-0.1.0/src/jaxatari/games/mods/frostbite_mods.py +48 -0
  92. jaxtari-0.1.0/src/jaxatari/games/mods/gravitar/gravitar_mod_plugins.py +166 -0
  93. jaxtari-0.1.0/src/jaxatari/games/mods/gravitar_mods.py +53 -0
  94. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/kangaroo_mod_plugins.py +1189 -0
  95. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/cactus.npy +0 -0
  96. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/cactus_tall.npy +0 -0
  97. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/chicken.npy +0 -0
  98. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/danger_sign.npy +0 -0
  99. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/dragon.npy +0 -0
  100. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/fireball.npy +0 -0
  101. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/flame_0.npy +0 -0
  102. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/flame_1.npy +0 -0
  103. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/honey_bee.npy +0 -0
  104. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/kangaroo_rope_climb.npy +0 -0
  105. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/polarbear.npy +0 -0
  106. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/snake.npy +0 -0
  107. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/tank.npy +0 -0
  108. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/tank_15x8.npy +0 -0
  109. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo/sprites/wasp.npy +0 -0
  110. jaxtari-0.1.0/src/jaxatari/games/mods/kangaroo_mods.py +86 -0
  111. jaxtari-0.1.0/src/jaxatari/games/mods/montezuma_revenge/montezuma_revenge_mod_plugins.py +285 -0
  112. jaxtari-0.1.0/src/jaxatari/games/mods/montezuma_revenge_mods.py +50 -0
  113. jaxtari-0.1.0/src/jaxatari/games/mods/mspacman/mspacman_mod_plugins.py +364 -0
  114. jaxtari-0.1.0/src/jaxatari/games/mods/mspacman_mods.py +53 -0
  115. jaxtari-0.1.0/src/jaxatari/games/mods/pacman/pacman_mod_plugins.py +269 -0
  116. jaxtari-0.1.0/src/jaxatari/games/mods/pacman_mods.py +40 -0
  117. jaxtari-0.1.0/src/jaxatari/games/mods/phoenix/phoenix_mod_plugins.py +138 -0
  118. jaxtari-0.1.0/src/jaxatari/games/mods/phoenix_mods.py +48 -0
  119. jaxtari-0.1.0/src/jaxatari/games/mods/pong/pong_mod_plugins.py +165 -0
  120. jaxtari-0.1.0/src/jaxatari/games/mods/pong_mods.py +36 -0
  121. jaxtari-0.1.0/src/jaxatari/games/mods/qbert/qbert_mod_plugins.py +673 -0
  122. jaxtari-0.1.0/src/jaxatari/games/mods/qbert_mods.py +55 -0
  123. jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/seaquest_mod_plugins.py +123 -0
  124. jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/sprites/fireball.npy +0 -0
  125. jaxtari-0.1.0/src/jaxatari/games/mods/seaquest/sprites/mine.npy +0 -0
  126. jaxtari-0.1.0/src/jaxatari/games/mods/seaquest_mods.py +115 -0
  127. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/__init__.py +0 -0
  128. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/skiing_mod_plugins.py +390 -0
  129. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skier_fallen.npy +0 -0
  130. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_0.npy +0 -0
  131. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_1.npy +0 -0
  132. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_2.npy +0 -0
  133. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_3.npy +0 -0
  134. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_4.npy +0 -0
  135. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_5.npy +0 -0
  136. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_6.npy +0 -0
  137. jaxtari-0.1.0/src/jaxatari/games/mods/skiing/sprites/blue_skiier_7.npy +0 -0
  138. jaxtari-0.1.0/src/jaxatari/games/mods/skiing_mods.py +53 -0
  139. jaxtari-0.1.0/src/jaxatari/games/mods/spaceinvaders/spaceinvaders_mod_plugins.py +108 -0
  140. jaxtari-0.1.0/src/jaxatari/games/mods/spaceinvaders_mods.py +37 -0
  141. jaxtari-0.1.0/src/jaxatari/games/mods/tennis/tennis_mod_plugins.py +376 -0
  142. jaxtari-0.1.0/src/jaxatari/games/mods/tennis_mods.py +54 -0
  143. jaxtari-0.1.0/src/jaxatari/games/mods/venture/venture_mod_plugins.py +148 -0
  144. jaxtari-0.1.0/src/jaxatari/games/mods/venture_mods.py +43 -0
  145. jaxtari-0.1.0/src/jaxatari/games/mods/videopinball/videopinball_mod_plugins.py +75 -0
  146. jaxtari-0.1.0/src/jaxatari/games/mods/videopinball_mods.py +31 -0
  147. jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/__init__.py +0 -0
  148. jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/core.py +224 -0
  149. jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/renderer.py +931 -0
  150. jaxtari-0.1.0/src/jaxatari/games/montezuma_revenge/rooms.py +1094 -0
  151. jaxtari-0.1.0/src/jaxatari/games/mspacman_mazes.py +285 -0
  152. jaxtari-0.1.0/src/jaxatari/games/timepilot_levels.py +177 -0
  153. jaxtari-0.1.0/src/jaxatari/games/videopinball_constants.py +1738 -0
  154. jaxtari-0.1.0/src/jaxatari/gym_wrapper.py +369 -0
  155. jaxtari-0.1.0/src/jaxatari/install_sprites.py +155 -0
  156. jaxtari-0.1.0/src/jaxatari/modification.py +1024 -0
  157. jaxtari-0.1.0/src/jaxatari/py.typed +0 -0
  158. jaxtari-0.1.0/src/jaxatari/renderers.py +15 -0
  159. jaxtari-0.1.0/src/jaxatari/rendering/__init__.py +0 -0
  160. jaxtari-0.1.0/src/jaxatari/rendering/jax_rendering_utils.py +1355 -0
  161. jaxtari-0.1.0/src/jaxatari/spaces.py +386 -0
  162. jaxtari-0.1.0/src/jaxatari/wrappers.py +1038 -0
@@ -0,0 +1,12 @@
1
+ env
2
+ *__pycache__*
3
+
4
+ .venv
5
+ .vscode
6
+ .idea
7
+ *.bc
8
+ results/
9
+ .DS_Store
10
+ docs/build/
11
+ .python-version
12
+
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.