mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-11 12:21:09 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user