(self, args: ModelArgs)
| 199 | |
| 200 | class 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, |
nothing calls this directly
no test coverage detected