MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / TransformerBlock

Class TransformerBlock

inference/models/llama.py:238–261  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

236
237
238class TransformerBlock(nn.Module):
239 def __init__(self, layer_id: int, args):
240 super().__init__()
241 self.n_heads = args.num_attention_heads
242 self.dim = args.hidden_size
243 self.head_dim = args.hidden_size // args.num_attention_heads
244 self.self_attn = LlamaAttentionFused(args)
245 self.mlp = LlamaMLP(args)
246 self.layer_id = layer_id
247 self.input_layernorm = RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
248 self.post_attention_layernorm = RMSNorm(args.hidden_size, eps=args.rms_norm_eps)
249
250 def forward(
251 self,
252 x: torch.Tensor,
253 start_pos: int,
254 freqs_cis: torch.Tensor,
255 mask: Optional[torch.Tensor],
256 ):
257 h = x + self.self_attn.forward(
258 self.input_layernorm(x), start_pos, freqs_cis, mask
259 )
260 out = h + self.mlp.forward(self.post_attention_layernorm(h))
261 return out
262
263
264class Transformer(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected