Perform a forward pass through the TransformerBlock. Args: x (torch.Tensor): Input tensor. start_pos (int): Starting position for attention caching. freqs_cis (torch.Tensor): Precomputed cosine and sine frequencies. mask (torch.Tensor
(
self,
x: torch.Tensor,
start_pos: int,
freqs_cis: torch.Tensor,
mask: Optional[torch.Tensor],
)
| 384 | self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps) |
| 385 | |
| 386 | def forward( |
| 387 | self, |
| 388 | x: torch.Tensor, |
| 389 | start_pos: int, |
| 390 | freqs_cis: torch.Tensor, |
| 391 | mask: Optional[torch.Tensor], |
| 392 | ): |
| 393 | """ |
| 394 | Perform a forward pass through the TransformerBlock. |
| 395 | |
| 396 | Args: |
| 397 | x (torch.Tensor): Input tensor. |
| 398 | start_pos (int): Starting position for attention caching. |
| 399 | freqs_cis (torch.Tensor): Precomputed cosine and sine frequencies. |
| 400 | mask (torch.Tensor, optional): Masking tensor for attention. Defaults to None. |
| 401 | |
| 402 | Returns: |
| 403 | torch.Tensor: Output tensor after applying attention and feedforward layers. |
| 404 | |
| 405 | """ |
| 406 | h = x + self.attention( |
| 407 | self.attention_norm(x), start_pos, freqs_cis, mask |
| 408 | ) |
| 409 | out = h + self.feed_forward(self.ffn_norm(h)) |
| 410 | return out |
| 411 | |
| 412 | |
| 413 | class Transformer(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected