mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
relative imports
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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,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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user