using frozen transformer embedding: runs

- use frozen transformer encoder where possible (obs, worldmodel, tokenizer, etc)
- it's frozen
- a huge vocab
This commit is contained in:
wassname
2023-11-19 10:50:48 +08:00
parent 7098b67955
commit 328d281651
8 changed files with 116 additions and 55 deletions
+3 -3
View File
@@ -1,14 +1,14 @@
_target_: src.models.tokenizer.Tokenizer
vocab_size: 32000 # change to llm vocab dim
embed_dim: 2048 # change this to whatever the embedding dimension is in your pretrained llm 2048 for llama. 2560 for stablelm
vocab_size: ${..world_model.vocab_size}
embed_dim: ${..world_model.embed_dim}
encoder:
_target_: src.models.tokenizer.Encoder
config:
_target_: src.models.tokenizer.EncoderDecoderConfig
resolution: 64
in_channels: 3
z_channels: 32000
z_channels: ${...vocab_size}
ch: 64
ch_mult: [1, 1, 1, 1, 1]
num_res_blocks: 2
+5 -6
View File
@@ -1,9 +1,8 @@
_target_: src.models.TransformerConfig
max_blocks: 10 # this is the rollout length when training policy
num_layers: 1
num_heads: 1
embed_dim: ${..tokenizer.embed_dim}
dropout: 0.1
model_name: "PY007/TinyLlama-1.1B-intermediate-step-715k-1.5T"
rank: 32
tokens_per_block: 17 # how much info we can encode
dropout: 0.1
rank: 32 # lora rank
model_name: "PY007/TinyLlama-1.1B-intermediate-step-715k-1.5T"
vocab_size: 32000 # change to llm vocab dim
embed_dim: 2048 # change this to whatever the embedding dimension is in your pretrained llm 2048 for llama. 2560 for stablelm
+3
View File
@@ -354,3 +354,6 @@ idea
- yes I am bypassing it by passing in the input_embeds... but maybe I shouldn't
- [x] use same embedding everywhere. e.g. model embedding in encoder decoder?
- Our embedings is (embed_tokens): Embedding(32000, 2048). So we would need to encode to 32000!
ok we need to freeze it, and change dtype
+3 -3
View File
@@ -23,15 +23,15 @@ class TokenizerEncoderOutput:
class Tokenizer(nn.Module):
def __init__(self, embedding, 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 = embedding # nn.Embedding(vocab_size, embed_dim) # TODO: use model embed?
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:
+16 -13
View File
@@ -16,23 +16,27 @@ from .kv_caching import KeysValues, KVCache
@dataclass
class TransformerConfig:
max_blocks: int
num_layers: int
num_heads: int
embed_dim: int
tokens_per_block: int
# 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
dropout: float = 0.1
rank: int = 32
@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):
@@ -41,13 +45,14 @@ class Transformer(nn.Module):
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.model.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)
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) == self.config.num_layers
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,
@@ -60,8 +65,6 @@ class Transformer(nn.Module):
# 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()
# k_size = (x.shape[0], x.shape[1], x.shape[1], 1)
# v_size = past_keys_values[0]._v_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
+17 -6
View File
@@ -33,20 +33,26 @@ class WorldModel(nn.Module):
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)
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
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)
# nn.ReLU(),
# nn.Linear(config.embed_dim, config.embed_dim)
)
self.head_observations = Head(
@@ -79,9 +85,15 @@ 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)
self.transformer = Transformer(config)
def __repr__(self) -> str:
return "world_model"
@@ -92,7 +104,6 @@ class WorldModel(nn.Module):
assert num_steps <= self.config.max_tokens
prev_steps = 0 if past_keys_values is None else past_keys_values.size
# TODO: replace wth model embedder?
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)
+61 -21
View File
@@ -1,4 +1,4 @@
from functools import partial
from functools import partial
from pathlib import Path
import hydra
@@ -11,51 +11,91 @@ 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=cfg.tokenizer.vocab_size, act_vocab_size=test_env.num_actions, config=instantiate(cfg.world_model))
tokenizer = Tokenizer(embedding=embedding, **cfg.tokenizer.embedding)
actor_critic = ActorCritic(**cfg.actor_critic, act_vocab_size=test_env.num_actions)
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()
+8 -3
View File
@@ -81,10 +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=cfg.tokenizer.vocab_size, act_vocab_size=env.num_actions, config=instantiate(cfg.world_model))
embedding = world_model.transformer.embedding
tokenizer = Tokenizer(embedding=embedding, **cfg.tokenizer.embedding)
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')