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

Class Transformer

inference/models/mpt.py:257–286  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

255
256
257class Transformer(nn.Module):
258 def __init__(self, params):
259 super().__init__()
260 self.params = params
261 self.vocab_size = params.vocab_size
262 self.n_layers = params.n_layers
263
264 self.wte = SharedEmbedding(params.vocab_size, params.d_model)
265
266 self.blocks = torch.nn.ModuleList()
267 for layer_id in range(params.n_layers):
268 self.blocks.append(MPTBlock(layer_id, params))
269
270 self.norm_f = LPLayerNorm(params.d_model, eps=1e-6)
271
272 @torch.inference_mode()
273 def forward(self, tokens: torch.Tensor, start_pos: int):
274 _bsz, seqlen = tokens.shape
275 h = self.wte(tokens)
276
277 mask = None
278 if seqlen > 1:
279 mask = torch.full(
280 (1, 1, seqlen, seqlen), float("-inf"), device=tokens.device
281 )
282 mask = torch.triu(mask, diagonal=start_pos + 1).type_as(h)
283 for layer in self.blocks:
284 h = layer(h, start_pos, mask)
285 h = self.norm_f(h)
286 return h
287
288
289class MPTForCausalLM(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected