| from torch import nn | |
| from core.norm import RMSNorm | |
| from core.attention import Attention | |
| from core.mlp import FeedForward, MOEFeedForward | |
| class Block(nn.Module): | |
| def __init__(self, layer_id: int, config: "LMConfig"): | |
| super().__init__() | |
| self.self_attn = Attention(config) | |
| self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) | |
| self.mlp = FeedForward(config) if not config.use_moe else MOEFeedForward(config) | |
| def forward(self, hidden_states, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None): | |
| residual = hidden_states | |
| hidden_states, present_key_value = self.self_attn( | |
| self.input_layernorm(hidden_states), position_embeddings, | |
| past_key_value, use_cache, attention_mask | |
| ) | |
| hidden_states += residual | |
| hidden_states = hidden_states + self.mlp(self.post_attention_layernorm(hidden_states)) | |
| return hidden_states, present_key_value | |