| 10 | get_head_mask) |
| 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)) |
| 46 | reshaped = reshaped.reshape(x.shape[:-2] + (-1,) + reshaped.shape[-1:]) |
| 47 | return reshaped |
| 48 | |
| 49 | def _merge_heads(self, tensor, num_attention_heads, attn_head_size): |
| 50 | """ |
| 51 | Merges attn_head_size dim and num_attn_heads dim into n_ctx |
| 52 | """ |
| 53 | if len(tensor.shape) == 5: |
| 54 | tensor = tensor.permute(0, 1, 3, 2, 4).contiguous() |
| 55 | elif len(tensor.shape) == 4: |
| 56 | tensor = tensor.permute(0, 2, 1, 3).contiguous() |
| 57 | else: |
| 58 | raise ValueError(f"Input tensor rank should be one of [4, 5], but is: {len(tensor.shape)}") |
| 59 | new_shape = tensor.size()[:-2] + (num_attention_heads * attn_head_size,) |
| 60 | return tensor.view(new_shape) |
| 61 | |
| 62 | def _attn( |
| 63 | self, |
| 64 | query, |
| 65 | key, |
| 66 | value, |
| 67 | attention_mask=None, |
| 68 | head_mask=None, |
| 69 | ): |