| 196 | |
| 197 | |
| 198 | class MossBlock(Module): |
| 199 | def __init__(self, config): |
| 200 | super(MossBlock, self).__init__() |
| 201 | self.config = config |
| 202 | inner_dim = config.n_inner if config.n_inner is not None else 4 * config.n_embd |
| 203 | self.ln_1 = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon) |
| 204 | self.attn = MossAttention(config) |
| 205 | self.mlp = MossMLP(inner_dim, config) |
| 206 | |
| 207 | def execute( |
| 208 | self, |
| 209 | hidden_states: Optional[jt.Var], |
| 210 | layer_past: Optional[Tuple[jt.Var]] = None, |
| 211 | attention_mask: Optional[jt.Var] = None, |
| 212 | head_mask: Optional[jt.Var] = None, |
| 213 | use_cache: Optional[bool] = False, |
| 214 | ) -> Union[Tuple[jt.Var], Optional[Tuple[jt.Var, Tuple[jt.Var, ...]]]]: |
| 215 | residual = hidden_states |
| 216 | hidden_states = self.ln_1(hidden_states) |
| 217 | attn_outputs = self.attn( |
| 218 | hidden_states, |
| 219 | layer_past=layer_past, |
| 220 | attention_mask=attention_mask, |
| 221 | head_mask=head_mask, |
| 222 | use_cache=use_cache |
| 223 | ) |
| 224 | attn_output = attn_outputs[0] # output_attn: a, present |
| 225 | outputs = attn_outputs[1:] |
| 226 | |
| 227 | feed_forward_hidden_states = self.mlp(hidden_states) |
| 228 | hidden_states = attn_output + feed_forward_hidden_states + residual |
| 229 | |
| 230 | if use_cache: |
| 231 | outputs = (hidden_states,) + outputs |
| 232 | else: |
| 233 | outputs = (hidden_states,) + outputs[1:] |
| 234 | |
| 235 | return outputs # hidden_states, present |
| 236 | |
| 237 | |
| 238 | class MossModel(Module): |