| 255 | |
| 256 | |
| 257 | class Transformer(nn.Module): |
| 258 | def __init__(self, params): |
| 259 | super().__init__() |
| 260 | self.params = params |
| 261 | self.vocab_size = params.vocab_size |
| 262 | self.n_layers = params.n_layers |
| 263 | |
| 264 | self.wte = SharedEmbedding(params.vocab_size, params.d_model) |
| 265 | |
| 266 | self.blocks = torch.nn.ModuleList() |
| 267 | for layer_id in range(params.n_layers): |
| 268 | self.blocks.append(MPTBlock(layer_id, params)) |
| 269 | |
| 270 | self.norm_f = LPLayerNorm(params.d_model, eps=1e-6) |
| 271 | |
| 272 | @torch.inference_mode() |
| 273 | def forward(self, tokens: torch.Tensor, start_pos: int): |
| 274 | _bsz, seqlen = tokens.shape |
| 275 | h = self.wte(tokens) |
| 276 | |
| 277 | mask = None |
| 278 | if seqlen > 1: |
| 279 | mask = torch.full( |
| 280 | (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device |
| 281 | ) |
| 282 | mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h) |
| 283 | for layer in self.blocks: |
| 284 | h = layer(h, start_pos, mask) |
| 285 | h = self.norm_f(h) |
| 286 | return h |
| 287 | |
| 288 | |
| 289 | class MPTForCausalLM(nn.Module): |