He-Tag/smallm-125m
2390
1import torch2import torch.nn as nn3import torch.nn.functional as F4from transformers import GenerationMixin, PreTrainedModel5from transformers.cache_utils import Cache, DynamicCache6from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast7 8from .configuration_smallm import SmallmConfig9 10 11def rms_norm(x):12 return F.rms_norm(x, (x.size(-1),))13 14 15class RotaryEmbedding(nn.Module):16 def __init__(self, head_dim, max_seq_len, base):17 super().__init__()18 self.head_dim = head_dim19 self.max_seq_len = max_seq_len20 self.base = base21 self.tables = None22 23 def angles(self, device):24 if self.tables is None or self.tables[0].device != device:25 steps = torch.arange(0, self.head_dim, 2, dtype=torch.float32, device=device)26 freqs = self.base ** (-steps / self.head_dim)27 positions = torch.arange(self.max_seq_len, dtype=torch.float32, device=device)28 angles = torch.outer(positions, freqs)29 self.tables = (angles.cos(), angles.sin())30 return self.tables31 32 def forward(self, x, offset=0):33 length = x.size(1)34 cos_table, sin_table = self.angles(x.device)35 cos = cos_table[offset:offset + length].view(1, length, 1, -1)36 sin = sin_table[offset:offset + length].view(1, length, 1, -1)37 left, right = x.float().chunk(2, dim=-1)38 rotated = torch.cat([left * cos - right * sin, left * sin + right * cos], dim=-1)39 return rotated.type_as(x)40 41 42class Attention(nn.Module):43 def __init__(self, config, layer_idx):44 super().__init__()45 self.layer_idx = layer_idx46 self.n_head = config.n_head47 self.head_dim = config.dim // config.n_head48 self.qkv = nn.Parameter(torch.empty(3, config.dim, config.dim))49 self.out = nn.Parameter(torch.empty(config.dim, config.dim))50 self.value_lambda = nn.Parameter(torch.tensor(0.5))51 52 def forward(self, x, first_value, rotary, attention_mask, past_key_values, offset):53 batch, length, dim = x.shape54 projected = F.linear(x, self.qkv.flatten(end_dim=1))55 q, k, v = projected.view(batch, length, 3, self.n_head, self.head_dim).unbind(dim=2)56 q, k = rms_norm(q), rms_norm(k)57 q, k = rotary(q, offset), rotary(k, offset)58 if first_value is None:59 first_value = v60 v = torch.lerp(v, first_value, self.value_lambda)61 q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)62 if past_key_values is not None:63 k, v = past_key_values.update(k, v, self.layer_idx)64 attended = F.scaled_dot_product_attention(65 q, k, v,66 attn_mask=attention_mask,67 is_causal=attention_mask is None and q.size(2) == k.size(2),68 )69 attended = attended.transpose(1, 2).reshape(batch, length, dim)70 return F.linear(attended, self.out), first_value71 72 73class Mlp(nn.Module):74 def __init__(self, config):75 super().__init__()76 self.relu2_clamp = config.relu2_clamp77 self.up = nn.Parameter(torch.empty(config.mlp_mult * config.dim, config.dim))78 self.down = nn.Parameter(torch.empty(config.dim, config.mlp_mult * config.dim))79 80 def forward(self, x):81 hidden = F.relu(F.linear(x, self.up))82 return F.linear(hidden.clamp(max=self.relu2_clamp).square(), self.down)83 84 85class Block(nn.Module):86 def __init__(self, config, layer_idx):87 super().__init__()88 self.attention = Attention(config, layer_idx)89 self.mlp = Mlp(config)90 self.residual_lambda = nn.Parameter(torch.tensor([1.0, 0.0]))91 92 def forward(self, x, first_value, embedded, rotary, attention_mask, past_key_values, offset):93 x = self.residual_lambda[0] * x + self.residual_lambda[1] * embedded94 attended, first_value = self.attention(95 rms_norm(x), first_value, rotary, attention_mask, past_key_values, offset96 )97 x = x + attended98 x = x + self.mlp(rms_norm(x))99 return x, first_value100 101 102class SmallmPreTrainedModel(PreTrainedModel):103 config_class = SmallmConfig104 base_model_prefix = "model"105 supports_gradient_checkpointing = False106 _no_split_modules = ["Block"]107 108 def _init_weights(self, module):109 bound = self.config.dim ** -0.5110 if isinstance(module, Attention):111 module.qkv.data.uniform_(-bound, bound)112 module.out.data.zero_()113 module.value_lambda.data.fill_(0.5)114 elif isinstance(module, Mlp):115 module.up.data.uniform_(-bound, bound)116 module.down.data.zero_()117 elif isinstance(module, Block):118 module.residual_lambda.data.copy_(torch.tensor([1.0, 0.0]))119 elif isinstance(module, nn.Embedding):120 module.weight.data.uniform_(-0.5 * bound, 0.5 * bound)121 elif isinstance(module, nn.Linear):122 module.weight.data.zero_()123 124 125class SmallmModel(SmallmPreTrainedModel):126 def __init__(self, config):127 super().__init__(config)128 self.embed_tokens = nn.Embedding(config.vocab_size, config.dim)129 self.rotary = RotaryEmbedding(config.dim // config.n_head, config.max_seq_len,130 config.rope_theta)131 self.layers = nn.ModuleList([Block(config, i) for i in range(config.n_layer)])132 self.n_skip = config.n_layer // 2133 self.skip_weights = nn.Parameter(torch.ones(self.n_skip))134 self.post_init()135 136 def get_input_embeddings(self):137 return self.embed_tokens138 139 def set_input_embeddings(self, value):140 self.embed_tokens = value141 142 def forward(self, input_ids=None, attention_mask=None, past_key_values=None,143 inputs_embeds=None, use_cache=None, **kwargs):144 if inputs_embeds is None:145 inputs_embeds = self.embed_tokens(input_ids)146 use_cache = self.config.use_cache if use_cache is None else use_cache147 if use_cache and past_key_values is None:148 past_key_values = DynamicCache()149 length = inputs_embeds.size(1)150 offset = past_key_values.get_seq_length() if past_key_values is not None else 0151 mask = self.build_attention_mask(attention_mask, length, offset + length,152 inputs_embeds.device)153 154 x = rms_norm(inputs_embeds)155 embedded, first_value = x, None156 skipped = []157 for i, layer in enumerate(self.layers):158 if i >= self.n_skip:159 x = x + self.skip_weights[i - self.n_skip] * skipped.pop()160 x, first_value = layer(x, first_value, embedded, self.rotary, mask,161 past_key_values if use_cache else None, offset)162 if i < self.n_skip:163 skipped.append(x)164 return BaseModelOutputWithPast(165 last_hidden_state=rms_norm(x),166 past_key_values=past_key_values if use_cache else None,167 )168 169 @staticmethod170 def build_attention_mask(attention_mask, query_length, key_length, device):171 if attention_mask is not None and attention_mask.dim() == 4:172 return attention_mask173 padded = attention_mask is not None and not bool(attention_mask.all())174 if not padded and query_length == key_length:175 return None176 keys = torch.arange(key_length, device=device)177 queries = torch.arange(key_length - query_length, key_length, device=device)178 allowed = (keys[None, :] <= queries[:, None])[None]179 if attention_mask is not None:180 allowed = allowed & attention_mask[:, None, :].bool()181 allowed = allowed | (queries[None, :, None] == keys[None, None, :])182 return allowed.unsqueeze(1)183 184 185class SmallmForCausalLM(SmallmPreTrainedModel, GenerationMixin):186 _tied_weights_keys = []187 188 def __init__(self, config):189 super().__init__(config)190 self.model = SmallmModel(config)191 self.lm_head = nn.Linear(config.dim, config.vocab_size, bias=False)192 self.post_init()193 194 def get_input_embeddings(self):195 return self.model.embed_tokens196 197 def set_input_embeddings(self, value):198 self.model.embed_tokens = value199 200 def get_output_embeddings(self):201 return self.lm_head202 203 def forward(self, input_ids=None, attention_mask=None, past_key_values=None,204 inputs_embeds=None, labels=None, use_cache=None, logits_to_keep=0, **kwargs):205 outputs = self.model(206 input_ids=input_ids,207 attention_mask=attention_mask,208 past_key_values=past_key_values,209 inputs_embeds=inputs_embeds,210 use_cache=use_cache,211 )212 hidden = outputs.last_hidden_state213 if isinstance(logits_to_keep, int) and logits_to_keep > 0:214 hidden = hidden[:, -logits_to_keep:]215 elif isinstance(logits_to_keep, torch.Tensor):216 hidden = hidden[:, logits_to_keep]217 logits = self.lm_head(hidden)218 logits = self.config.softcap * torch.tanh(logits.float() / self.config.softcap)219 220 loss = None221 if labels is not None:222 loss = F.cross_entropy(223 logits[:, :-1].reshape(-1, logits.size(-1)),224 labels[:, 1:].reshape(-1).to(logits.device),225 )226 return CausalLMOutputWithPast(227 loss=loss,228 logits=logits,229 past_key_values=outputs.past_key_values,230 )231 232 233__all__ = ["SmallmConfig", "SmallmModel", "SmallmForCausalLM", "SmallmPreTrainedModel"]234 