CoolFace
Modelpublic

He-Tag/smallm-125m

sourceHugging Faceapache-2.0updated 6d agoView on Hugging Face
2likes390downloads
modeling_smallm.py234 linesDownload Raw Back to root
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