relative imports

This commit is contained in:
wassname
2023-11-17 12:30:16 +08:00
parent 1af7aa74fa
commit a778296e40
12 changed files with 74 additions and 33 deletions
+9
View File
@@ -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]
+4 -4
View File
@@ -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):
+5 -5
View File
@@ -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:
+1 -1
View File
@@ -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]
+1 -1
View File
@@ -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
+31 -1
View File
@@ -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:
+3 -1
View File
@@ -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):
+7 -7
View File
@@ -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
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
+8 -8
View File
@@ -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:
+1 -1
View File
@@ -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