MCPcopy Create free account
hub / github.com/microsoft/BitNet / Transformer

Class Transformer

gpu/model.py:246–296  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

244
245
246class Transformer(nn.Module):
247 def __init__(self, args: ModelArgs):
248 super().__init__()
249 assert args.vocab_size > 0
250
251 self.tok_embeddings = nn.Embedding(
252 num_embeddings=args.vocab_size,
253 embedding_dim=args.dim,
254 )
255
256 self.layers = nn.ModuleList()
257 for _ in range(args.n_layers):
258 self.layers.append(TransformerBlock(args))
259
260 self.norm = RMSNorm(args.dim, eps=args.norm_eps)
261
262 self.output = nn.Linear(
263 args.dim,
264 args.vocab_size,
265 bias=False,
266 )
267
268 @torch.no_grad()
269 def forward_with_attn_bias(
270 self,
271 token_values: torch.Tensor,
272 attn_bias: AttnBias,
273 cache: list[LayerCache],
274 ) -> torch.Tensor:
275 h = self.tok_embeddings(token_values)
276
277 for i, layer in enumerate(self.layers):
278 h = layer(h, cache[i], attn_bias)
279
280 logits = self.output(self.norm(h))
281 return logits.float()
282
283 def forward(
284 self,
285 token_values: torch.Tensor,
286 token_lengths: torch.Tensor,
287 start_pos: torch.Tensor,
288 cache: list[LayerCache],
289 kv_padding: int,
290 ) -> torch.Tensor:
291 attn_bias = AttnBias.from_seqlens(
292 q_seqlen=token_lengths.tolist(),
293 kv_seqlen=(start_pos + token_lengths).tolist(),
294 kv_padding=kv_padding,
295 )
296 return self.forward_with_attn_bias(token_values, attn_bias, cache)
297
298
299def make_cache(

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected