diff --git a/config/tokenizer/default.yaml b/config/tokenizer/default.yaml index 54a4dd7..39fe6ae 100644 --- a/config/tokenizer/default.yaml +++ b/config/tokenizer/default.yaml @@ -17,4 +17,4 @@ encoder: dropout: 0.0 decoder: _target_: models.tokenizer.Decoder - config: ${..encoder.config} \ No newline at end of file + config: ${..encoder.config} diff --git a/config/trainer.yaml b/config/trainer.yaml index 9e14d17..fe220a7 100644 --- a/config/trainer.yaml +++ b/config/trainer.yaml @@ -54,20 +54,20 @@ training: should: True learning_rate: 0.0001 tokenizer: - batch_num_samples: 256 + batch_num_samples: 64 grad_acc_steps: 1 max_grad_norm: 10.0 start_after_epochs: 5 steps_per_epoch: 200 world_model: - batch_num_samples: 64 + batch_num_samples: 16 grad_acc_steps: 1 max_grad_norm: 10.0 weight_decay: 0.01 start_after_epochs: 25 steps_per_epoch: 200 actor_critic: - batch_num_samples: 64 + batch_num_samples: 16 grad_acc_steps: 1 max_grad_norm: 10.0 start_after_epochs: 50 diff --git a/poetry.lock b/poetry.lock index 9ad1a2b..cf1ee10 100644 --- a/poetry.lock +++ b/poetry.lock @@ -31,37 +31,32 @@ testing = ["bitsandbytes", "datasets", "deepspeed", "evaluate", "parameterized", [[package]] name = "ale-py" -version = "0.8.1" +version = "0.7.5" description = "The Arcade Learning Environment (ALE) - a platform for AI research." optional = false python-versions = ">=3.7" files = [ - {file = "ale_py-0.8.1-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:b2aa2f69a4169742800615970efe6914fa856e33eaf7fa9133c0e06a617a80e2"}, - {file = "ale_py-0.8.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:6f2f6b92c8fd6189654979bbf0b305dbe0ecf82176c47f244d8c1cbc36286b89"}, - {file = "ale_py-0.8.1-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:c9b168eb88c87d0f3e2a778e6c5cdde4ad951d1ca8a6dc3d3679fd45398df7d1"}, - {file = "ale_py-0.8.1-cp310-cp310-win_amd64.whl", hash = "sha256:5fcc31f495de79ee1d6bfc0f4b7c4619948851e679bbf010035e25f23146a687"}, - {file = "ale_py-0.8.1-cp311-cp311-macosx_10_15_x86_64.whl", hash = "sha256:0856ca777473ec4ae8a59f3af9580259adb0fd4a47d586a125a440c62e82fc10"}, - {file = "ale_py-0.8.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:f10b1df8774bbe3b00365748b5e0e07cf35f6a703bbaff991bc7b3b2247dccc9"}, - {file = "ale_py-0.8.1-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0006d80dfe7745eb5a93444492337203c8bc7eb594a2c24c6a651c5c5b0eaf09"}, - {file = "ale_py-0.8.1-cp311-cp311-win_amd64.whl", hash = "sha256:9773eea7505484e024beb2fff0f3bfd363db151bdb9799d70995448e196b1ded"}, - {file = "ale_py-0.8.1-cp37-cp37m-macosx_10_15_x86_64.whl", hash = "sha256:87557db05be0e04130e2ec1bf909d3bb0b0bc034645d4f664e6baa573fe32191"}, - {file = "ale_py-0.8.1-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:ae2ba24557e0ce541ea3be13b148db2a9cfa730d83537b4cbed5e10449826e51"}, - {file = "ale_py-0.8.1-cp37-cp37m-win_amd64.whl", hash = "sha256:ade5c32af567629164a6b49378978c728a15dc4db07ad6b679e8832d4fd3ea1f"}, - {file = "ale_py-0.8.1-cp38-cp38-macosx_10_15_x86_64.whl", hash = "sha256:0ffecb5c956749596030e464827642945162170a132d093c3d4fa2d7e5725c18"}, - {file = "ale_py-0.8.1-cp38-cp38-macosx_11_0_arm64.whl", hash = "sha256:7cd74b7ee0248ef11a086c9764e142e71defd40ec8989a99232bfd2d9e8023be"}, - {file = "ale_py-0.8.1-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:eadf9f3990b4ff2f9e5ca35889f5e2e95cddd6a353d9d857d9b4601a6e1c4e7c"}, - {file = "ale_py-0.8.1-cp38-cp38-win_amd64.whl", hash = "sha256:817adf9a3a82c4923c731e634520a5ecf296aca0367f5c69959a96b32119d831"}, - {file = "ale_py-0.8.1-cp39-cp39-macosx_10_15_x86_64.whl", hash = "sha256:2d9fcfa06c74a613c5419e942ef4d3e0959533f52e94d2d4bda61d07fbfffeee"}, - {file = "ale_py-0.8.1-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:f278036f9b6066062abcdf0987a0ec5a8e0f22a2c7cfac925e39378d4343d490"}, - {file = "ale_py-0.8.1-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b00f74e27815131c1a2791f3d48114363fa2708e19f09ce6b7b614cb14c9d469"}, - {file = "ale_py-0.8.1-cp39-cp39-win_amd64.whl", hash = "sha256:d49b550a2d9c25b63c343aa680fd81f253a3714cdc0e1835640933ebff1798ff"}, + {file = "ale_py-0.7.5-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:c68d3960f0e14738362069798ff5baa85ecb337038168cb61bd4855125deb894"}, + {file = "ale_py-0.7.5-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:003725a8b6f02353066564ab417f4228adc886434aa4c8d314489a7f79fe9dfd"}, + {file = "ale_py-0.7.5-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1a30270d4aec820527be7a643aad381d0d1c932dc476740f63cc649dd200e065"}, + {file = "ale_py-0.7.5-cp310-cp310-win_amd64.whl", hash = "sha256:1a72f8a311887fe3800964976c7ce3567dd701e7d9b77209d3a05c5cdf9c6af0"}, + {file = "ale_py-0.7.5-cp37-cp37m-macosx_10_15_x86_64.whl", hash = "sha256:910c5443a836bf0bc6b03c45fefeeb9096254f807d24e247753e214062c12cae"}, + {file = "ale_py-0.7.5-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1f97301b89a8b49b9d2cc034f8a7df55eb6b1930126fb6cb6c6d6239bab8a89d"}, + {file = "ale_py-0.7.5-cp37-cp37m-win_amd64.whl", hash = "sha256:38e9f6c2b1cca246253f41028b8ad92d725beec8d4672c0fe247a217a4c7dc4f"}, + {file = "ale_py-0.7.5-cp38-cp38-macosx_10_15_x86_64.whl", hash = "sha256:48c5abf2c59bed0b46f17ac60ec6c940e6e7d4abc6ff7e5296246286f80e4a23"}, + {file = "ale_py-0.7.5-cp38-cp38-macosx_11_0_arm64.whl", hash = "sha256:b2680b9a3c577865572dd9b8f24303189bb22c67b1b9668dcad3c20f59285e78"}, + {file = "ale_py-0.7.5-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:dec87873c6fe65c71129e20c52ce35dd4507ca409fa1233f4ae9c511144a9c99"}, + {file = "ale_py-0.7.5-cp38-cp38-win_amd64.whl", hash = "sha256:1f8694a091e13be45c8704a089462b1519f7112a85e1c1b61c842c69a2608a98"}, + {file = "ale_py-0.7.5-cp39-cp39-macosx_10_15_x86_64.whl", hash = "sha256:4c7e7fed3ce2b305d2db9ca2f87106e3df6da2c4fc2363103782ffd5975b4549"}, + {file = "ale_py-0.7.5-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:da5be311285a00e587b886805ea56762d44bf4f1da672229c1fbaff209bfcdc2"}, + {file = "ale_py-0.7.5-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e513799fbb5128fa045cfb9a31cffa7f7c188dbe885665188b6c4cb50e9d983c"}, + {file = "ale_py-0.7.5-cp39-cp39-win_amd64.whl", hash = "sha256:9976ceea0251ccad7e562e2021c337d5e9a9b24b5f9eecfe82573c2279337016"}, ] [package.dependencies] importlib-metadata = {version = ">=4.10.0", markers = "python_version < \"3.10\""} importlib-resources = "*" numpy = "*" -typing-extensions = {version = "*", markers = "python_version < \"3.11\""} [package.extras] test = ["gym (>=0.23,<1.0)", "pytest (>=7.0)"] @@ -701,32 +696,32 @@ test = ["black", "coverage[toml]", "ddt (>=1.1.1,!=1.4.3)", "mock", "mypy", "pre [[package]] name = "gym" -version = "0.26.2" -description = "Gym: A universal API for reinforcement learning environments" +version = "0.22.0" +description = "Gym: A universal API for reinforcement learning environments." optional = false -python-versions = ">=3.6" +python-versions = ">=3.7" files = [ - {file = "gym-0.26.2.tar.gz", hash = "sha256:e0d882f4b54f0c65f203104c24ab8a38b039f1289986803c7d02cdbe214fbcc4"}, + {file = "gym-0.22.0.tar.gz", hash = "sha256:339144c89951758187c378111919bc0e2f1695f9e9d9699e3f19279a6398148d"}, ] [package.dependencies] +ale-py = {version = ">=0.7.4,<0.8.0", optional = true, markers = "extra == \"atari\""} autorom = {version = ">=0.4.2,<0.5.0", extras = ["accept-rom-license"], optional = true, markers = "extra == \"accept-rom-license\""} cloudpickle = ">=1.2.0" gym_notices = ">=0.0.4" -importlib_metadata = {version = ">=4.8.0", markers = "python_version < \"3.10\""} +importlib_metadata = {version = ">=4.10.0", markers = "python_version < \"3.10\""} numpy = ">=1.18.0" [package.extras] accept-rom-license = ["autorom[accept-rom-license] (>=0.4.2,<0.5.0)"] -all = ["ale-py (>=0.8.0,<0.9.0)", "box2d-py (==2.3.5)", "imageio (>=2.14.1)", "lz4 (>=3.1.0)", "matplotlib (>=3.0)", "moviepy (>=1.0.0)", "mujoco (==2.2)", "mujoco_py (>=2.1,<2.2)", "opencv-python (>=3.0)", "pygame (==2.1.0)", "pytest (==7.0.1)", "swig (==4.*)"] -atari = ["ale-py (>=0.8.0,<0.9.0)"] -box2d = ["box2d-py (==2.3.5)", "pygame (==2.1.0)", "swig (==4.*)"] +all = ["ale-py (>=0.7.4,<0.8.0)", "box2d-py (==2.3.5)", "box2d-py (==2.3.5)", "lz4 (>=3.1.0)", "lz4 (>=3.1.0)", "mujoco_py (>=1.50,<2.0)", "opencv-python (>=3.0)", "opencv-python (>=3.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "scipy (>=1.4.1)", "scipy (>=1.4.1)"] +atari = ["ale-py (>=0.7.4,<0.8.0)"] +box2d = ["box2d-py (==2.3.5)", "pygame (==2.1.0)"] classic-control = ["pygame (==2.1.0)"] -mujoco = ["imageio (>=2.14.1)", "mujoco (==2.2)"] -mujoco-py = ["mujoco_py (>=2.1,<2.2)"] -other = ["lz4 (>=3.1.0)", "matplotlib (>=3.0)", "moviepy (>=1.0.0)", "opencv-python (>=3.0)"] -testing = ["box2d-py (==2.3.5)", "imageio (>=2.14.1)", "lz4 (>=3.1.0)", "matplotlib (>=3.0)", "moviepy (>=1.0.0)", "mujoco (==2.2)", "mujoco_py (>=2.1,<2.2)", "opencv-python (>=3.0)", "pygame (==2.1.0)", "pytest (==7.0.1)", "swig (==4.*)"] -toy-text = ["pygame (==2.1.0)"] +mujoco = ["mujoco_py (>=1.50,<2.0)"] +nomujoco = ["box2d-py (==2.3.5)", "lz4 (>=3.1.0)", "opencv-python (>=3.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "pygame (==2.1.0)", "scipy (>=1.4.1)"] +other = ["lz4 (>=3.1.0)", "opencv-python (>=3.0)"] +toy-text = ["pygame (==2.1.0)", "scipy (>=1.4.1)"] [[package]] name = "gym-notices" @@ -2903,4 +2898,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p [metadata] lock-version = "2.0" python-versions = ">=3.9,<3.13" -content-hash = "914e609489e134bb4b844f424ea35ef15e1a6288e634ad4046889785123a1e1f" +content-hash = "3fb1a61c2bf4b6fa530dff25509331cb4a0dfe863f77355d01e886feb06c9b34" diff --git a/pyproject.toml b/pyproject.toml index 8208ce7..fd1121a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -18,14 +18,15 @@ bitsandbytes = {url = "https://github.com/TimDettmers/bitsandbytes/releases/down transformers = "4.34.0" tqdm = "^4.66.1" wandb = "^0.12.6" -ale-py = "^0.8.1" +ale-py = "^0.7.4" pygame = "^2.5.2" psutil = "^5.9.6" protobuf = "^3.10.0" opencv-python = "^4.8.1.78" hydra-core = "^1.3.2" -gym = {extras = ["accept-rom-license"], version = "^0.26.2"} torchvision = "^0.16.0" +numpy = ">=1.18.0" +gym = {version = "0.22.0", extras = ["accept-rom-license", "atari"]} [[tool.poetry.source]] name = "pytorch" diff --git a/research_journal.md b/research_journal.md index 5fe1ea2..57c3c0b 100644 --- a/research_journal.md +++ b/research_journal.md @@ -8,7 +8,10 @@ Try IRIs but with pretrained transformer with LoRA adapter ```sh poetry install . ./.venv/bin/activate -python src/main.py env.train.id=BreakoutNoFrameskip-v4 common.device=cuda:0 wandb.mode=online +python src/main.py env.train.id=BreakoutNoFrameskip-v4 common.device=cuda:0 wandb.mode=offline + +# or for quick debug +WANDB_MODE=disabled python -m pdb src/main.py env.train.id=BreakoutNoFrameskip-v4 ``` diff --git a/src/envs/wrappers.py b/src/envs/wrappers.py index b1054a1..eca6f88 100644 --- a/src/envs/wrappers.py +++ b/src/envs/wrappers.py @@ -64,7 +64,7 @@ class NoopResetEnv(gym.Wrapper): if self.override_num_noops is not None: noops = self.override_num_noops else: - noops = self.unwrapped.np_random.randint(1, self.noop_max + 1) + noops = self.unwrapped.np_random.integers(1, self.noop_max + 1) assert noops > 0 obs = None for _ in range(noops):