Merge branch 'full_ft'

This commit is contained in:
wassname
2024-04-27 07:04:40 +08:00
32 changed files with 2537 additions and 1219 deletions
+5 -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):
@@ -34,4 +34,5 @@ class Agent(nn.Module):
input_ac = obs if self.actor_critic.use_original_obs else torch.clamp(self.tokenizer.encode_decode(obs, should_preprocess=True, should_postprocess=True), 0, 1)
logits_actions = self.actor_critic(input_ac).logits_actions[:, -1] / temperature
act_token = Categorical(logits=logits_actions).sample() if should_sample else logits_actions.argmax(dim=-1)
# FIXME, is this really just an action and doesn't have an obs in?
return act_token
+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
+2 -1
View File
@@ -65,12 +65,13 @@ class WorldModelEnv:
token = action.clone().detach() if isinstance(action, torch.Tensor) else torch.tensor(action, dtype=torch.long)
token = token.reshape(-1, 1).to(self.device) # (B, 1)
for k in range(num_passes): # assumption that there is only one action token.
# FIXME: hold on we are ONLY passing in the action token! should it not be obs too
outputs_wm = self.world_model(token, past_keys_values=self.keys_values_wm)
output_sequence.append(outputs_wm.output_sequence)
if k == 0:
reward = Categorical(logits=outputs_wm.logits_rewards).sample().float().cpu().numpy().reshape(-1) - 1 # (B,)
done = Categorical(logits=outputs_wm.logits_ends).sample().cpu().numpy().astype(bool).reshape(-1) # (B,)
+33 -2
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,26 @@ 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)
env = ResizeObsWrapper(env, (size, size))
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:
@@ -64,7 +95,7 @@ class NoopResetEnv(gym.Wrapper):
if self.override_num_noops is not None:
noops = self.override_num_noops
else:
noops = self.unwrapped.np_random.randint(1, self.noop_max + 1)
noops = self.unwrapped.np_random.integers(1, self.noop_max + 1)
assert noops > 0
obs = None
for _ in range(noops):
+4 -4
View File
@@ -4,14 +4,14 @@ from PIL import Image
import torch
from torchvision.transforms.functional import InterpolationMode, resize
from agent import Agent
from envs import SingleProcessEnv, WorldModelEnv
from game.keymap import get_keymap_and_action_names
from src.agent import Agent
from src.envs import SingleProcessEnv, WorldModelEnv
from src.game.keymap import get_keymap_and_action_names
class AgentEnv:
def __init__(self, agent: Agent, env: SingleProcessEnv, keymap_name: str, do_reconstruction: bool) -> None:
assert isinstance(env, SingleProcessEnv) or isinstance(env, WorldModelEnv)
assert isinstance(env, SingleProcessEnv) or isinstance(env, WorldModelEnv), f"{env}"
self.agent = agent
self.env = env
_, self.action_names = get_keymap_and_action_names(keymap_name)
+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: 1,
pygame.K_d: 2,
pygame.K_w: 3,
pygame.K_s: 4,
pygame.K_SPACE: 5,
pygame.K_TAB: 6,
pygame.K_r: 7,
pygame.K_t: 8,
pygame.K_f: 9,
pygame.K_p: 10,
pygame.K_1: 11,
pygame.K_2: 12,
pygame.K_3: 13,
pygame.K_4: 14,
pygame.K_5: 15,
pygame.K_6: 16,
}
+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):
+19 -15
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
@@ -34,24 +34,26 @@ 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 = 1
s = 1
self.use_original_obs = use_original_obs
self.conv1 = nn.Conv2d(3, 32, 3, stride=1, padding=1)
self.conv1 = nn.Conv2d(3, 32//s, 3, stride=1, padding=1)
self.maxp1 = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(32, 32, 3, stride=1, padding=1)
self.conv2 = nn.Conv2d(32//s, 32//s, 3, stride=1, padding=1)
self.maxp2 = nn.MaxPool2d(2, 2)
self.conv3 = nn.Conv2d(32, 64, 3, stride=1, padding=1)
self.conv3 = nn.Conv2d(32//s, 64//s, 3, stride=1, padding=1)
self.maxp3 = nn.MaxPool2d(2, 2)
self.conv4 = nn.Conv2d(64, 64, 3, stride=1, padding=1)
self.conv4 = nn.Conv2d(64//s, 64//shrink, 3, stride=1, padding=1)
self.maxp4 = nn.MaxPool2d(2, 2)
self.lstm_dim = 512
self.lstm = nn.LSTMCell(1024, self.lstm_dim)
self.lstm_dim = lstm_dim
self.lstm = nn.LSTMCell(1024//shrink, self.lstm_dim)
self.hx, self.cx = None, None
self.critic_linear = nn.Linear(512, 1)
self.actor_linear = nn.Linear(512, act_vocab_size)
self.critic_linear = nn.Linear(self.lstm_dim, 1)
self.actor_linear = nn.Linear(self.lstm_dim, act_vocab_size)
def __repr__(self) -> str:
return "actor_critic"
@@ -85,7 +87,7 @@ class ActorCritic(nn.Module):
x = F.relu(self.maxp2(self.conv2(x)))
x = F.relu(self.maxp3(self.conv3(x)))
x = F.relu(self.maxp4(self.conv4(x)))
x = torch.flatten(x, start_dim=1)
x = torch.flatten(x, start_dim=1) # [b=32, 64//shrink, 4, 4]
if mask_padding is None:
self.hx, self.cx = self.lstm(x, (self.hx, self.cx))
@@ -146,6 +148,8 @@ class ActorCritic(nn.Module):
outputs_ac = self(obs)
action_token = Categorical(logits=outputs_ac.logits_actions).sample()
# FIXME shouldn't we pass in obs too?
obs, reward, done, _ = wm_env.step(action_token, should_predict_next_obs=(k < horizon - 1))
all_actions.append(action_token)
+2 -1
View File
@@ -31,7 +31,8 @@ class Head(Slicer):
self.head_module = head_module
def forward(self, x: torch.Tensor, num_steps: int, prev_steps: int) -> torch.Tensor:
x_sliced = x[:, self.compute_slice(num_steps, prev_steps)] # x is (B, T, E)
s = self.compute_slice(num_steps, prev_steps)
x_sliced = x[:, s] # x is (B, T, E)
return self.head_module(x_sliced)
+7 -5
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
@@ -23,15 +23,15 @@ class TokenizerEncoderOutput:
class Tokenizer(nn.Module):
def __init__(self, vocab_size: int, embed_dim: int, encoder: Encoder, decoder: Decoder, with_lpips: bool = True) -> None:
def __init__(self, transformer_embedding: nn.Embedding, vocab_size: int, embed_dim: int, encoder: Encoder, decoder: Decoder, with_lpips: bool = True) -> None:
super().__init__()
self.vocab_size = vocab_size
self.encoder = encoder
self.pre_quant_conv = torch.nn.Conv2d(encoder.config.z_channels, embed_dim, 1)
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.embedding = transformer_embedding # pretrained transformer embedding
self.post_quant_conv = torch.nn.Conv2d(embed_dim, decoder.config.z_channels, 1)
self.decoder = decoder
self.embedding.weight.data.uniform_(-1.0 / vocab_size, 1.0 / vocab_size)
# self.embedding.weight.data.uniform_(-1.0 / vocab_size, 1.0 / vocab_size)
self.lpips = LPIPS().eval() if with_lpips else None
def __repr__(self) -> str:
@@ -46,6 +46,8 @@ class Tokenizer(nn.Module):
def compute_loss(self, batch: Batch, **kwargs: Any) -> LossWithIntermediateLosses:
assert self.lpips is not None
observations = self.preprocess_input(rearrange(batch['observations'], 'b t c h w -> (b t) c h w'))
# TODO: in the delta-IRIS paper (https://openreview.net/forum?id=o8IDoZggqO) they encode(x0, a0, x1) -> z1 and decode(x0, a0, z1). In esense the tokens only need to encode the change
# note they also do dynamics(x0, a0, z1) -> z2. decode(x1, a1, z2) -> x2
z, z_quantized, reconstructions = self(observations, should_preprocess=False, should_postprocess=False)
# Codebook loss. Notes:
+138 -94
View File
@@ -2,119 +2,163 @@
# Credits to https://github.com/karpathy/minGPT
# """
# from dataclasses import dataclass
# import math
# from typing import Optional
# from einops import rearrange
# import torch
# import torch.nn as nn
# from torch.nn import functional as F
from dataclasses import dataclass
import math
from typing import Optional
from contextlib import contextmanager
from einops import rearrange
import torch
import torch.nn as nn
from torch.nn import functional as F
from loguru import logger
# from .kv_caching import KeysValues, KVCache
# @dataclass
# class TransformerConfig:
# tokens_per_block: int
# max_blocks: int
# attention: str
@dataclass
class TransformerConfig:
# model_name: str = "stabilityai/stablelm-3b-4e1t"
# https://huggingface.co/PY007/TinyLlama-1.1B-intermediate-step-715k-1.5T
vocab_size: int = 32000
embed_dim: int = 2048
max_blocks: int = 20
tokens_per_block: int = 17
model_name: str = "PY007/TinyLlama-1.1B-intermediate-step-715k-1.5T"
dropout: float = 0.1
rank: int = 32
# num_layers: int
# num_heads: int
# embed_dim: int
# embed_pdrop: float
# resid_pdrop: float
# attn_pdrop: float
# @property
# def max_tokens(self):
# return self.tokens_per_block * self.max_blocks
@property
def max_tokens(self):
return self.tokens_per_block * self.max_blocks
def freeze(n: nn.Module):
for p in n.parameters():
p.requires_grad = False
return n
# class Transformer(nn.Module):
# def __init__(self, config: TransformerConfig) -> None:
# super().__init__()
# self.config = config
# self.drop = nn.Dropout(config.embed_pdrop)
# self.blocks = nn.ModuleList([Block(config) for _ in range(config.num_layers)])
# self.ln_f = nn.LayerNorm(config.embed_dim)
class Transformer(nn.Module):
def __init__(self, config: TransformerConfig) -> None:
super().__init__()
self.config = config
self.model = load_pretrained_model(config)
self.ln_f = nn.Linear(self.model.config.vocab_size, config.embed_dim)
self.embedding = freeze(self.model.base_model.embed_tokens.to(torch.float)) # HACK: custom path to embeddings layer for model
# def generate_empty_keys_values(self, n: int, max_tokens: int) -> KeysValues:
# device = self.ln_f.weight.device # Assumption that all submodules are on the same device
# return KeysValues(n, self.config.num_heads, max_tokens, self.config.embed_dim, self.config.num_layers, device)
def generate_empty_keys_values(self, n: int, max_tokens: int) -> KeysValues:
device = self.ln_f.weight.device # Assumption that all submodules are on the same device
return KeysValues(n, 1, max_tokens, self.config.embed_dim, 1, device)
# def forward(self, sequences: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor:
# assert past_keys_values is None or len(past_keys_values) == len(self.blocks)
# x = self.drop(sequences)
# for i, block in enumerate(self.blocks):
# x = block(x, None if past_keys_values is None else past_keys_values[i])
def forward(self, sequences: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor:
assert past_keys_values is None or len(past_keys_values) == 1
with torch.cuda.amp.autocast(dtype=torch.bfloat16):
outputs = self.model(
inputs_embeds=sequences,
return_dict=True,
output_hidden_states=True,
)
x = outputs.logits
x = self.ln_f(x)
# fake it, since it's used to keep track of steps
if past_keys_values is not None:
k_size = past_keys_values[0]._k_cache._cache.size()
v_size = (k_size[0], k_size[1], x.shape[1], k_size[3])
past_keys_values[0].update(torch.rand(v_size), torch.rand(v_size))
return x
# x = self.ln_f(x)
# return x
# class Block(nn.Module):
# def __init__(self, config: TransformerConfig) -> None:
# super().__init__()
# self.ln1 = nn.LayerNorm(config.embed_dim)
# self.ln2 = nn.LayerNorm(config.embed_dim)
# self.attn = SelfAttention(config)
# self.mlp = nn.Sequential(
# nn.Linear(config.embed_dim, 4 * config.embed_dim),
# nn.GELU(),
# nn.Linear(4 * config.embed_dim, config.embed_dim),
# nn.Dropout(config.resid_pdrop),
# )
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
from peft import PeftModel, LoraConfig
import peft
# def forward(self, x: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor:
# x_attn = self.attn(self.ln1(x), past_keys_values)
# x = x + x_attn
# x = x + self.mlp(self.ln2(x))
# return x
def load_pretrained_model(config, device="cuda:0"):
tokenizer = AutoTokenizer.from_pretrained(config.model_name, trust_remote_code=True)
tokenizer.padding_side = "left"
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_quant_type="nf4",
bnb_4bit_use_double_quant=True,
)
base_model = AutoModelForCausalLM.from_pretrained(
config.model_name,
device_map={"": device},
quantization_config=bnb_config,
torch_dtype=torch.bfloat16,
trust_remote_code=True
)
peft_config = peft.LoraConfig(
peft.TaskType.CAUSAL_LM,
inference_mode=False,
r=config.rank,
lora_alpha=config.rank*2, # Adjusting the LoRA rank is essential, and so is selecting an apt alpha value. A good heuristic is setting alpha at twice the rank's value. https://magazine.sebastianraschka.com/p/practical-tips-for-finetuning-llms
lora_dropout=config.dropout,
# TODO: If you're incorporating LoRA, ensure it's applied across all layers, not just to the Key and Value matrices, to maximize model performance.
target_modules=[
"self_attn.q_proj",
"self_attn.k_proj",
"self_attn.v_proj",
"self_attn.o_proj",
"mlp.gate_proj",
"mlp.up_proj",
"mlp.down_proj",
# "wte", "embed_tokens",
# "lm_head",
],
# bias="lora_only",
# tune the embedding layer and prediction head
modules_to_save = ["lm_head",], # we want the classifier parameters to be trained too when fine-tuning the base model on our custom dataset. To ensure that the classifier parameters are also trained, we specify modules_to_save.
)
base_model_peft = base_model
# base_model_peft = peft.get_peft_model(base_model, peft_config)
# base_model_peft.add_adapter(adapter_name="dynamics", peft_config=peft_config) # make and set an adapter
disable_causal_mask_always()
# print(base_model_peft.print_trainable_parameters())
logger.debug(f"loaded model {base_model_peft}")
return base_model_peft
@contextmanager
def set_adapter(model, adapter_name):
old_adapter_name = model.active_adapter
try:
if adapter_name is not None:
model.set_adapter(adapter_name)
yield model
else:
with model.disable_adapter():
yield model
finally:
model.set_adapter(old_adapter_name)
# class SelfAttention(nn.Module):
# def __init__(self, config: TransformerConfig) -> None:
# super().__init__()
# assert config.embed_dim % config.num_heads == 0
# assert config.attention in ('causal', 'block_causal')
# self.num_heads = config.num_heads
# self.key = nn.Linear(config.embed_dim, config.embed_dim)
# self.query = nn.Linear(config.embed_dim, config.embed_dim)
# self.value = nn.Linear(config.embed_dim, config.embed_dim)
# self.attn_drop = nn.Dropout(config.attn_pdrop)
# self.resid_drop = nn.Dropout(config.resid_pdrop)
# self.proj = nn.Linear(config.embed_dim, config.embed_dim)
def disable_causal_mask_always():
import transformers.models.llama.modeling_llama as modeling
# causal_mask = torch.tril(torch.ones(config.max_tokens, config.max_tokens))
# block_causal_mask = torch.max(causal_mask, torch.block_diag(*[torch.ones(config.tokens_per_block, config.tokens_per_block) for _ in range(config.max_blocks)]))
# self.register_buffer('mask', causal_mask if config.attention == 'causal' else block_causal_mask)
decoder_fn = modeling._make_causal_mask
# def forward(self, x: torch.Tensor, kv_cache: Optional[KVCache] = None) -> torch.Tensor:
# B, T, C = x.size()
# if kv_cache is not None:
# b, nh, L, c = kv_cache.shape
# assert nh == self.num_heads and b == B and c * nh == C
# else:
# L = 0
def encoder_fn(*args, **kwargs):
return torch.zeros_like(decoder_fn(*args, **kwargs))
# q = self.query(x).view(B, T, self.num_heads, C // self.num_heads).transpose(1, 2) # (B, nh, T, hs)
# k = self.key(x).view(B, T, self.num_heads, C // self.num_heads).transpose(1, 2) # (B, nh, T, hs)
# v = self.value(x).view(B, T, self.num_heads, C // self.num_heads).transpose(1, 2) # (B, nh, T, hs)
modeling._make_causal_mask = encoder_fn
# if kv_cache is not None:
# kv_cache.update(k, v)
# k, v = kv_cache.get()
@contextmanager
def disable_causal_mask():
import transformers.models.llama.modeling_llama as modeling
# att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
# att = att.masked_fill(self.mask[L:L + T, :L + T] == 0, float('-inf'))
# att = F.softmax(att, dim=-1)
# att = self.attn_drop(att)
# y = att @ v
# y = rearrange(y, 'b h t e -> b t (h e)')
decoder_fn = modeling._make_causal_mask
# y = self.resid_drop(self.proj(y))
def encoder_fn(*args, **kwargs):
return torch.zeros_like(decoder_fn(*args, **kwargs))
# return y
try:
modeling._make_causal_mask = encoder_fn
yield
finally:
modeling._make_causal_mask = decoder_fn
+39 -15
View File
@@ -6,13 +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 .bigvae import BigVAE, BigVAEConfig
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
@@ -27,21 +26,35 @@ class WorldModel(nn.Module):
def __init__(self, obs_vocab_size: int, act_vocab_size: int, config: BigVAEConfig) -> None:
super().__init__()
self.obs_vocab_size, self.act_vocab_size = obs_vocab_size, act_vocab_size
self.config = config
self.transformer = BigVAE(config)
self.config = config
all_but_last_obs_tokens_pattern = torch.ones(config.tokens_per_block)
all_but_last_obs_tokens_pattern[-2] = 0
act_tokens_pattern = torch.zeros(self.config.tokens_per_block)
act_tokens_pattern[-1] = 1
obs_tokens_pattern = 1 - act_tokens_pattern
self.transformer = Transformer(config)
transformer_embedding = self.transformer.embedding
self.pos_emb = nn.Embedding(config.max_tokens, config.embed_dim)
self.act_emb = nn.Embedding(act_vocab_size, config.embed_dim)
# FIXME: having slices is unclear. maybe it's better just to have obs and action embeddings?
self.embedder = Embedder(
max_blocks=config.max_blocks,
block_masks=[act_tokens_pattern, obs_tokens_pattern],
embedding_tables=nn.ModuleList([nn.Embedding(act_vocab_size, config.embed_dim), nn.Embedding(obs_vocab_size, config.embed_dim)])
embedding_tables=nn.ModuleList([self.act_emb, transformer_embedding])
)
# why have this? Well I worry that the transformer can't adapt, since so much is frozen
# TODO: If I get the dynamics model working, maybe try without it
self.post_embed = nn.Sequential(
nn.Linear(config.embed_dim, config.embed_dim),
nn.ReLU(),
nn.Linear(config.embed_dim, config.embed_dim),
# nn.ReLU(),
# nn.Linear(config.embed_dim, config.embed_dim)
)
self.head_observations = Head(
@@ -74,19 +87,28 @@ class WorldModel(nn.Module):
)
)
self.apply(init_weights)
# don't apply to transformer or transformer/obs embeddings
self.act_emb.apply(init_weights)
self.pos_emb.apply(init_weights)
self.post_embed.apply(init_weights)
self.head_observations.apply(init_weights)
self.head_rewards.apply(init_weights)
self.head_ends.apply(init_weights)
def __repr__(self) -> str:
return "world_model"
def forward(self, tokens: torch.LongTensor, past_keys_values: Optional[KeysValues] = None) -> WorldModelOutput:
num_steps = tokens.size(1) # (B, T)
num_steps = tokens.size(1) # (B=8, T=170) where often the last 10 are actons
assert num_steps <= self.config.max_tokens
prev_steps = 0 if past_keys_values is None else past_keys_values.size
sequences = self.embedder(tokens, num_steps, prev_steps) + self.pos_emb(prev_steps + torch.arange(num_steps, device=tokens.device))
# [batch=8, num_steps=170, embed_size=2048]
sequences = self.post_embed(sequences)
x = self.transformer(sequences, past_keys_values)
logits_observations = self.head_observations(x, num_steps=num_steps, prev_steps=prev_steps)
@@ -97,11 +119,13 @@ class WorldModel(nn.Module):
def compute_loss(self, batch: Batch, tokenizer: Tokenizer, **kwargs: Any) -> LossWithIntermediateLosses:
with torch.no_grad():
obs_tokens = tokenizer.encode(batch['observations'], should_preprocess=True).tokens # (BL, K)
# with torch.no_grad():
# [B=8, S=10, Colors=3, H=64, W=64] -> [B=8, S=10, 16]
obs_tokens = tokenizer.encode(batch['observations'], should_preprocess=True).tokens # (BL, K)
act_tokens = rearrange(batch['actions'], 'b l -> b l 1')
tokens = rearrange(torch.cat((obs_tokens, act_tokens), dim=2), 'b l k1 -> b (l k1)') # (B, L(K+1))
# So first 10 are observation, the last 10 tokens are actions
outputs = self(tokens)
+67 -26
View File
@@ -1,4 +1,4 @@
from functools import partial
from functools import partial
from pathlib import Path
import hydra
@@ -6,55 +6,96 @@ 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
from src.models.tokenizer import Tokenizer
@hydra.main(config_path="../config", config_name="trainer")
def main(cfg: DictConfig):
device = torch.device(cfg.common.device)
assert cfg.mode in ('episode_replay', 'agent_in_env', 'agent_in_world_model', 'play_in_world_model')
assert cfg.mode in (
"episode_replay",
"agent_in_env",
"agent_in_world_model",
"play_in_world_model",
)
env_fn = partial(instantiate, config=cfg.env.test)
test_env = SingleProcessEnv(env_fn)
if cfg.mode.startswith('agent_in_'):
if cfg.mode.startswith("agent_in_"):
h, w, _ = test_env.env.unwrapped.observation_space.shape
else:
h, w = 64, 64
multiplier = 800 // h
size = [h * multiplier, w * multiplier]
if cfg.mode == 'episode_replay':
env = EpisodeReplayEnv(replay_keymap_name=cfg.env.keymap, episode_dir=Path('media/episodes'))
keymap = 'episode_replay'
if cfg.mode == "episode_replay":
env = EpisodeReplayEnv(
replay_keymap_name=cfg.env.keymap, episode_dir=Path("media/episodes")
)
keymap = "episode_replay"
else:
tokenizer = instantiate(cfg.tokenizer)
world_model = WorldModel(obs_vocab_size=tokenizer.vocab_size, act_vocab_size=test_env.num_actions, config=instantiate(cfg.world_model))
actor_critic = ActorCritic(**cfg.actor_critic, act_vocab_size=test_env.num_actions)
# tokenizer = instantiate(cfg.tokenizer)
world_model = WorldModel(
obs_vocab_size=cfg.tokenizer.vocab_size,
act_vocab_size=test_env.num_actions,
config=instantiate(cfg.world_model),
)
transformer_embedding = world_model.transformer.embedding
tokenizer = Tokenizer(
transformer_embedding=transformer_embedding,
vocab_size=cfg.tokenizer.vocab_size,
embed_dim=cfg.tokenizer.embed_dim,
encoder=instantiate(cfg.tokenizer.encoder),
decoder=instantiate(cfg.tokenizer.decoder),
)
actor_critic = ActorCritic(
**cfg.actor_critic, act_vocab_size=test_env.num_actions
)
agent = Agent(tokenizer, world_model, actor_critic).to(device)
agent.load(Path('checkpoints/last.pt'), device)
agent.load(Path("checkpoints/last.pt"), device)
if cfg.mode == 'play_in_world_model':
env = WorldModelEnv(tokenizer=agent.tokenizer, world_model=agent.world_model, device=device, env=env_fn())
if cfg.mode == "play_in_world_model":
env = WorldModelEnv(
tokenizer=agent.tokenizer,
world_model=agent.world_model,
device=device,
env=env_fn(),
)
keymap = cfg.env.keymap
elif cfg.mode == 'agent_in_env':
env = AgentEnv(agent, test_env, cfg.env.keymap, do_reconstruction=cfg.reconstruction)
keymap = 'empty'
elif cfg.mode == "agent_in_env":
env = AgentEnv(
agent, test_env, cfg.env.keymap, do_reconstruction=cfg.reconstruction
)
keymap = "empty"
if cfg.reconstruction:
size[1] *= 3
elif cfg.mode == 'agent_in_world_model':
wm_env = WorldModelEnv(tokenizer=agent.tokenizer, world_model=agent.world_model, device=device, env=env_fn())
elif cfg.mode == "agent_in_world_model":
wm_env = WorldModelEnv(
tokenizer=agent.tokenizer,
world_model=agent.world_model,
device=device,
env=env_fn(),
)
env = AgentEnv(agent, wm_env, cfg.env.keymap, do_reconstruction=False)
keymap = 'empty'
keymap = "empty"
game = Game(env, keymap_name=keymap, size=size, fps=cfg.fps, verbose=bool(cfg.header), record_mode=bool(cfg.save_mode))
game = Game(
env,
keymap_name=keymap,
size=size,
fps=cfg.fps,
verbose=bool(cfg.header),
record_mode=bool(cfg.save_mode),
)
game.run()
+27 -13
View File
@@ -14,14 +14,15 @@ 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
from src.models.tokenizer import Tokenizer
class Trainer:
@@ -46,6 +47,7 @@ class Trainer:
self.reconstructions_dir = self.media_dir / 'reconstructions'
if not cfg.common.resume:
print('cwd', Path.cwd())
config_dir = Path('config')
config_path = config_dir / 'trainer.yaml'
config_dir.mkdir(exist_ok=False, parents=False)
@@ -79,8 +81,15 @@ class Trainer:
assert self.cfg.training.should or self.cfg.evaluation.should
env = train_env if self.cfg.training.should else test_env
tokenizer = instantiate(cfg.tokenizer)
world_model = WorldModel(obs_vocab_size=tokenizer.vocab_size, act_vocab_size=env.num_actions, config=instantiate(cfg.world_model))
world_model = WorldModel(obs_vocab_size=cfg.tokenizer.vocab_size, act_vocab_size=env.num_actions, config=instantiate(cfg.world_model))
transformer_embedding = world_model.transformer.embedding
tokenizer = Tokenizer(
transformer_embedding=transformer_embedding,
vocab_size=cfg.tokenizer.vocab_size,
embed_dim=cfg.tokenizer.embed_dim,
encoder=instantiate(cfg.tokenizer.encoder),
decoder=instantiate(cfg.tokenizer.decoder),
)
actor_critic = ActorCritic(**cfg.actor_critic, act_vocab_size=env.num_actions)
self.agent = Agent(tokenizer, world_model, actor_critic).to(self.device)
print(f'{sum(p.numel() for p in self.agent.tokenizer.parameters())} parameters in agent.tokenizer')
@@ -88,7 +97,11 @@ class Trainer:
print(f'{sum(p.numel() for p in self.agent.actor_critic.parameters())} parameters in agent.actor_critic')
self.optimizer_tokenizer = torch.optim.Adam(self.agent.tokenizer.parameters(), lr=cfg.training.learning_rate)
self.optimizer_world_model = configure_optimizer(self.agent.world_model, cfg.training.learning_rate, cfg.training.world_model.weight_decay)
# self.optimizer_world_model = configure_optimizer([self.agent.tokenizer, self.agent.world_model], cfg.training.learning_rate, cfg.training.world_model.weight_decay)
self.optimizer_world_model = torch.optim.Adam(
list(self.agent.tokenizer.parameters())+list(self.agent.world_model.parameters()),
lr=cfg.training.learning_rate
)
self.optimizer_actor_critic = torch.optim.Adam(self.agent.actor_critic.parameters(), lr=cfg.training.learning_rate)
if cfg.initialization.path_to_checkpoint is not None:
@@ -136,10 +149,10 @@ class Trainer:
if epoch > cfg_tokenizer.start_after_epochs:
metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, **cfg_tokenizer)
self.agent.tokenizer.eval()
if epoch > cfg_world_model.start_after_epochs:
metrics_world_model = self.train_component(self.agent.world_model, self.optimizer_world_model, sequence_length=self.cfg.common.sequence_length, sample_from_start=True, tokenizer=self.agent.tokenizer, **cfg_world_model)
self.agent.tokenizer.eval()
self.agent.world_model.eval()
if epoch > cfg_actor_critic.start_after_epochs:
@@ -169,6 +182,7 @@ class Trainer:
if max_grad_norm is not None:
torch.nn.utils.clip_grad_norm_(component.parameters(), max_grad_norm)
optimizer.step()
metrics = {f'{str(component)}/train/total_loss': loss_total_epoch, **intermediate_losses}
@@ -190,7 +204,7 @@ class Trainer:
if epoch > cfg_world_model.start_after_epochs:
metrics_world_model = self.eval_component(self.agent.world_model, cfg_world_model.batch_num_samples, sequence_length=self.cfg.common.sequence_length, tokenizer=self.agent.tokenizer)
if epoch > cfg_actor_critic.start_after_epochs:
if epoch > cfg_world_model.start_after_epochs:
self.inspect_imagination(epoch)
if cfg_tokenizer.save_reconstructions:
+39 -22
View File
@@ -3,47 +3,64 @@ import cv2
from pathlib import Path
import random
import shutil
from loguru import logger
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
def configure_optimizer(model, learning_rate, weight_decay, *blacklist_module_names):
def configure_optimizer(models, learning_rate, weight_decay, *blacklist_module_names):
"""Credits to https://github.com/karpathy/minGPT"""
# FIXME: check this is still good for LoRA
# separate out all parameters to those that will and won't experience regularizing weight decay
decay = set()
no_decay = set()
decay_params = []
no_decay_params = []
param_dict = {}
whitelist_weight_modules = (torch.nn.Linear, torch.nn.Conv1d)
blacklist_weight_modules = (torch.nn.LayerNorm, torch.nn.Embedding)
for mn, m in model.named_modules():
for pn, p in m.named_parameters():
fpn = '%s.%s' % (mn, pn) if mn else pn # full param name
if any([fpn.startswith(module_name) for module_name in blacklist_module_names]):
no_decay.add(fpn)
elif 'bias' in pn:
# all biases will not be decayed
no_decay.add(fpn)
elif pn.endswith('weight') and isinstance(m, whitelist_weight_modules):
# weights of whitelist modules will be weight decayed
decay.add(fpn)
elif pn.endswith('weight') and isinstance(m, blacklist_weight_modules):
# weights of blacklist modules will NOT be weight decayed
no_decay.add(fpn)
blacklist_weight_modules = tuple(ALL_LAYERNORM_LAYERS+[torch.nn.Embedding])
for model in models:
for mn, m in model.named_modules():
for pn, p in m.named_parameters():
fpn = '%s.%s' % (mn, pn) if mn else pn # full param name
if any([fpn.startswith(module_name) for module_name in blacklist_module_names]):
no_decay.add(fpn)
no_decay_params.append(p)
elif 'bias' in pn:
# all biases will not be decayed
no_decay.add(fpn)
no_decay_params.append(p)
elif pn.endswith('weight') and isinstance(m, whitelist_weight_modules):
# weights of whitelist modules will be weight decayed
decay.add(fpn)
decay_params.append(p)
elif pn.endswith('weight') and isinstance(m, blacklist_weight_modules):
# weights of blacklist modules will NOT be weight decayed
no_decay.add(fpn)
no_decay_params.append(p)
else:
logger.warning(f"Parameter {fpn} of module {mn} not handled!")
# raise NotImplementedError(f"Parameter {fpn} of module {m} not handled!")
decay.add(fpn)
decay_params.append(p)
# validate that we considered every parameter
param_dict = {pn: p for pn, p in model.named_parameters()}
# validate that we considered every parameter
param_dict.update({pn: p for pn, p in model.named_parameters()})
inter_params = decay & no_decay
union_params = decay | no_decay
# logger.debug(f"decay {decay} no_decay {no_decay}")
assert len(inter_params) == 0, f"parameters {str(inter_params)} made it into both decay/no_decay sets!"
assert len(param_dict.keys() - union_params) == 0, f"parameters {str(param_dict.keys() - union_params)} were not separated into either decay/no_decay set!"
# create the pytorch optimizer object
optim_groups = [
{"params": [param_dict[pn] for pn in sorted(list(decay))], "weight_decay": weight_decay},
{"params": [param_dict[pn] for pn in sorted(list(no_decay))], "weight_decay": 0.0},
{"params": no_decay_params, "weight_decay": weight_decay},
{"params": decay_params, "weight_decay": 0.0},
]
optimizer = torch.optim.AdamW(optim_groups, lr=learning_rate)
return optimizer