abs imports, fix play for crafter

This commit is contained in:
wassname
2023-11-18 07:44:05 +08:00
parent 1dabd79251
commit a341e7aaad
8 changed files with 52 additions and 23 deletions
+1 -1
View File
@@ -7,7 +7,7 @@ defaults:
- datasets: default
wandb:
mode: offline
mode: online
project: iris
entity: null
name: null
+4
View File
@@ -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
+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 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):
+26 -1
View File
@@ -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 = {
}
}
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',
}
+5 -5
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 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
+6 -6
View File
@@ -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
+5 -5
View File
@@ -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")
+1 -1
View File
@@ -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