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

Class Transformer

inference/models/llama.py:264–300  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

262
263
264class Transformer(nn.Module):
265 def __init__(self, params):
266 super().__init__()
267 self.params = params
268 self.vocab_size = params.vocab_size
269 self.n_layers = params.num_hidden_layers
270
271 self.embed_tokens = nn.Embedding(params.vocab_size, params.hidden_size)
272
273 self.layers = torch.nn.ModuleList()
274 for layer_id in range(params.num_hidden_layers):
275 self.layers.append(TransformerBlock(layer_id, params))
276
277 self.norm = RMSNorm(params.hidden_size, eps=params.rms_norm_eps)
278
279 self.freqs_cis = precompute_freqs_cis(
280 self.params.hidden_size // self.params.num_attention_heads,
281 self.params.max_position_embeddings * 2,
282 )
283
284 @torch.inference_mode()
285 def forward(self, tokens: torch.Tensor, start_pos: int):
286 _bsz, seqlen = tokens.shape
287 h = self.embed_tokens(tokens)
288 self.freqs_cis = self.freqs_cis.to(h.device)
289 freqs_cis = self.freqs_cis[start_pos : start_pos + seqlen]
290
291 mask = None
292 if seqlen > 1:
293 mask = torch.full(
294 (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device
295 )
296 mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h)
297 for layer in self.layers:
298 h = layer(h, start_pos, freqs_cis, mask)
299 h = self.norm(h)
300 return h
301
302
303class LlamaForCausalLM(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected