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

Method __init__

models_jittor/model.py:13–42  ·  view source on GitHub ↗
(self, config)

Source from the content-addressed store, hash-verified

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))

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected