diff --git a/research_journal.md b/research_journal.md index 5e2b17e..8bdf3fe 100644 --- a/research_journal.md +++ b/research_journal.md @@ -221,5 +221,14 @@ How long to train? `600*10//6/24` = 41 days - 600 epochs * 10 minutes / 6 to get hours, 24 to get days +# 2023-11-17 07:59:44 +so I've got it working with these times. But maybe it's too small + +Epoch 148 / 600 + +Experience collection (train_dataset): 100%|████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [00:10<00:00, 19.25it/s] +Training tokenizer: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [00:59<00:00, 3.34it/s] +Training world_model: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 200/200 [01:15<00:00, 2.64it/s] +Training actor_critic: 100%|██████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 20/20 [03:20<00:00, 10.05s/it] diff --git a/src/agent.py b/src/agent.py index f5dd1d3..a8d943c 100644 --- a/src/agent.py +++ b/src/agent.py @@ -4,10 +4,10 @@ import torch from torch.distributions.categorical import Categorical import torch.nn as nn -from models.actor_critic import ActorCritic -from models.tokenizer import Tokenizer -from models.world_model import WorldModel -from utils import extract_state_dict +from .models.actor_critic import ActorCritic +from .models.tokenizer import Tokenizer +from .models.world_model import WorldModel +from .utils import extract_state_dict class Agent(nn.Module): diff --git a/src/collector.py b/src/collector.py index 85e0150..8c5947e 100644 --- a/src/collector.py +++ b/src/collector.py @@ -8,11 +8,11 @@ import torch from tqdm import tqdm import wandb -from agent import Agent -from dataset import EpisodesDataset -from envs import SingleProcessEnv, MultiProcessEnv -from episode import Episode -from utils import EpisodeDirManager, RandomHeuristic +from src.agent import Agent +from src.dataset import EpisodesDataset +from src.envs import SingleProcessEnv, MultiProcessEnv +from src.episode import Episode +from src.utils import EpisodeDirManager, RandomHeuristic class Collector: diff --git a/src/dataset.py b/src/dataset.py index 59d9f30..5cc2d6e 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -7,7 +7,7 @@ from typing import Dict, List, Optional, Tuple import psutil import torch -from episode import Episode +from src.episode import Episode Batch = Dict[str, torch.Tensor] diff --git a/src/envs/__init__.py b/src/envs/__init__.py index 5afe68b..4d85674 100644 --- a/src/envs/__init__.py +++ b/src/envs/__init__.py @@ -1,4 +1,4 @@ from .multi_process_env import MultiProcessEnv -from .wrappers import make_atari, ResizeObsWrapper +from .wrappers import make_atari, make_crafter, make_env, ResizeObsWrapper from .single_process_env import SingleProcessEnv from .world_model_env import WorldModelEnv diff --git a/src/envs/wrappers.py b/src/envs/wrappers.py index eca6f88..4183483 100644 --- a/src/envs/wrappers.py +++ b/src/envs/wrappers.py @@ -7,11 +7,22 @@ from typing import Tuple import gym import numpy as np from PIL import Image +import crafter + + +def make_env(id, size=64, max_episode_steps=None, noop_max=30, frame_skip=4, done_on_life_loss=False, clip_reward=False): + if id.startswith('Crafter'): + return make_crafter(id, size=size, max_episode_steps=max_episode_steps, done_on_life_loss=done_on_life_loss) + if id.startswith('MiniHack'): + return make_minihack(size=size, max_episode_steps=max_episode_steps, done_on_life_loss=done_on_life_loss) + else: + return make_atari(id, size, max_episode_steps, noop_max, frame_skip, done_on_life_loss, clip_reward) def make_atari(id, size=64, max_episode_steps=None, noop_max=30, frame_skip=4, done_on_life_loss=False, clip_reward=False): env = gym.make(id) - assert 'NoFrameskip' in env.spec.id or 'Frameskip' not in env.spec + print(env.spec) + assert 'NoFrameskip' in env.spec.id or 'Frameskip' not in str(env.spec) env = ResizeObsWrapper(env, (size, size)) if clip_reward: env = RewardClippingWrapper(env) @@ -24,6 +35,25 @@ def make_atari(id, size=64, max_episode_steps=None, noop_max=30, frame_skip=4, d env = EpisodicLifeEnv(env) return env +def make_crafter(id, size=64, max_episode_steps=None, done_on_life_loss=False): + # https://github.com/danijar/dreamerv2/blob/07d906e9c4322c6fc2cd6ed23e247ccd6b7c8c41/dreamerv2/common/envs.py#L242 + # https://github.com/footoredo/torchbeast/blob/12939569cc46b6a8616e4c25b138d97248cc8581/torchbeast/atari_wrappers.py#L301 + env = gym.make(id) + return env + + +def make_minihack(id, size=64, max_episode_steps=None, noop_max=30, frame_skip=4, done_on_life_loss=False, clip_reward=False): + # https://github.com/facebookresearch/minihack/blob/47065748f04714c49ba5b52fb74d166228c7acc1/minihack/agent/common/envs/wrapper.py#L117 + # https://github.com/roger-creus/SOFE/blob/5551a115a9c7e1d632cf6996bf5dcabde59cdcc5/e3b/minihack/torchbeast/src/utils.py#L110 + env = gym.make(id, + # https://minihack.readthedocs.io/en/latest/getting-started/observation_spaces.html + observation_keys=("pixel_crop"), + # obs_crop_h=9, + # obs_crop_w=9, + ) + env = ResizeObsWrapper(env, (size, size)) + return env + class ResizeObsWrapper(gym.ObservationWrapper): def __init__(self, env: gym.Env, size: Tuple[int, int]) -> None: diff --git a/src/main.py b/src/main.py index be4177b..6a84b94 100644 --- a/src/main.py +++ b/src/main.py @@ -2,7 +2,9 @@ import hydra from omegaconf import DictConfig from trainer import Trainer - +from loguru import logger +import sys +logger.add(sys.stderr, format="{time} {level} {message}", filter="my_module", level="INFO") @hydra.main(config_path="../config", config_name="trainer") def main(cfg: DictConfig): diff --git a/src/models/actor_critic.py b/src/models/actor_critic.py index 75b4261..13eb51d 100644 --- a/src/models/actor_critic.py +++ b/src/models/actor_critic.py @@ -10,11 +10,11 @@ import torch.nn as nn import torch.nn.functional as F from tqdm import tqdm -from dataset import Batch -from envs.world_model_env import WorldModelEnv -from models.tokenizer import Tokenizer -from models.world_model import WorldModel -from utils import compute_lambda_returns, LossWithIntermediateLosses +from ..dataset import Batch +from ..envs.world_model_env import WorldModelEnv +from ..models.tokenizer import Tokenizer +from ..models.world_model import WorldModel +from ..utils import compute_lambda_returns, LossWithIntermediateLosses @dataclass @@ -34,7 +34,7 @@ class ImagineOutput: class ActorCritic(nn.Module): - def __init__(self, act_vocab_size, use_original_obs: bool = False) -> None: + def __init__(self, act_vocab_size, use_original_obs: bool = False, lstm_dim = 16) -> None: super().__init__() shrink = 8 s = 2 @@ -48,7 +48,7 @@ class ActorCritic(nn.Module): self.conv4 = nn.Conv2d(64//s, 64//shrink, 3, stride=1, padding=1) self.maxp4 = nn.MaxPool2d(2, 2) - self.lstm_dim = 16 + self.lstm_dim = lstm_dim self.lstm = nn.LSTMCell(1024//shrink, self.lstm_dim) self.hx, self.cx = None, None diff --git a/src/models/tokenizer/tokenizer.py b/src/models/tokenizer/tokenizer.py index 6528fa4..1b8723d 100644 --- a/src/models/tokenizer/tokenizer.py +++ b/src/models/tokenizer/tokenizer.py @@ -9,10 +9,10 @@ from einops import rearrange import torch import torch.nn as nn -from dataset import Batch +from src.dataset import Batch from .lpips import LPIPS from .nets import Encoder, Decoder -from utils import LossWithIntermediateLosses +from src.utils import LossWithIntermediateLosses @dataclass diff --git a/src/models/world_model.py b/src/models/world_model.py index 194119e..349423c 100644 --- a/src/models/world_model.py +++ b/src/models/world_model.py @@ -6,12 +6,12 @@ import torch import torch.nn as nn import torch.nn.functional as F -from dataset import Batch +from ..dataset import Batch from .kv_caching import KeysValues from .slicer import Embedder, Head from .tokenizer import Tokenizer from .transformer import Transformer, TransformerConfig -from utils import init_weights, LossWithIntermediateLosses +from ..utils import init_weights, LossWithIntermediateLosses @dataclass diff --git a/src/trainer.py b/src/trainer.py index 8b53aa5..5091fc3 100644 --- a/src/trainer.py +++ b/src/trainer.py @@ -14,14 +14,14 @@ import torch.nn as nn from tqdm import tqdm import wandb -from agent import Agent -from collector import Collector -from envs import SingleProcessEnv, MultiProcessEnv -from episode import Episode -from make_reconstructions import make_reconstructions_from_batch -from models.actor_critic import ActorCritic -from models.world_model import WorldModel -from utils import configure_optimizer, EpisodeDirManager, set_seed +from src.agent import Agent +from src.collector import Collector +from src.envs import SingleProcessEnv, MultiProcessEnv +from src.episode import Episode +from src.make_reconstructions import make_reconstructions_from_batch +from src.models.actor_critic import ActorCritic +from src.models.world_model import WorldModel +from src.utils import configure_optimizer, EpisodeDirManager, set_seed class Trainer: diff --git a/src/utils.py b/src/utils.py index ed8e31a..1aaf3b2 100644 --- a/src/utils.py +++ b/src/utils.py @@ -8,7 +8,7 @@ import numpy as np import torch import torch.nn as nn -from episode import Episode +from .episode import Episode from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS