diff --git a/config/tokenizer/default.yaml b/config/tokenizer/default.yaml index 2df2b4d..9029ef4 100644 --- a/config/tokenizer/default.yaml +++ b/config/tokenizer/default.yaml @@ -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 diff --git a/config/world_model/default.yaml b/config/world_model/default.yaml index 76894ac..2b631a9 100644 --- a/config/world_model/default.yaml +++ b/config/world_model/default.yaml @@ -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 diff --git a/research_journal.md b/research_journal.md index d7bbb80..fd40272 100644 --- a/research_journal.md +++ b/research_journal.md @@ -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 diff --git a/src/models/tokenizer/tokenizer.py b/src/models/tokenizer/tokenizer.py index 332c1fe..2504b42 100644 --- a/src/models/tokenizer/tokenizer.py +++ b/src/models/tokenizer/tokenizer.py @@ -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: diff --git a/src/models/transformer.py b/src/models/transformer.py index 89a2b25..22a949a 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -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 diff --git a/src/models/world_model.py b/src/models/world_model.py index cb3e798..9999ff4 100644 --- a/src/models/world_model.py +++ b/src/models/world_model.py @@ -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) diff --git a/src/play.py b/src/play.py index 28c3f07..e3486a5 100644 --- a/src/play.py +++ b/src/play.py @@ -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() diff --git a/src/trainer.py b/src/trainer.py index c337d1f..2b4e1b5 100644 --- a/src/trainer.py +++ b/src/trainer.py @@ -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')