| 252 | |
| 253 | |
| 254 | class Transformer(nn.Module): |
| 255 | def __init__(self, params): |
| 256 | super().__init__() |
| 257 | self.params = params |
| 258 | self.vocab_size = params.vocab_size |
| 259 | self.n_layers = params.n_layer |
| 260 | |
| 261 | self.word_embeddings = nn.Embedding(params.vocab_size, params.hidden_size) |
| 262 | |
| 263 | self.h = torch.nn.ModuleList() |
| 264 | for layer_id in range(params.n_layer): |
| 265 | self.h.append(TransformerBlock(layer_id, params)) |
| 266 | |
| 267 | self.ln_f = nn.LayerNorm(params.hidden_size, eps=params.layer_norm_epsilon) |
| 268 | |
| 269 | @torch.inference_mode() |
| 270 | def forward(self, tokens: torch.Tensor, start_pos: int): |
| 271 | _bsz, seqlen = tokens.shape |
| 272 | h = self.word_embeddings(tokens) |
| 273 | |
| 274 | mask = None |
| 275 | if seqlen > 1: |
| 276 | mask = torch.full( |
| 277 | (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device |
| 278 | ) |
| 279 | mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h) |
| 280 | for layer in self.h: |
| 281 | h = layer(h, start_pos, mask) |
| 282 | h = self.ln_f(h) |
| 283 | return h |
| 284 | |
| 285 | |
| 286 | class FalconForCausalLM(nn.Module): |