| 225 | |
| 226 | |
| 227 | class TransformerBlock(nn.Module): |
| 228 | def __init__(self, layer_id: int, args): |
| 229 | super().__init__() |
| 230 | self.n_heads = args.n_head |
| 231 | self.dim = args.hidden_size |
| 232 | self.head_dim = args.hidden_size // args.n_head |
| 233 | self.self_attention = FalconAttentionFused(args) |
| 234 | self.mlp = FalconMLP(dim=args.hidden_size) |
| 235 | self.layer_id = layer_id |
| 236 | self.input_layernorm = nn.LayerNorm( |
| 237 | args.hidden_size, eps=args.layer_norm_epsilon |
| 238 | ) |
| 239 | # self.post_attention_layernorm = nn.LayerNorm(args.dim, eps=args.norm_eps) |
| 240 | |
| 241 | def forward( |
| 242 | self, |
| 243 | x: torch.Tensor, |
| 244 | start_pos: int, |
| 245 | mask: Optional[torch.Tensor], |
| 246 | ): |
| 247 | layernorm_output = self.input_layernorm(x) |
| 248 | h_attn = x + self.self_attention.forward(layernorm_output, start_pos, mask) |
| 249 | h_mlp = self.mlp(layernorm_output) |
| 250 | out = h_attn + h_mlp |
| 251 | return out |
| 252 | |
| 253 | |
| 254 | class Transformer(nn.Module): |