mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-04 16:24:07 +08:00
Fix embedding path in Transformer model
This commit is contained in:
@@ -46,7 +46,7 @@ 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
|
||||
self.embedding = freeze(self.model.base_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
|
||||
|
||||
+1
-1
@@ -149,11 +149,11 @@ class Trainer:
|
||||
|
||||
if epoch > cfg_tokenizer.start_after_epochs:
|
||||
metrics_tokenizer = self.train_component(self.agent.tokenizer, self.optimizer_tokenizer, sequence_length=1, sample_from_start=True, **cfg_tokenizer)
|
||||
self.agent.tokenizer.eval()
|
||||
|
||||
if epoch > cfg_world_model.start_after_epochs:
|
||||
metrics_world_model = self.train_component(self.agent.world_model, self.optimizer_world_model, sequence_length=self.cfg.common.sequence_length, sample_from_start=True, tokenizer=self.agent.tokenizer, **cfg_world_model)
|
||||
self.agent.world_model.eval()
|
||||
self.agent.tokenizer.eval()
|
||||
|
||||
if epoch > cfg_actor_critic.start_after_epochs:
|
||||
metrics_actor_critic = self.train_component(self.agent.actor_critic, self.optimizer_actor_critic, sequence_length=1 + self.cfg.training.actor_critic.burn_in, sample_from_start=False, tokenizer=self.agent.tokenizer, world_model=self.agent.world_model, **cfg_actor_critic)
|
||||
|
||||
Reference in New Issue
Block a user