fix debug launch

This commit is contained in:
wassname
2023-11-12 14:44:10 +08:00
parent 7dfded5bde
commit a6b850bd96
3 changed files with 25 additions and 27 deletions
+3 -3
View File
@@ -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",
]
}
]
+1 -1
View File
@@ -7,7 +7,7 @@ defaults:
- datasets: default
wandb:
mode: online
mode: offline
project: iris
entity: null
name: null
+21 -23
View File
@@ -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