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