Fix embedding path in Transformer model

This commit is contained in:
wassname
2023-11-24 09:24:12 +08:00
parent c3878e4a4c
commit 9155ecca95
2 changed files with 2 additions and 2 deletions
+1 -1
View File
@@ -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
View File
@@ -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)