This commit is contained in:
wassname
2023-11-19 08:48:36 +08:00
parent 36ea66d0ae
commit 7098b67955
7 changed files with 17 additions and 12 deletions
+3 -3
View File
@@ -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
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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!
+2 -2
View File
@@ -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)
+1
View File
@@ -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)
+3 -2
View File
@@ -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)
+5 -2
View File
@@ -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')