| 11 | |
| 12 | class MossAttention(Module): |
| 13 | def __init__(self, config): |
| 14 | super(MossAttention, self).__init__() |
| 15 | |
| 16 | max_positions = config.n_positions |
| 17 | self.register_buffer( |
| 18 | "causal_mask", |
| 19 | jt.tril(jt.ones((max_positions, max_positions), dtype=jt.bool)).view( |
| 20 | 1, 1, max_positions, max_positions |
| 21 | ), |
| 22 | ) |
| 23 | |
| 24 | self.attn_dropout = nn.Dropout(config.attn_pdrop) |
| 25 | self.resid_dropout = nn.Dropout(config.resid_pdrop) |
| 26 | |
| 27 | self.embed_dim = config.n_embd |
| 28 | self.num_attention_heads = config.n_head |
| 29 | self.head_dim = self.embed_dim // self.num_attention_heads |
| 30 | if self.head_dim * self.num_attention_heads != self.embed_dim: |
| 31 | raise ValueError( |
| 32 | f"embed_dim must be divisible by num_attention_heads (got `embed_dim`: {self.embed_dim} and" |
| 33 | f" `num_attention_heads`: {self.num_attention_heads})." |
| 34 | ) |
| 35 | self.scale_attn = jt.sqrt(jt.float32(self.head_dim)) |
| 36 | self.qkv_proj = nn.Linear(self.embed_dim, self.embed_dim * 3, bias=False) |
| 37 | jt.float16 |
| 38 | |
| 39 | self.out_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False) |
| 40 | self.rotary_dim = None |
| 41 | if config.rotary_dim is not None: |
| 42 | self.rotary_dim = config.rotary_dim |
| 43 | |
| 44 | def _split_heads(self, x, n_head, dim_head, mp_num): |
| 45 | reshaped = x.reshape(x.shape[:-1] + (n_head // mp_num, dim_head)) |