mirror of
https://github.com/wassname/iris_bigvae.git
synced 2026-09-09 11:24:31 +08:00
fix debug launch
This commit is contained in:
Vendored
+3
-3
@@ -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
@@ -7,7 +7,7 @@ defaults:
|
||||
- datasets: default
|
||||
|
||||
wandb:
|
||||
mode: online
|
||||
mode: offline
|
||||
project: iris
|
||||
entity: null
|
||||
name: null
|
||||
|
||||
+21
-23
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user