| 244 | |
| 245 | |
| 246 | class Transformer(nn.Module): |
| 247 | def __init__(self, args: ModelArgs): |
| 248 | super().__init__() |
| 249 | assert args.vocab_size > 0 |
| 250 | |
| 251 | self.tok_embeddings = nn.Embedding( |
| 252 | num_embeddings=args.vocab_size, |
| 253 | embedding_dim=args.dim, |
| 254 | ) |
| 255 | |
| 256 | self.layers = nn.ModuleList() |
| 257 | for _ in range(args.n_layers): |
| 258 | self.layers.append(TransformerBlock(args)) |
| 259 | |
| 260 | self.norm = RMSNorm(args.dim, eps=args.norm_eps) |
| 261 | |
| 262 | self.output = nn.Linear( |
| 263 | args.dim, |
| 264 | args.vocab_size, |
| 265 | bias=False, |
| 266 | ) |
| 267 | |
| 268 | @torch.no_grad() |
| 269 | def forward_with_attn_bias( |
| 270 | self, |
| 271 | token_values: torch.Tensor, |
| 272 | attn_bias: AttnBias, |
| 273 | cache: list[LayerCache], |
| 274 | ) -> torch.Tensor: |
| 275 | h = self.tok_embeddings(token_values) |
| 276 | |
| 277 | for i, layer in enumerate(self.layers): |
| 278 | h = layer(h, cache[i], attn_bias) |
| 279 | |
| 280 | logits = self.output(self.norm(h)) |
| 281 | return logits.float() |
| 282 | |
| 283 | def forward( |
| 284 | self, |
| 285 | token_values: torch.Tensor, |
| 286 | token_lengths: torch.Tensor, |
| 287 | start_pos: torch.Tensor, |
| 288 | cache: list[LayerCache], |
| 289 | kv_padding: int, |
| 290 | ) -> torch.Tensor: |
| 291 | attn_bias = AttnBias.from_seqlens( |
| 292 | q_seqlen=token_lengths.tolist(), |
| 293 | kv_seqlen=(start_pos + token_lengths).tolist(), |
| 294 | kv_padding=kv_padding, |
| 295 | ) |
| 296 | return self.forward_with_attn_bias(token_values, attn_bias, cache) |
| 297 | |
| 298 | |
| 299 | def make_cache( |
nothing calls this directly
no outgoing calls
no test coverage detected