JAXtari 0.1.0__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.
- jaxatari/__init__.py +32 -0
- jaxatari/core.py +207 -0
- jaxatari/environment.py +289 -0
- jaxatari/games/__init__.py +0 -0
- jaxatari/games/amidar_mazes.py +23 -0
- jaxatari/games/jax_airraid.py +1168 -0
- jaxatari/games/jax_alien.py +3108 -0
- jaxatari/games/jax_amidar.py +1303 -0
- jaxatari/games/jax_asterix.py +1114 -0
- jaxatari/games/jax_asteroids.py +1420 -0
- jaxatari/games/jax_atlantis.py +1592 -0
- jaxatari/games/jax_bankheist.py +1644 -0
- jaxatari/games/jax_beamrider.py +4909 -0
- jaxatari/games/jax_berzerk.py +2258 -0
- jaxatari/games/jax_blackjack.py +1021 -0
- jaxatari/games/jax_breakout.py +1065 -0
- jaxatari/games/jax_casino.py +355 -0
- jaxatari/games/jax_casino_blackjack.py +1174 -0
- jaxatari/games/jax_casino_five_stud_poker.py +499 -0
- jaxatari/games/jax_casino_poker_solitaire.py +421 -0
- jaxatari/games/jax_centipede.py +2569 -0
- jaxatari/games/jax_choppercommand.py +2211 -0
- jaxatari/games/jax_donkeykong.py +2532 -0
- jaxatari/games/jax_enduro.py +1769 -0
- jaxatari/games/jax_fishingderby.py +1879 -0
- jaxatari/games/jax_flagcapture.py +718 -0
- jaxatari/games/jax_freeway.py +723 -0
- jaxatari/games/jax_frostbite.py +3742 -0
- jaxatari/games/jax_galaxian.py +1779 -0
- jaxatari/games/jax_gravitar.py +4124 -0
- jaxatari/games/jax_hangman.py +767 -0
- jaxatari/games/jax_hauntedhouse.py +1524 -0
- jaxatari/games/jax_humancannonball.py +1022 -0
- jaxatari/games/jax_kangaroo.py +2367 -0
- jaxatari/games/jax_kingkong.py +2631 -0
- jaxatari/games/jax_klax.py +1161 -0
- jaxatari/games/jax_lasergates.py +3511 -0
- jaxatari/games/jax_montezumarevenge.py +1016 -0
- jaxatari/games/jax_mspacman.py +1687 -0
- jaxatari/games/jax_namethisgame.py +1649 -0
- jaxatari/games/jax_pacman.py +1301 -0
- jaxatari/games/jax_phoenix.py +2883 -0
- jaxatari/games/jax_pong.py +590 -0
- jaxatari/games/jax_qbert.py +1581 -0
- jaxatari/games/jax_riverraid.py +2209 -0
- jaxatari/games/jax_seaquest.py +2856 -0
- jaxatari/games/jax_sirlancelot.py +2911 -0
- jaxatari/games/jax_skiing.py +1242 -0
- jaxatari/games/jax_slotmachine.py +1474 -0
- jaxatari/games/jax_spaceinvaders.py +1427 -0
- jaxatari/games/jax_spacewar.py +1128 -0
- jaxatari/games/jax_surround.py +765 -0
- jaxatari/games/jax_tennis.py +1737 -0
- jaxatari/games/jax_tetris.py +787 -0
- jaxatari/games/jax_timepilot.py +1934 -0
- jaxatari/games/jax_tron.py +2927 -0
- jaxatari/games/jax_turmoil.py +2145 -0
- jaxatari/games/jax_venture.py +2136 -0
- jaxatari/games/jax_videocheckers.py +1461 -0
- jaxatari/games/jax_videocube.py +1396 -0
- jaxatari/games/jax_videopinball.py +4381 -0
- jaxatari/games/jax_wordzapper.py +2115 -0
- jaxatari/games/kangaroo_levels.py +249 -0
- jaxatari/games/mods/__init__.py +0 -0
- jaxatari/games/mods/alien/alien_mod_plugins.py +180 -0
- jaxatari/games/mods/alien_mods.py +39 -0
- jaxatari/games/mods/asteroids/__init__.py +0 -0
- jaxatari/games/mods/asteroids/asteroids_mod_plugins.py +345 -0
- jaxatari/games/mods/asteroids_mods.py +30 -0
- jaxatari/games/mods/atlantis/atlantis_mod_plugins.py +114 -0
- jaxatari/games/mods/atlantis_mods.py +39 -0
- jaxatari/games/mods/bankheist/bankheist_mod_plugins.py +668 -0
- jaxatari/games/mods/bankheist_mods.py +61 -0
- jaxatari/games/mods/beamrider/beamrider_mod_plugins.py +3623 -0
- jaxatari/games/mods/beamrider_mods.py +41 -0
- jaxatari/games/mods/breakout/breakout_mod_plugins.py +141 -0
- jaxatari/games/mods/breakout_mods.py +44 -0
- jaxatari/games/mods/enduro/enduro_mod_plugins.py +216 -0
- jaxatari/games/mods/enduro_mods.py +36 -0
- jaxatari/games/mods/fishingderby/fishingderby_mod_plugins.py +82 -0
- jaxatari/games/mods/fishingderby_mods.py +39 -0
- jaxatari/games/mods/freeway/freeway_mod_plugins.py +234 -0
- jaxatari/games/mods/freeway_mods.py +40 -0
- jaxatari/games/mods/frostbite/__init__.py +0 -0
- jaxatari/games/mods/frostbite/frostbite_mod_plugins.py +190 -0
- jaxatari/games/mods/frostbite_mods.py +48 -0
- jaxatari/games/mods/gravitar/gravitar_mod_plugins.py +166 -0
- jaxatari/games/mods/gravitar_mods.py +53 -0
- jaxatari/games/mods/kangaroo/kangaroo_mod_plugins.py +1189 -0
- jaxatari/games/mods/kangaroo/sprites/cactus.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/cactus_tall.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/chicken.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/danger_sign.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/dragon.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/fireball.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/flame_0.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/flame_1.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/honey_bee.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/kangaroo_rope_climb.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/polarbear.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/snake.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/tank.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/tank_15x8.npy +0 -0
- jaxatari/games/mods/kangaroo/sprites/wasp.npy +0 -0
- jaxatari/games/mods/kangaroo_mods.py +86 -0
- jaxatari/games/mods/montezuma_revenge/montezuma_revenge_mod_plugins.py +285 -0
- jaxatari/games/mods/montezuma_revenge_mods.py +50 -0
- jaxatari/games/mods/mspacman/mspacman_mod_plugins.py +364 -0
- jaxatari/games/mods/mspacman_mods.py +53 -0
- jaxatari/games/mods/pacman/pacman_mod_plugins.py +269 -0
- jaxatari/games/mods/pacman_mods.py +40 -0
- jaxatari/games/mods/phoenix/phoenix_mod_plugins.py +138 -0
- jaxatari/games/mods/phoenix_mods.py +48 -0
- jaxatari/games/mods/pong/pong_mod_plugins.py +165 -0
- jaxatari/games/mods/pong_mods.py +36 -0
- jaxatari/games/mods/qbert/qbert_mod_plugins.py +673 -0
- jaxatari/games/mods/qbert_mods.py +55 -0
- jaxatari/games/mods/seaquest/seaquest_mod_plugins.py +123 -0
- jaxatari/games/mods/seaquest/sprites/fireball.npy +0 -0
- jaxatari/games/mods/seaquest/sprites/mine.npy +0 -0
- jaxatari/games/mods/seaquest_mods.py +115 -0
- jaxatari/games/mods/skiing/__init__.py +0 -0
- jaxatari/games/mods/skiing/skiing_mod_plugins.py +390 -0
- jaxatari/games/mods/skiing/sprites/blue_skier_fallen.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_0.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_1.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_2.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_3.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_4.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_5.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_6.npy +0 -0
- jaxatari/games/mods/skiing/sprites/blue_skiier_7.npy +0 -0
- jaxatari/games/mods/skiing_mods.py +53 -0
- jaxatari/games/mods/spaceinvaders/spaceinvaders_mod_plugins.py +108 -0
- jaxatari/games/mods/spaceinvaders_mods.py +37 -0
- jaxatari/games/mods/tennis/tennis_mod_plugins.py +376 -0
- jaxatari/games/mods/tennis_mods.py +54 -0
- jaxatari/games/mods/venture/venture_mod_plugins.py +148 -0
- jaxatari/games/mods/venture_mods.py +43 -0
- jaxatari/games/mods/videopinball/videopinball_mod_plugins.py +75 -0
- jaxatari/games/mods/videopinball_mods.py +31 -0
- jaxatari/games/montezuma_revenge/__init__.py +0 -0
- jaxatari/games/montezuma_revenge/core.py +224 -0
- jaxatari/games/montezuma_revenge/renderer.py +931 -0
- jaxatari/games/montezuma_revenge/rooms.py +1094 -0
- jaxatari/games/mspacman_mazes.py +285 -0
- jaxatari/games/timepilot_levels.py +177 -0
- jaxatari/games/videopinball_constants.py +1738 -0
- jaxatari/gym_wrapper.py +369 -0
- jaxatari/install_sprites.py +155 -0
- jaxatari/modification.py +1024 -0
- jaxatari/py.typed +0 -0
- jaxatari/renderers.py +15 -0
- jaxatari/rendering/__init__.py +0 -0
- jaxatari/rendering/jax_rendering_utils.py +1355 -0
- jaxatari/spaces.py +386 -0
- jaxatari/wrappers.py +1038 -0
- jaxtari-0.1.0.dist-info/METADATA +409 -0
- jaxtari-0.1.0.dist-info/RECORD +162 -0
- jaxtari-0.1.0.dist-info/WHEEL +4 -0
- jaxtari-0.1.0.dist-info/entry_points.txt +2 -0
- jaxtari-0.1.0.dist-info/licenses/LICENSE +21 -0
jaxatari/__init__.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
from platformdirs import user_data_dir
|
|
3
|
+
|
|
4
|
+
# 1. Define the path (Must match the installer script exactly)
|
|
5
|
+
# appname="jaxatari", appauthor="mycompany" (or whatever you used)
|
|
6
|
+
DATA_DIR = Path(user_data_dir("jaxatari"))
|
|
7
|
+
MARKER_FILE = DATA_DIR / ".ownership_confirmed"
|
|
8
|
+
ALT_SPRITES_MARKER_FILE = DATA_DIR / ".alternative_sprites_installed"
|
|
9
|
+
|
|
10
|
+
def check_ownership():
|
|
11
|
+
"""
|
|
12
|
+
Verifies that the user has accepted the license and confirmed ownership
|
|
13
|
+
of the original hardware/software by looking for the marker file.
|
|
14
|
+
"""
|
|
15
|
+
if not (MARKER_FILE.exists() or ALT_SPRITES_MARKER_FILE.exists()):
|
|
16
|
+
# Raise a clear, blocking error
|
|
17
|
+
raise RuntimeError(
|
|
18
|
+
"\n"
|
|
19
|
+
"❌ SPRITES NOT INSTALLED\n"
|
|
20
|
+
"----------------------------------------------------\n"
|
|
21
|
+
"JaxAtari needs sprite assets before environments can start.\n"
|
|
22
|
+
"You can either confirm your ownership of the original Atari 2600 ROMs and install sprites,\n"
|
|
23
|
+
"or continue with replacement/custom sprites.\n\n"
|
|
24
|
+
"Please run the following command in your terminal:\n\n"
|
|
25
|
+
" .venv/bin/install-sprites\n"
|
|
26
|
+
" or\n"
|
|
27
|
+
" python3 scripts/install_sprites.py\n"
|
|
28
|
+
"----------------------------------------------------\n"
|
|
29
|
+
)
|
|
30
|
+
|
|
31
|
+
# ... rest of your package imports ...
|
|
32
|
+
from jaxatari.core import make, list_available_games
|
jaxatari/core.py
ADDED
|
@@ -0,0 +1,207 @@
|
|
|
1
|
+
import importlib
|
|
2
|
+
import inspect
|
|
3
|
+
import warnings
|
|
4
|
+
|
|
5
|
+
from jaxatari.environment import JaxEnvironment
|
|
6
|
+
from jaxatari.renderers import JAXGameRenderer
|
|
7
|
+
from jaxatari.modification import apply_modifications
|
|
8
|
+
from jaxatari.wrappers import JaxatariWrapper
|
|
9
|
+
|
|
10
|
+
from . import check_ownership
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _warn_deprecated_obs_to_flat_array(env: JaxEnvironment) -> None:
|
|
14
|
+
"""Warn if legacy obs_to_flat_array is present on the environment."""
|
|
15
|
+
if hasattr(env, "obs_to_flat_array") and callable(getattr(env, "obs_to_flat_array")):
|
|
16
|
+
warnings.warn(
|
|
17
|
+
"Environment exposes deprecated obs_to_flat_array(). "
|
|
18
|
+
"Observations should now be flax.struct.dataclasses using ObjectObservation "
|
|
19
|
+
"for objects or plain arrays for observations like lives, score, etc. "
|
|
20
|
+
"Depending on legacy obs_to_flat_array might lead to unforseen issues with wrappers.",
|
|
21
|
+
DeprecationWarning,
|
|
22
|
+
stacklevel=2,
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
# Map of game names to their module paths (commented out games are WIP and will be supported in the near future)
|
|
28
|
+
GAME_MODULES = {
|
|
29
|
+
"amidar": "jaxatari.games.jax_amidar",
|
|
30
|
+
"airraid": "jaxatari.games.jax_airraid",
|
|
31
|
+
"alien": "jaxatari.games.jax_alien",
|
|
32
|
+
"asterix": "jaxatari.games.jax_asterix",
|
|
33
|
+
"asteroids": "jaxatari.games.jax_asteroids",
|
|
34
|
+
"atlantis": "jaxatari.games.jax_atlantis",
|
|
35
|
+
"bankheist": "jaxatari.games.jax_bankheist",
|
|
36
|
+
"beamrider": "jaxatari.games.jax_beamrider",
|
|
37
|
+
"berzerk": "jaxatari.games.jax_berzerk",
|
|
38
|
+
"blackjack": "jaxatari.games.jax_blackjack",
|
|
39
|
+
"breakout": "jaxatari.games.jax_breakout",
|
|
40
|
+
"casino": "jaxatari.games.jax_casino",
|
|
41
|
+
"centipede": "jaxatari.games.jax_centipede",
|
|
42
|
+
"choppercommand": "jaxatari.games.jax_choppercommand",
|
|
43
|
+
"donkeykong": "jaxatari.games.jax_donkeykong",
|
|
44
|
+
"enduro": "jaxatari.games.jax_enduro",
|
|
45
|
+
"fishingderby": "jaxatari.games.jax_fishingderby",
|
|
46
|
+
"flagcapture": "jaxatari.games.jax_flagcapture",
|
|
47
|
+
"freeway": "jaxatari.games.jax_freeway",
|
|
48
|
+
"frostbite": "jaxatari.games.jax_frostbite",
|
|
49
|
+
"galaxian": "jaxatari.games.jax_galaxian",
|
|
50
|
+
"gravitar": "jaxatari.games.jax_gravitar",
|
|
51
|
+
"hangman": "jaxatari.games.jax_hangman",
|
|
52
|
+
"hauntedhouse": "jaxatari.games.jax_hauntedhouse",
|
|
53
|
+
"humancannonball": "jaxatari.games.jax_humancannonball",
|
|
54
|
+
"kangaroo": "jaxatari.games.jax_kangaroo",
|
|
55
|
+
"kingkong": "jaxatari.games.jax_kingkong",
|
|
56
|
+
"klax": "jaxatari.games.jax_klax",
|
|
57
|
+
"lasergates": "jaxatari.games.jax_lasergates",
|
|
58
|
+
"namethisgame": "jaxatari.games.jax_namethisgame",
|
|
59
|
+
"phoenix": "jaxatari.games.jax_phoenix",
|
|
60
|
+
"pong": "jaxatari.games.jax_pong",
|
|
61
|
+
"qbert": "jaxatari.games.jax_qbert",
|
|
62
|
+
"riverraid": "jaxatari.games.jax_riverraid",
|
|
63
|
+
"seaquest": "jaxatari.games.jax_seaquest",
|
|
64
|
+
"sirlancelot": "jaxatari.games.jax_sirlancelot",
|
|
65
|
+
"skiing": "jaxatari.games.jax_skiing",
|
|
66
|
+
"slotmachine": "jaxatari.games.jax_slotmachine",
|
|
67
|
+
"spaceinvaders": "jaxatari.games.jax_spaceinvaders",
|
|
68
|
+
"spacewar": "jaxatari.games.jax_spacewar",
|
|
69
|
+
"surround": "jaxatari.games.jax_surround",
|
|
70
|
+
"tennis": "jaxatari.games.jax_tennis",
|
|
71
|
+
"tetris": "jaxatari.games.jax_tetris",
|
|
72
|
+
"timepilot": "jaxatari.games.jax_timepilot",
|
|
73
|
+
"tron": "jaxatari.games.jax_tron",
|
|
74
|
+
"turmoil": "jaxatari.games.jax_turmoil",
|
|
75
|
+
"venture": "jaxatari.games.jax_venture",
|
|
76
|
+
"videocheckers": "jaxatari.games.jax_videocheckers",
|
|
77
|
+
"videocube": "jaxatari.games.jax_videocube",
|
|
78
|
+
"videopinball": "jaxatari.games.jax_videopinball",
|
|
79
|
+
"wordzapper": "jaxatari.games.jax_wordzapper",
|
|
80
|
+
"mspacman": "jaxatari.games.jax_mspacman",
|
|
81
|
+
"montezumarevenge": "jaxatari.games.jax_montezumarevenge",
|
|
82
|
+
"pacman": "jaxatari.games.jax_pacman",
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
# Mod modules registry: for each game, provide the Controller class path
|
|
86
|
+
MOD_MODULES = {
|
|
87
|
+
"pong": "jaxatari.games.mods.pong_mods.PongEnvMod",
|
|
88
|
+
"kangaroo": "jaxatari.games.mods.kangaroo_mods.KangarooEnvMod",
|
|
89
|
+
"freeway": "jaxatari.games.mods.freeway_mods.FreewayEnvMod",
|
|
90
|
+
"breakout": "jaxatari.games.mods.breakout_mods.BreakoutEnvMod",
|
|
91
|
+
"seaquest": "jaxatari.games.mods.seaquest_mods.SeaquestEnvMod",
|
|
92
|
+
"videopinball": "jaxatari.games.mods.videopinball_mods.VideoPinballEnvMod",
|
|
93
|
+
'tennis': "jaxatari.games.mods.tennis_mods.TennisEnvMod",
|
|
94
|
+
"fishingderby": "jaxatari.games.mods.fishingderby_mods.FishingDerbyEnvMod",
|
|
95
|
+
"atlantis": "jaxatari.games.mods.atlantis_mods.AtlantisEnvMod",
|
|
96
|
+
"bankheist": "jaxatari.games.mods.bankheist_mods.BankHeistEnvMod",
|
|
97
|
+
"montezumarevenge": "jaxatari.games.mods.montezuma_revenge_mods.MontezumaRevengeEnvMod",
|
|
98
|
+
"frostbite": "jaxatari.games.mods.frostbite_mods.FrostbiteEnvMod",
|
|
99
|
+
"gravitar": "jaxatari.games.mods.gravitar_mods.GravitarEnvMod",
|
|
100
|
+
"phoenix": "jaxatari.games.mods.phoenix_mods.PhoenixEnvMod",
|
|
101
|
+
"enduro": "jaxatari.games.mods.enduro_mods.EnduroEnvMod",
|
|
102
|
+
"qbert": "jaxatari.games.mods.qbert_mods.QbertEnvMod",
|
|
103
|
+
"mspacman": "jaxatari.games.mods.mspacman_mods.MsPacmanEnvMod",
|
|
104
|
+
"beamrider": "jaxatari.games.mods.beamrider_mods.BeamRiderEnvMod",
|
|
105
|
+
"venture": "jaxatari.games.mods.venture_mods.VentureEnvMod",
|
|
106
|
+
"spaceinvaders": "jaxatari.games.mods.spaceinvaders_mods.SpaceInvadersEnvMod",
|
|
107
|
+
"skiing": "jaxatari.games.mods.skiing_mods.SkiingEnvMod",
|
|
108
|
+
"alien": "jaxatari.games.mods.alien_mods.AlienEnvMod",
|
|
109
|
+
"asteroids": "jaxatari.games.mods.asteroids_mods.AsteroidsEnvMod"
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def list_available_games() -> list[str]:
|
|
114
|
+
"""Lists all available, registered games."""
|
|
115
|
+
return list(GAME_MODULES.keys())
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def make(game_name: str,
|
|
119
|
+
mode: int = 0,
|
|
120
|
+
difficulty: int = 0,
|
|
121
|
+
mods_config: list = None, # deprecated, output warning if its used
|
|
122
|
+
mods: list = None,
|
|
123
|
+
allow_conflicts: bool = False
|
|
124
|
+
) -> JaxEnvironment:
|
|
125
|
+
"""
|
|
126
|
+
Creates and returns a JaxAtari game environment instance.
|
|
127
|
+
This is the main entry point for creating environments.
|
|
128
|
+
|
|
129
|
+
If 'mods' is provided, this function applies the
|
|
130
|
+
full two-stage modding pipeline:
|
|
131
|
+
1. Pre-scans for constant overrides.
|
|
132
|
+
2. Instantiates the base env with modded constants.
|
|
133
|
+
3. Applies the internal 'JaxAtariModController'.
|
|
134
|
+
4. Wraps the env with the 'JaxAtariModWrapper'.
|
|
135
|
+
|
|
136
|
+
Args:
|
|
137
|
+
game_name: Name of the game to load (e.g., "pong").
|
|
138
|
+
mode: Game mode.
|
|
139
|
+
difficulty: Game difficulty.
|
|
140
|
+
mods: List of modifications to apply (default: None).
|
|
141
|
+
allow_conflicts: Whether to allow conflicting mods (default: False).
|
|
142
|
+
Returns:
|
|
143
|
+
An instance of the specified game environment.
|
|
144
|
+
"""
|
|
145
|
+
|
|
146
|
+
check_ownership() # Ensure ownership confirmed
|
|
147
|
+
|
|
148
|
+
if isinstance(game_name, str):
|
|
149
|
+
game_name_clean = game_name.lower().replace("_", "").replace("-", "")
|
|
150
|
+
for key in GAME_MODULES:
|
|
151
|
+
if key.lower().replace("_", "").replace("-", "") == game_name_clean:
|
|
152
|
+
game_name = key
|
|
153
|
+
break
|
|
154
|
+
|
|
155
|
+
if mods_config is not None:
|
|
156
|
+
warnings.warn(
|
|
157
|
+
"'mods_config' is deprecated and will be removed in future versions. "
|
|
158
|
+
"Please use 'mods' instead.",
|
|
159
|
+
DeprecationWarning
|
|
160
|
+
)
|
|
161
|
+
mods = mods_config
|
|
162
|
+
|
|
163
|
+
if game_name not in GAME_MODULES:
|
|
164
|
+
raise NotImplementedError(
|
|
165
|
+
f"The game '{game_name}' does not exist. Available games: {list_available_games()}"
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
try:
|
|
169
|
+
# 1. Load the base environment class
|
|
170
|
+
module = importlib.import_module(GAME_MODULES[game_name])
|
|
171
|
+
env_class = None
|
|
172
|
+
for _, obj in inspect.getmembers(module):
|
|
173
|
+
if inspect.isclass(obj) and issubclass(obj, JaxEnvironment) and obj is not JaxEnvironment:
|
|
174
|
+
env_class = obj
|
|
175
|
+
break
|
|
176
|
+
if env_class is None:
|
|
177
|
+
raise ImportError(f"No JaxEnvironment subclass found in {GAME_MODULES[game_name]}")
|
|
178
|
+
|
|
179
|
+
# 2. Mods need default consts for pre-scan; otherwise a single env_class() is enough.
|
|
180
|
+
if mods:
|
|
181
|
+
try:
|
|
182
|
+
base_consts = env_class().consts
|
|
183
|
+
env = apply_modifications(
|
|
184
|
+
game_name=game_name,
|
|
185
|
+
mods_config=mods,
|
|
186
|
+
allow_conflicts=allow_conflicts,
|
|
187
|
+
base_consts=base_consts,
|
|
188
|
+
env_class=env_class,
|
|
189
|
+
MOD_MODULES=MOD_MODULES
|
|
190
|
+
)
|
|
191
|
+
_warn_deprecated_obs_to_flat_array(env)
|
|
192
|
+
return env
|
|
193
|
+
except NotImplementedError as e:
|
|
194
|
+
# Mod module not defined for this game - fall back to base environment
|
|
195
|
+
warnings.warn(
|
|
196
|
+
f"Mods requested for '{game_name}' but no mod module is available. "
|
|
197
|
+
f"Creating base environment without mods. Error: {e}",
|
|
198
|
+
UserWarning
|
|
199
|
+
)
|
|
200
|
+
|
|
201
|
+
env = env_class()
|
|
202
|
+
_warn_deprecated_obs_to_flat_array(env)
|
|
203
|
+
return env
|
|
204
|
+
|
|
205
|
+
except (ImportError, NotImplementedError) as e:
|
|
206
|
+
# Only wrap registration/import errors - let intentional errors (ValueError, etc.) propagate
|
|
207
|
+
raise ImportError(f"Failed to load game '{game_name}': {e}") from e
|
jaxatari/environment.py
ADDED
|
@@ -0,0 +1,289 @@
|
|
|
1
|
+
from enum import Enum
|
|
2
|
+
from typing import Tuple, Generic, TypeVar
|
|
3
|
+
import jax.numpy as jnp
|
|
4
|
+
import jax.random as jrandom
|
|
5
|
+
import warnings
|
|
6
|
+
from jaxatari.spaces import Space
|
|
7
|
+
from flax import struct
|
|
8
|
+
|
|
9
|
+
EnvObs = TypeVar("EnvObs")
|
|
10
|
+
EnvState = TypeVar("EnvState")
|
|
11
|
+
EnvInfo = TypeVar("EnvInfo")
|
|
12
|
+
EnvConstants = TypeVar("EnvConstants")
|
|
13
|
+
|
|
14
|
+
class JAXAtariAction:
|
|
15
|
+
"""
|
|
16
|
+
"Namespace" for Atari action integer constants.
|
|
17
|
+
These are directly usable in JAX arrays.
|
|
18
|
+
"""
|
|
19
|
+
NOOP: int = 0
|
|
20
|
+
FIRE: int = 1
|
|
21
|
+
UP: int = 2
|
|
22
|
+
RIGHT: int = 3
|
|
23
|
+
LEFT: int = 4
|
|
24
|
+
DOWN: int = 5
|
|
25
|
+
UPRIGHT: int = 6
|
|
26
|
+
UPLEFT: int = 7
|
|
27
|
+
DOWNRIGHT: int = 8
|
|
28
|
+
DOWNLEFT: int = 9
|
|
29
|
+
UPFIRE: int = 10
|
|
30
|
+
RIGHTFIRE: int = 11
|
|
31
|
+
LEFTFIRE: int = 12
|
|
32
|
+
DOWNFIRE: int = 13
|
|
33
|
+
UPRIGHTFIRE: int = 14
|
|
34
|
+
UPLEFTFIRE: int = 15
|
|
35
|
+
DOWNRIGHTFIRE: int = 16
|
|
36
|
+
DOWNLEFTFIRE: int = 17
|
|
37
|
+
|
|
38
|
+
@classmethod
|
|
39
|
+
def get_all_values(cls) -> jnp.ndarray:
|
|
40
|
+
# For fixed action sets, explicit listing is safest and clearest.
|
|
41
|
+
return jnp.array([
|
|
42
|
+
cls.NOOP, cls.FIRE, cls.UP, cls.RIGHT, cls.LEFT, cls.DOWN,
|
|
43
|
+
cls.UPRIGHT, cls.UPLEFT, cls.DOWNRIGHT, cls.DOWNLEFT,
|
|
44
|
+
cls.UPFIRE, cls.RIGHTFIRE, cls.LEFTFIRE, cls.DOWNFIRE,
|
|
45
|
+
cls.UPRIGHTFIRE, cls.UPLEFTFIRE, cls.DOWNRIGHTFIRE, cls.DOWNLEFTFIRE
|
|
46
|
+
], dtype=jnp.int32)
|
|
47
|
+
|
|
48
|
+
@struct.dataclass
|
|
49
|
+
class ObjectObservation:
|
|
50
|
+
"""
|
|
51
|
+
Dataclass for object centric observations of objects in jaxatari environments.
|
|
52
|
+
Can hold 1 to N objects of the same type (for example 12 sharks in seaquest or 1 player ship in asteroids).
|
|
53
|
+
Should always be instantiated via the create() classmethod to ensure proper default handling.
|
|
54
|
+
Attributes:
|
|
55
|
+
x: x position of the object.
|
|
56
|
+
y: y position of the object.
|
|
57
|
+
width: width of the object.
|
|
58
|
+
height: height of the object.
|
|
59
|
+
active: whether the object is currently active.
|
|
60
|
+
"""
|
|
61
|
+
x: jnp.ndarray # obligatory (int8)
|
|
62
|
+
y: jnp.ndarray # obligatory (int8)
|
|
63
|
+
width: jnp.ndarray # obligatory (int8)
|
|
64
|
+
height: jnp.ndarray # obligatory (int8)
|
|
65
|
+
|
|
66
|
+
# --- Additional attributes (will be set to 0 if not used) ---
|
|
67
|
+
active: jnp.ndarray = struct.field(default_factory=lambda: jnp.array(1)) # whether the object is currently active (0 or 1)
|
|
68
|
+
visual_id: jnp.ndarray = struct.field(default_factory=lambda: jnp.array(0)) # visual identifier of the object (different color sprites for example)
|
|
69
|
+
state: jnp.ndarray = struct.field(default_factory=lambda: jnp.array(0)) # state of the object, for example is the ghost in pacman vulnerable [blinking] or not [static] (format depends on game, see the game docs)
|
|
70
|
+
orientation: jnp.ndarray = struct.field(default_factory=lambda: jnp.array(0)) # angle of the object (format depends on game, see the game docs)
|
|
71
|
+
|
|
72
|
+
@classmethod
|
|
73
|
+
def create(cls, x, y, width, height, active=None, visual_id=None, state=None, orientation=None):
|
|
74
|
+
# Helper to handle defaults
|
|
75
|
+
if active is None: active = jnp.ones_like(x, dtype=jnp.int32)
|
|
76
|
+
if visual_id is None: visual_id = jnp.zeros_like(x, dtype=jnp.int32)
|
|
77
|
+
if state is None: state = jnp.zeros_like(x, dtype=jnp.int32)
|
|
78
|
+
if orientation is None: orientation = jnp.zeros_like(x, dtype=jnp.int32)
|
|
79
|
+
return cls(x=x, y=y, width=width, height=height, active=active, visual_id=visual_id, state=state, orientation=orientation)
|
|
80
|
+
|
|
81
|
+
def __repr__(self):
|
|
82
|
+
try:
|
|
83
|
+
# Handle scalar case (0-d arrays)
|
|
84
|
+
if self.x.ndim == 0:
|
|
85
|
+
try:
|
|
86
|
+
# Try to get concrete values for cleaner output
|
|
87
|
+
x, y = int(self.x), int(self.y)
|
|
88
|
+
w, h = int(self.width), int(self.height)
|
|
89
|
+
act = int(self.active)
|
|
90
|
+
ori = float(self.orientation)
|
|
91
|
+
st = int(self.state)
|
|
92
|
+
vid = int(self.visual_id)
|
|
93
|
+
status = "ACTIVE" if act else "INACTIVE"
|
|
94
|
+
return (f"Object(Single, {status}): Pos=({x}, {y}) | Size=({w}, {h}) | "
|
|
95
|
+
f"Ori={ori:.1f} | State={st} | VisID={vid}")
|
|
96
|
+
except:
|
|
97
|
+
# Fallback for Tracers
|
|
98
|
+
return f"Object(Single): Pos=({self.x}, {self.y}) | Active={self.active}"
|
|
99
|
+
|
|
100
|
+
# Handle vector case (1-d arrays)
|
|
101
|
+
n = self.x.shape[0]
|
|
102
|
+
lines = [f"ObjectGroup(count={n}):"]
|
|
103
|
+
|
|
104
|
+
# Limit print length if too huge
|
|
105
|
+
limit = min(n, 20)
|
|
106
|
+
|
|
107
|
+
for i in range(limit):
|
|
108
|
+
try:
|
|
109
|
+
# Try to extract concrete values
|
|
110
|
+
act = int(self.active[i])
|
|
111
|
+
status = "ACTIVE" if act else " - " # Dim inactive ones
|
|
112
|
+
|
|
113
|
+
x, y = int(self.x[i]), int(self.y[i])
|
|
114
|
+
w, h = int(self.width[i]), int(self.height[i])
|
|
115
|
+
ori = float(self.orientation[i])
|
|
116
|
+
st = int(self.state[i])
|
|
117
|
+
vid = int(self.visual_id[i])
|
|
118
|
+
|
|
119
|
+
# Formatted table row
|
|
120
|
+
line = (f" [{i:2d}] {status} | Pos: ({x:3d}, {y:3d}) | Size: ({w:2d}, {h:2d}) | "
|
|
121
|
+
f"Ori: {ori:5.1f} | State: {st:2d} | VisID: {vid:2d}")
|
|
122
|
+
except:
|
|
123
|
+
# Fallback for Tracers
|
|
124
|
+
line = f" [{i}] Active={self.active[i]} | Pos=({self.x[i]}, {self.y[i]})"
|
|
125
|
+
|
|
126
|
+
lines.append(line)
|
|
127
|
+
|
|
128
|
+
if n > limit:
|
|
129
|
+
lines.append(f" ... ({n - limit} more objects) ...")
|
|
130
|
+
|
|
131
|
+
return "\n".join(lines)
|
|
132
|
+
except Exception as e:
|
|
133
|
+
return f"ObjectObservation(Error in __repr__: {e})"
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
class JaxEnvironment(Generic[EnvState, EnvObs, EnvInfo, EnvConstants]):
|
|
137
|
+
"""
|
|
138
|
+
Abstract class for a JAX environment.
|
|
139
|
+
Generics:
|
|
140
|
+
EnvState: The type of the environment state.
|
|
141
|
+
EnvObs: The type of the observation.
|
|
142
|
+
EnvInfo: The type of the additional information.
|
|
143
|
+
EnvConstants: The type of the environment constants.
|
|
144
|
+
"""
|
|
145
|
+
|
|
146
|
+
def __init__(self, consts: EnvConstants = None):
|
|
147
|
+
if consts is not None:
|
|
148
|
+
# Check for legacy NamedTuple usage (has _fields but is not a PyTreeNode)
|
|
149
|
+
is_named_tuple = isinstance(consts, tuple) and hasattr(consts, '_fields')
|
|
150
|
+
# Check if it's a Flax PyTreeNode (flax.struct.dataclass instances)
|
|
151
|
+
try:
|
|
152
|
+
from flax import struct
|
|
153
|
+
is_flax_node = isinstance(consts, struct.PyTreeNode)
|
|
154
|
+
except (ImportError, AttributeError):
|
|
155
|
+
is_flax_node = False
|
|
156
|
+
|
|
157
|
+
if is_named_tuple and not is_flax_node:
|
|
158
|
+
warnings.warn(
|
|
159
|
+
f"Performance Warning: {self.__class__.__name__}.consts is a 'NamedTuple'. "
|
|
160
|
+
"This prevents JAX from treating constants as static metadata, potentially causing excessive recompilation. "
|
|
161
|
+
"Future versions will require 'flax.struct.PyTreeNode' (and the states/observations/info to flax.struct.dataclass/PyTreeNode). "
|
|
162
|
+
"Please refactor your constants class.",
|
|
163
|
+
UserWarning,
|
|
164
|
+
stacklevel=2
|
|
165
|
+
)
|
|
166
|
+
|
|
167
|
+
self.consts = consts
|
|
168
|
+
|
|
169
|
+
# --- MODDING INFRASTRUCTURE ---
|
|
170
|
+
# Functional: Tracks which renderer methods mods have patched.
|
|
171
|
+
# Used by wrappers to safely transfer patches during renderer swaps.
|
|
172
|
+
self._patched_renderer_methods = []
|
|
173
|
+
|
|
174
|
+
# Functional: Explicit registry of jitted callables that must be invalidated
|
|
175
|
+
# when renderer hot-swaps occur (e.g., native downscaling).
|
|
176
|
+
self._jit_invalidation_targets = []
|
|
177
|
+
# Functional: mutation epoch + tripwire controls for detecting risky
|
|
178
|
+
# post-trace monkeypatching.
|
|
179
|
+
self._jit_mutation_epoch = 0
|
|
180
|
+
self._jit_tripwire_enabled = True
|
|
181
|
+
|
|
182
|
+
# Informational: Structured audit log of every change made by the mod system.
|
|
183
|
+
# Machine-parseable: dict of category -> set of names that were changed.
|
|
184
|
+
# Categories: "attribute", "method", "constant", "asset".
|
|
185
|
+
self._mod_history = {
|
|
186
|
+
"attribute": set(),
|
|
187
|
+
"method": set(),
|
|
188
|
+
"constant": set(),
|
|
189
|
+
"asset": set(),
|
|
190
|
+
}
|
|
191
|
+
|
|
192
|
+
def reset(self, key: jrandom.PRNGKey=None) -> Tuple[EnvObs, EnvState]:
|
|
193
|
+
"""
|
|
194
|
+
Resets the environment to the initial state.
|
|
195
|
+
Returns: The initial observation and the initial environment state.
|
|
196
|
+
|
|
197
|
+
"""
|
|
198
|
+
raise NotImplementedError("Abstract method")
|
|
199
|
+
|
|
200
|
+
def step(
|
|
201
|
+
self, state: EnvState, action
|
|
202
|
+
) -> Tuple[EnvObs, EnvState, float, bool, EnvInfo]:
|
|
203
|
+
"""
|
|
204
|
+
Takes a step in the environment.
|
|
205
|
+
Args:
|
|
206
|
+
state: The current environment state.
|
|
207
|
+
action: The action to take.
|
|
208
|
+
|
|
209
|
+
Returns: The observation, the new environment state, the reward, whether the state is terminal, and additional info.
|
|
210
|
+
|
|
211
|
+
"""
|
|
212
|
+
raise NotImplementedError("Abstract method")
|
|
213
|
+
|
|
214
|
+
def render(self, state: EnvState) -> Tuple[jnp.ndarray]:
|
|
215
|
+
"""
|
|
216
|
+
Renders the environment state to a single image.
|
|
217
|
+
Args:
|
|
218
|
+
state: The environment state.
|
|
219
|
+
|
|
220
|
+
Returns: A single image of the environment state.
|
|
221
|
+
|
|
222
|
+
"""
|
|
223
|
+
raise NotImplementedError("Abstract method")
|
|
224
|
+
|
|
225
|
+
def action_space(self) -> Space:
|
|
226
|
+
"""
|
|
227
|
+
Returns the action space of the environment as an array containing the actions that can be taken.
|
|
228
|
+
Returns: The action space of the environment as an array.
|
|
229
|
+
"""
|
|
230
|
+
raise NotImplementedError("Abstract method")
|
|
231
|
+
|
|
232
|
+
def observation_space(self) -> Space:
|
|
233
|
+
"""
|
|
234
|
+
Returns the observation space of the environment.
|
|
235
|
+
Returns: The observation space of the environment.
|
|
236
|
+
"""
|
|
237
|
+
raise NotImplementedError("Abstract method")
|
|
238
|
+
|
|
239
|
+
def image_space(self) -> Space:
|
|
240
|
+
"""
|
|
241
|
+
Returns the image space of the environment.
|
|
242
|
+
Returns: The image space of the environment.
|
|
243
|
+
"""
|
|
244
|
+
raise NotImplementedError("Abstract method")
|
|
245
|
+
|
|
246
|
+
def _get_observation(self, state: EnvState) -> EnvObs:
|
|
247
|
+
"""
|
|
248
|
+
Converts the environment state to the observation by filtering out non-relevant information.
|
|
249
|
+
Args:
|
|
250
|
+
state: The environment state.
|
|
251
|
+
|
|
252
|
+
Returns: observation
|
|
253
|
+
|
|
254
|
+
"""
|
|
255
|
+
raise NotImplementedError("Abstract method")
|
|
256
|
+
|
|
257
|
+
def _get_info(self, state: EnvState, all_rewards: jnp.array = None) -> EnvInfo:
|
|
258
|
+
"""
|
|
259
|
+
Extracts information from the environment state that is not relevant for the agent.
|
|
260
|
+
Args:
|
|
261
|
+
state: The environment state.
|
|
262
|
+
|
|
263
|
+
Returns: info
|
|
264
|
+
|
|
265
|
+
"""
|
|
266
|
+
raise NotImplementedError("Abstract method")
|
|
267
|
+
|
|
268
|
+
def _get_reward(self, previous_state: EnvState, state: EnvState) -> float:
|
|
269
|
+
"""
|
|
270
|
+
Calculates the reward from the environment state.
|
|
271
|
+
Args:
|
|
272
|
+
previous_state: The previous environment state.
|
|
273
|
+
state: The environment state.
|
|
274
|
+
|
|
275
|
+
Returns: reward
|
|
276
|
+
|
|
277
|
+
"""
|
|
278
|
+
raise NotImplementedError("Abstract method")
|
|
279
|
+
|
|
280
|
+
def _get_done(self, state: EnvState) -> bool:
|
|
281
|
+
"""
|
|
282
|
+
Determines if the environment state is a terminal state
|
|
283
|
+
Args:
|
|
284
|
+
state: The environment state.
|
|
285
|
+
|
|
286
|
+
Returns: True if the state is terminal, False otherwise.
|
|
287
|
+
|
|
288
|
+
"""
|
|
289
|
+
raise NotImplementedError("Abstract method")
|
|
File without changes
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
# Original amidar maze
|
|
2
|
+
import jax.numpy as jnp
|
|
3
|
+
|
|
4
|
+
WIDTH = 160
|
|
5
|
+
HEIGHT = 210
|
|
6
|
+
PATH_THICKNESS_HORIZONTAL = 5
|
|
7
|
+
PATH_THICKNESS_VERTICAL = 4
|
|
8
|
+
MAX_ENEMIES = 6
|
|
9
|
+
|
|
10
|
+
PATH_CORNERS = jnp.array([[16, 14], [40, 14], [56, 14], [72, 14], [84, 14], [100, 14], [116, 14], [140, 14], [16, 44], [32, 44], [40, 44], [52, 44], [56, 44], [64, 44], [72, 44], [84, 44], [92, 44], [100, 44], [104, 44], [116, 44], [124, 44], [140, 44], [16, 74], [28, 74], [32, 74], [52, 74], [60, 74], [64, 74], [92, 74], [96, 74], [104, 74], [124, 74], [128, 74], [140, 74], [16, 104], [28, 104], [36, 104], [60, 104], [72, 104], [84, 104], [96, 104], [120, 104], [128, 104], [140, 104], [16, 134], [36, 134], [40, 134], [64, 134], [72, 134], [84, 134], [92, 134], [116, 134], [120, 134], [140, 134], [16, 164], [40, 164], [64, 164], [92, 164], [116, 164], [140, 164]], dtype=jnp.int32)
|
|
11
|
+
HORIZONTAL_PATH_EDGES = jnp.array([[[16, 14], [40, 14]], [[16, 44], [32, 44]], [[16, 74], [28, 74]], [[16, 104], [28, 104]], [[16, 134], [36, 134]], [[16, 164], [40, 164]], [[28, 74], [32, 74]], [[28, 104], [36, 104]], [[32, 44], [40, 44]], [[32, 74], [52, 74]], [[36, 104], [60, 104]], [[36, 134], [40, 134]], [[40, 14], [56, 14]], [[40, 44], [52, 44]], [[40, 134], [64, 134]], [[40, 164], [64, 164]], [[52, 44], [56, 44]], [[52, 74], [60, 74]], [[56, 14], [72, 14]], [[56, 44], [64, 44]], [[60, 74], [64, 74]], [[60, 104], [72, 104]], [[64, 44], [72, 44]], [[64, 74], [92, 74]], [[64, 134], [72, 134]], [[64, 164], [92, 164]], [[72, 14], [84, 14]], [[72, 44], [84, 44]], [[72, 104], [84, 104]], [[72, 134], [84, 134]], [[84, 14], [100, 14]], [[84, 44], [92, 44]], [[84, 104], [96, 104]], [[84, 134], [92, 134]], [[92, 44], [100, 44]], [[92, 74], [96, 74]], [[92, 134], [116, 134]], [[92, 164], [116, 164]], [[96, 74], [104, 74]], [[96, 104], [120, 104]], [[100, 14], [116, 14]], [[100, 44], [104, 44]], [[104, 44], [116, 44]], [[104, 74], [124, 74]], [[116, 14], [140, 14]], [[116, 44], [124, 44]], [[116, 134], [120, 134]], [[116, 164], [140, 164]], [[120, 104], [128, 104]], [[120, 134], [140, 134]], [[124, 44], [140, 44]], [[124, 74], [128, 74]], [[128, 74], [140, 74]], [[128, 104], [140, 104]]], dtype=jnp.int32)
|
|
12
|
+
VERTICAL_PATH_EDGES = jnp.array([[[16, 14], [16, 44]], [[16, 44], [16, 74]], [[16, 74], [16, 104]], [[16, 104], [16, 134]], [[16, 134], [16, 164]], [[28, 74], [28, 104]], [[32, 44], [32, 74]], [[36, 104], [36, 134]], [[40, 14], [40, 44]], [[40, 134], [40, 164]], [[52, 44], [52, 74]], [[56, 14], [56, 44]], [[60, 74], [60, 104]], [[64, 44], [64, 74]], [[64, 134], [64, 164]], [[72, 14], [72, 44]], [[72, 104], [72, 134]], [[84, 14], [84, 44]], [[84, 104], [84, 134]], [[92, 44], [92, 74]], [[92, 134], [92, 164]], [[96, 74], [96, 104]], [[100, 14], [100, 44]], [[104, 44], [104, 74]], [[116, 14], [116, 44]], [[116, 134], [116, 164]], [[120, 104], [120, 134]], [[124, 44], [124, 74]], [[128, 74], [128, 104]], [[140, 14], [140, 44]], [[140, 44], [140, 74]], [[140, 74], [140, 104]], [[140, 104], [140, 134]], [[140, 134], [140, 164]]], dtype=jnp.int32)
|
|
13
|
+
PATH_EDGES = jnp.concatenate((HORIZONTAL_PATH_EDGES, VERTICAL_PATH_EDGES), axis=0)
|
|
14
|
+
RECTANGLES = jnp.array([[1, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0], [0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0], [0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 1, 1, 0, 1, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0], [0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1, 0], [0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0], [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 1]], dtype=jnp.int32)
|
|
15
|
+
RECTANGLE_BOUNDS = jnp.array([[16, 14, 40, 44], [40, 14, 56, 44], [56, 14, 72, 44], [72, 14, 84, 44], [84, 14, 100, 44], [100, 14, 116, 44], [116, 14, 140, 44], [16, 44, 32, 74], [32, 44, 52, 74], [52, 44, 64, 74], [64, 44, 92, 74], [92, 44, 104, 74], [104, 44, 124, 74], [124, 44, 140, 74], [16, 74, 28, 104], [28, 74, 60, 104], [60, 74, 96, 104], [96, 74, 128, 104], [128, 74, 140, 104], [16, 104, 36, 134], [36, 104, 72, 134], [72, 104, 84, 134], [84, 104, 120, 134], [120, 104, 140, 134], [16, 134, 40, 164], [40, 134, 64, 164], [64, 134, 92, 164], [92, 134, 116, 164], [116, 134, 140, 164]], dtype=jnp.int32)
|
|
16
|
+
CORNER_RECTANGLES = jnp.array([0, 6, 24, 28], dtype=jnp.int32)
|
|
17
|
+
SHORT_PATHS = jnp.array([[23, 24, 6], [45, 46, 11], [11, 12, 16], [26, 27, 20], [28, 29, 35], [17, 18, 41], [51, 52, 46], [31, 32, 51]], dtype=jnp.int32)
|
|
18
|
+
|
|
19
|
+
INITIAL_PLAYER_POSITION = jnp.array([140, 89], dtype=jnp.int32)
|
|
20
|
+
INITIAL_ENEMY_POSITIONS = jnp.array([[16, 14], [16, 14], [44, 14], [16, 137], [52, 164], [16, 164]], dtype=jnp.int32)
|
|
21
|
+
PLAYER_STARTING_PATH = jnp.array([85], dtype=jnp.int32)
|
|
22
|
+
|
|
23
|
+
# Import these from this file where needed instead of modifying jax_amidar.py.
|