mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-12 12:34:10 +08:00
abs imports, fix play for crafter
This commit is contained in:
+1
-1
@@ -7,7 +7,7 @@ defaults:
|
||||
- datasets: default
|
||||
|
||||
wandb:
|
||||
mode: offline
|
||||
mode: online
|
||||
project: iris
|
||||
entity: null
|
||||
name: null
|
||||
|
||||
@@ -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
@@ -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
@@ -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',
|
||||
}
|
||||
|
||||
@@ -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,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
@@ -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
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user