mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
1
This commit is contained in:
@@ -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,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
@@ -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!
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user