diff --git a/config/trainer.yaml b/config/trainer.yaml index d1f92a6..390a3ac 100644 --- a/config/trainer.yaml +++ b/config/trainer.yaml @@ -7,7 +7,7 @@ defaults: - datasets: default wandb: - mode: offline + mode: online project: iris entity: null name: null diff --git a/research_journal.md b/research_journal.md index 2e6e327..0d40317 100644 --- a/research_journal.md +++ b/research_journal.md @@ -314,3 +314,7 @@ Ok so it's all just the - rollout, controller by max block size. 10x - the fact that actor_critic can use a larger batch, therefore 4-8x more samples - for each one it imagines 2 + +# 2023-11-18 06:17:55 + +It trained overnight, now I would like to view a replay diff --git a/src/agent.py b/src/agent.py index a8d943c..5534232 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 src.models.actor_critic import ActorCritic +from src.models.tokenizer import Tokenizer +from src.models.world_model import WorldModel +from src.utils import extract_state_dict class Agent(nn.Module): diff --git a/src/game/keymap.py b/src/game/keymap.py index 62aadc7..735eb8c 100644 --- a/src/game/keymap.py +++ b/src/game/keymap.py @@ -12,6 +12,10 @@ def get_keymap_and_action_names(name): if name == 'atari': return ATARI_KEYMAP, ATARI_ACTION_NAMES + + if name == 'atari/CrafterReward-v1': + env_id = name.split('atari/')[1] + return CRAFTER_KEYMAP, gym.make(env_id).action_names assert name.startswith('atari/') env_id = name.split('atari/')[1] @@ -100,4 +104,25 @@ EMPTY_ACTION_NAMES = [ ] EMPTY_KEYMAP = { -} \ No newline at end of file +} + +CRAFTER_KEYMAP = { + pygame.K_a: 'move_left', + pygame.K_d: 'move_right', + pygame.K_w: 'move_up', + pygame.K_s: 'move_down', + pygame.K_SPACE: 'do', + pygame.K_TAB: 'sleep', + + pygame.K_r: 'place_stone', + pygame.K_t: 'place_table', + pygame.K_f: 'place_furnace', + pygame.K_p: 'place_plant', + + pygame.K_1: 'make_wood_pickaxe', + pygame.K_2: 'make_stone_pickaxe', + pygame.K_3: 'make_iron_pickaxe', + pygame.K_4: 'make_wood_sword', + pygame.K_5: 'make_stone_sword', + pygame.K_6: 'make_iron_sword', +} diff --git a/src/models/actor_critic.py b/src/models/actor_critic.py index 976439a..ec1256d 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 src.dataset import Batch +from src.envs.world_model_env import WorldModelEnv +from src.models.tokenizer import Tokenizer +from src.models.world_model import WorldModel +from src.utils import compute_lambda_returns, LossWithIntermediateLosses @dataclass diff --git a/src/models/world_model.py b/src/models/world_model.py index 349423c..eecfe60 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 .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 src.dataset import Batch +from src.models.kv_caching import KeysValues +from src.models.slicer import Embedder, Head +from src.models.tokenizer import Tokenizer +from src.models.transformer import Transformer, TransformerConfig +from src.utils import init_weights, LossWithIntermediateLosses @dataclass diff --git a/src/play.py b/src/play.py index 7b18dee..7e3ab85 100644 --- a/src/play.py +++ b/src/play.py @@ -6,11 +6,11 @@ from hydra.utils import instantiate from omegaconf import DictConfig import torch -from agent import Agent -from envs import SingleProcessEnv, WorldModelEnv -from game import AgentEnv, EpisodeReplayEnv, Game -from models.actor_critic import ActorCritic -from models.world_model import WorldModel +from src.agent import Agent +from src.envs import SingleProcessEnv, WorldModelEnv +from src.game import AgentEnv, EpisodeReplayEnv, Game +from src.models.actor_critic import ActorCritic +from src.models.world_model import WorldModel @hydra.main(config_path="../config", config_name="trainer") diff --git a/src/utils.py b/src/utils.py index 1aaf3b2..d64273d 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 src.episode import Episode from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS