diff --git a/.vscode/launch.json b/.vscode/launch.json index fda7a32..ac43a5c 100644 --- a/.vscode/launch.json +++ b/.vscode/launch.json @@ -16,9 +16,9 @@ "args": [ "env.train.id=BreakoutNoFrameskip-v4", // # make it start early - "trainer.training.tokenizer.start_after_epochs=1", - "trainer.training.world_model.start_after_epochs=2", - "trainer.training.actor_critic.start_after_epochs=3", + "training.tokenizer.start_after_epochs=1", + "training.world_model.start_after_epochs=2", + "training.actor_critic.start_after_epochs=3", ] } ] diff --git a/config/trainer.yaml b/config/trainer.yaml index fe220a7..bd02bf0 100644 --- a/config/trainer.yaml +++ b/config/trainer.yaml @@ -7,7 +7,7 @@ defaults: - datasets: default wandb: - mode: online + mode: offline project: iris entity: null name: null diff --git a/src/models/transformer.py b/src/models/transformer.py index 2aaac69..dea058d 100644 --- a/src/models/transformer.py +++ b/src/models/transformer.py @@ -32,29 +32,6 @@ class TransformerConfig: def max_tokens(self): return self.tokens_per_block * self.max_blocks - -class Transformer(nn.Module): - def __init__(self, config: TransformerConfig) -> None: - super().__init__() - self.config = config - self.drop = nn.Dropout(config.embed_pdrop) - self.blocks = nn.ModuleList([Block(config) for _ in range(config.num_layers)]) - self.ln_f = nn.LayerNorm(config.embed_dim) - - 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 - return KeysValues(n, self.config.num_heads, max_tokens, self.config.embed_dim, self.config.num_layers, device) - - def forward(self, sequences: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor: - assert past_keys_values is None or len(past_keys_values) == len(self.blocks) - x = self.drop(sequences) - for i, block in enumerate(self.blocks): - x = block(x, None if past_keys_values is None else past_keys_values[i]) - - x = self.ln_f(x) - return x - - class Block(nn.Module): def __init__(self, config: TransformerConfig) -> None: super().__init__() @@ -118,3 +95,24 @@ class SelfAttention(nn.Module): y = self.resid_drop(self.proj(y)) return y + +class Transformer(nn.Module): + def __init__(self, config: TransformerConfig) -> None: + super().__init__() + self.config = config + self.drop = nn.Dropout(config.embed_pdrop) + self.blocks = nn.ModuleList([Block(config) for _ in range(config.num_layers)]) + self.ln_f = nn.LayerNorm(config.embed_dim) + + 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 + return KeysValues(n, self.config.num_heads, max_tokens, self.config.embed_dim, self.config.num_layers, device) + + def forward(self, sequences: torch.Tensor, past_keys_values: Optional[KeysValues] = None) -> torch.Tensor: + assert past_keys_values is None or len(past_keys_values) == len(self.blocks) + x = self.drop(sequences) + for i, block in enumerate(self.blocks): + x = block(x, None if past_keys_values is None else past_keys_values[i]) + + x = self.ln_f(x) + return x