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

Class Transformer

inference/models/falcon.py:254–283  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

252
253
254class Transformer(nn.Module):
255 def __init__(self, params):
256 super().__init__()
257 self.params = params
258 self.vocab_size = params.vocab_size
259 self.n_layers = params.n_layer
260
261 self.word_embeddings = nn.Embedding(params.vocab_size, params.hidden_size)
262
263 self.h = torch.nn.ModuleList()
264 for layer_id in range(params.n_layer):
265 self.h.append(TransformerBlock(layer_id, params))
266
267 self.ln_f = nn.LayerNorm(params.hidden_size, eps=params.layer_norm_epsilon)
268
269 @torch.inference_mode()
270 def forward(self, tokens: torch.Tensor, start_pos: int):
271 _bsz, seqlen = tokens.shape
272 h = self.word_embeddings(tokens)
273
274 mask = None
275 if seqlen > 1:
276 mask = torch.full(
277 (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device
278 )
279 mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h)
280 for layer in self.h:
281 h = layer(h, start_pos, mask)
282 h = self.ln_f(h)
283 return h
284
285
286class FalconForCausalLM(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected