From 7098b679559a7e696540d8b698fbd5921f67d28f Mon Sep 17 00:00:00 2001 From: wassname Date: Sun, 19 Nov 2023 08:48:36 +0800 Subject: [PATCH] 1 --- config/tokenizer/default.yaml | 6 +++--- config/world_model/default.yaml | 4 ++-- research_journal.md | 2 +- src/models/tokenizer/tokenizer.py | 4 ++-- src/models/world_model.py | 1 + src/play.py | 5 +++-- src/trainer.py | 7 +++++-- 7 files changed, 17 insertions(+), 12 deletions(-) diff --git a/config/tokenizer/default.yaml b/config/tokenizer/default.yaml index 64e2bac..2df2b4d 100644 --- a/config/tokenizer/default.yaml +++ b/config/tokenizer/default.yaml @@ -1,14 +1,14 @@ _target_: src.models.tokenizer.Tokenizer -vocab_size: 2048 -embed_dim: 2048 +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 encoder: _target_: src.models.tokenizer.Encoder config: _target_: src.models.tokenizer.EncoderDecoderConfig resolution: 64 in_channels: 3 - z_channels: 2048 + z_channels: 32000 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 08ede7e..76894ac 100644 --- a/config/world_model/default.yaml +++ b/config/world_model/default.yaml @@ -2,8 +2,8 @@ _target_: src.models.TransformerConfig max_blocks: 10 # this is the rollout length when training policy num_layers: 1 num_heads: 1 -embed_dim: 2048 # change this to whatever the embedding dimension is in your pretrained llm 2048 for llama. 2560 for stablelm +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 +tokens_per_block: 17 # how much info we can encode diff --git a/research_journal.md b/research_journal.md index b7e2a98..d7bbb80 100644 --- a/research_journal.md +++ b/research_journal.md @@ -352,5 +352,5 @@ So I tried just trainign the world model for 200 epochs. And with a post_embeddi idea - bypass embedding?, but wait dreamerv3 needed quant z... - yes I am bypassing it by passing in the input_embeds... but maybe I shouldn't -- use same embedding everywhere. e.g. model embedding in encoder decoder? +- [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! diff --git a/src/models/tokenizer/tokenizer.py b/src/models/tokenizer/tokenizer.py index e69c18e..332c1fe 100644 --- a/src/models/tokenizer/tokenizer.py +++ b/src/models/tokenizer/tokenizer.py @@ -23,12 +23,12 @@ 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, 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) # TODO: use model embed? + self.embedding = embedding # nn.Embedding(vocab_size, embed_dim) # TODO: use model embed? 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) diff --git a/src/models/world_model.py b/src/models/world_model.py index 723b733..cb3e798 100644 --- a/src/models/world_model.py +++ b/src/models/world_model.py @@ -92,6 +92,7 @@ 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 7e3ab85..28c3f07 100644 --- a/src/play.py +++ b/src/play.py @@ -33,8 +33,9 @@ def main(cfg: DictConfig): 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)) + # 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) agent = Agent(tokenizer, world_model, actor_critic).to(device) agent.load(Path('checkpoints/last.pt'), device) diff --git a/src/trainer.py b/src/trainer.py index af0460c..c337d1f 100644 --- a/src/trainer.py +++ b/src/trainer.py @@ -22,6 +22,7 @@ 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: @@ -80,8 +81,10 @@ 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)) + # 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) 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')