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

Class TransformerBlock

gpu/model.py:200–243  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

198
199
200class TransformerBlock(nn.Module):
201 def __init__(self, args: ModelArgs):
202 super().__init__()
203
204 assert args.dim % args.n_heads == 0
205 head_dim = args.dim // args.n_heads
206 if args.n_kv_heads is not None:
207 n_kv_heads = args.n_kv_heads
208 else:
209 n_kv_heads = args.n_heads
210
211 assert args.n_heads % n_kv_heads == 0
212
213 self.attention = Attention(
214 dim=args.dim,
215 head_dim=head_dim,
216 n_heads=args.n_heads,
217 n_kv_heads=n_kv_heads,
218 rope_theta=args.rope_theta,
219 norm_eps=args.norm_eps,
220 use_kernel=args.use_kernel,
221 )
222 self.feed_forward = FeedForward(
223 dim=args.dim,
224 hidden_dim=args.ffn_dim,
225 norm_eps=args.norm_eps,
226 use_kernel=args.use_kernel,
227 )
228 self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
229 self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)
230
231 def forward(
232 self,
233 x: torch.Tensor,
234 cache: LayerCache,
235 attn_bias: AttnBias,
236 ) -> torch.Tensor:
237 h = x + self.attention.forward(
238 self.attention_norm(x),
239 cache,
240 attn_bias,
241 )
242 out = h + self.feed_forward(self.ffn_norm(h))
243 return out
244
245
246class Transformer(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected