| 131 | return y |
| 132 | |
| 133 | class SpaceBlock(nn.Module): |
| 134 | def __init__(self, config): |
| 135 | super().__init__() |
| 136 | self.ln1 = nn.LayerNorm(config.n_embd) |
| 137 | self.ln2 = nn.LayerNorm(config.n_embd) |
| 138 | self.attn = SpaceSelfAttention(config) |
| 139 | self.mlp = nn.Sequential( |
| 140 | nn.Linear(config.n_embd, 4 * config.n_embd, bias=False), |
| 141 | nn.GELU(), |
| 142 | nn.Linear(4 * config.n_embd, config.n_embd, bias=False), |
| 143 | nn.Dropout(config.resid_pdrop), |
| 144 | ) |
| 145 | |
| 146 | def forward(self, x, attn_mask): |
| 147 | attn = self.attn(self.ln1(x),attn_mask) |
| 148 | x = x + attn |
| 149 | x = x + self.mlp(self.ln2(x)) |
| 150 | return x |
| 151 | |
| 152 | class CausalTimeSelfAttention(nn.Module): |
| 153 | def __init__(self, config): |