MCPcopy Create free account
hub / github.com/OpenMOSS/MOSS / MossAttention

Class MossAttention

models_jittor/model.py:12–175  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

10 get_head_mask)
11
12class 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 ):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected