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

Class MossBlock

models_jittor/model.py:198–235  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

196
197
198class 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
238class MossModel(Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected