CoolFace
Modelpublic

Krishna0812/Tiny_Stories

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes64downloads
modeling_tiny.py446 linesDownload Raw Back to root
1import math2import sys3from typing import Optional, Tuple, Union4 5# Monkeypatch safetensors to handle None/missing metadata which crashes transformers6try:7    import safetensors8    original_safe_open = safetensors.safe_open9 10    class SafeOpenWrapper:11        def __init__(self, original_obj):12            self.original_obj = original_obj13 14        def __enter__(self):15            self.original_obj.__enter__()16            return self17 18        def __exit__(self, exc_type, exc_val, exc_tb):19            return self.original_obj.__exit__(exc_type, exc_val, exc_tb)20 21        def metadata(self):22            meta = self.original_obj.metadata()23            if meta is None:24                return {"format": "pt"}25            return meta26 27        def __getattr__(self, name):28            return getattr(self.original_obj, name)29 30    def patched_safe_open(*args, **kwargs):31        f = original_safe_open(*args, **kwargs)32        return SafeOpenWrapper(f)33 34    safetensors.safe_open = patched_safe_open35 36    if "transformers.modeling_utils" in sys.modules:37        import transformers.modeling_utils38        transformers.modeling_utils.safe_open = patched_safe_open39except Exception:40    pass41 42 43import torch44import torch.nn as nn45import torch.nn.functional as F46 47from transformers import PreTrainedModel48from transformers.modeling_outputs import (49    BaseModelOutputWithPast,50    CausalLMOutputWithPast,51)52 53from .configuration_tiny import TinyConfig54 55 56def rotate_half(x: torch.Tensor) -> torch.Tensor:57    x1 = x[..., :x.shape[-1] // 2]58    x2 = x[..., x.shape[-1] // 2:]59    return torch.cat((-x2, x1), dim=-1)60 61 62def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):63    assert dim % 2 == 064    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))65    t = torch.arange(end)66    freqs = torch.outer(t, freqs).float()67    cos = torch.cos(freqs)68    sin = torch.sin(freqs)69    cos = torch.cat([cos, cos], dim=-1)70    sin = torch.cat([sin, sin], dim=-1)71    return cos, sin72 73 74def apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:75    T = x.shape[2]76    cos_t = cos[:T, :].unsqueeze(0).unsqueeze(1)77    sin_t = sin[:T, :].unsqueeze(0).unsqueeze(1)78    return (x * cos_t) + (rotate_half(x) * sin_t)79 80 81class RMSNorm(nn.Module):82    def __init__(self, dim: int, eps: float = 1e-5):83        super().__init__()84        self.eps = eps85        self.weight = nn.Parameter(torch.ones(dim))86 87    def forward(self, x: torch.Tensor) -> torch.Tensor:88        variance = x.pow(2).mean(-1, keepdim=True)89        return x * torch.rsqrt(variance + self.eps) * self.weight90 91 92class FeedForward(nn.Module):93    def __init__(self, config: TinyConfig):94        super().__init__()95        self.w1 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)96        self.w2 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)97        self.w3 = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)98        self.dropout = nn.Dropout(config.hidden_dropout) if config.hidden_dropout > 0.0 else None99 100    def forward(self, x: torch.Tensor) -> torch.Tensor:101        out = F.silu(self.w1(x)) * self.w2(x)102        out = self.w3(out)103        if self.dropout is not None:104            out = self.dropout(out)105        return out106 107 108class Attention(nn.Module):109    def __init__(self, config: TinyConfig):110        super().__init__()111        self.n_heads = config.num_attention_heads112        self.hidden_size = config.hidden_size113        self.head_dim = config.hidden_size // config.num_attention_heads114        115        assert self.n_heads * self.head_dim == self.hidden_size116        117        self.wq = nn.Linear(config.hidden_size, config.hidden_size, bias=False)118        self.wk = nn.Linear(config.hidden_size, config.hidden_size, bias=False)119        self.wv = nn.Linear(config.hidden_size, config.hidden_size, bias=False)120        self.wo = nn.Linear(config.hidden_size, config.hidden_size, bias=False)121        self.dropout_p = config.attention_dropout122 123    def forward(124        self,125        x: torch.Tensor,126        cos: torch.Tensor,127        sin: torch.Tensor,128        attention_mask: Optional[torch.Tensor] = None,129    ) -> torch.Tensor:130        B, T, C = x.shape131        q = self.wq(x)132        k = self.wk(x)133        v = self.wv(x)134        135        q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)136        k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)137        v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)138        139        q = apply_rotary_emb(q, cos, sin)140        k = apply_rotary_emb(k, cos, sin)141        142        dropout_p = self.dropout_p if self.training else 0.0143        144        if attention_mask is not None:145            if torch.all(attention_mask == 1):146                attn_mask = None147                is_causal = True148            else:149                causal_mask = torch.tril(torch.ones((T, T), dtype=torch.bool, device=x.device))150                padding_mask = attention_mask.to(torch.bool).unsqueeze(1).unsqueeze(2) # shape: (B, 1, 1, T)151                attn_mask = causal_mask.unsqueeze(0).unsqueeze(1) & padding_mask # shape: (B, 1, T, T)152                is_causal = False153        else:154            attn_mask = None155            is_causal = True156            157        out = F.scaled_dot_product_attention(158            q, k, v, 159            attn_mask=attn_mask, 160            dropout_p=dropout_p, 161            is_causal=is_causal162        )163        164        out = out.transpose(1, 2).contiguous().view(B, T, C)165        return self.wo(out)166 167 168class TransformerBlock(nn.Module):169    def __init__(self, config: TinyConfig):170        super().__init__()171        self.attention = Attention(config)172        self.feed_forward = FeedForward(config)173        self.attention_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)174        self.ffn_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)175 176    def forward(177        self,178        x: torch.Tensor,179        cos: torch.Tensor,180        sin: torch.Tensor,181        attention_mask: Optional[torch.Tensor] = None,182    ) -> torch.Tensor:183        x = x + self.attention(self.attention_norm(x), cos, sin, attention_mask)184        x = x + self.feed_forward(self.ffn_norm(x))185        return x186 187 188class TinyPreTrainedModel(PreTrainedModel):189    config_class = TinyConfig190    base_model_prefix = "model"191    supports_gradient_checkpointing = True192    _no_split_modules = ["TransformerBlock"]193 194    def _init_weights(self, module):195        if isinstance(module, nn.Linear):196            nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)197            if module.bias is not None:198                nn.init.zeros_(module.bias)199        elif isinstance(module, nn.Embedding):200            nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)201 202    def _set_gradient_checkpointing(self, module, value=False):203        if isinstance(module, (TinyModel, TinyForCausalLM)):204            module.gradient_checkpointing = value205 206 207class TinyModel(TinyPreTrainedModel):208    def __init__(self, config: TinyConfig):209        super().__init__(config)210        self.padding_idx = config.pad_token_id211        self.tok_embeddings = nn.Embedding(212            config.vocab_size, config.hidden_size, self.padding_idx213        )214        self.layers = nn.ModuleList(215            [TransformerBlock(config) for _ in range(config.num_hidden_layers)]216        )217        self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)218 219        cos, sin = precompute_freqs_cis(220            dim=config.hidden_size // config.num_attention_heads,221            end=config.max_position_embeddings * 2,222            theta=config.rope_theta,223        )224        self.register_buffer("cos", cos, persistent=False)225        self.register_buffer("sin", sin, persistent=False)226 227        self.gradient_checkpointing = False228        self.post_init()229 230    def get_input_embeddings(self):231        return self.tok_embeddings232 233    def set_input_embeddings(self, value):234        self.tok_embeddings = value235 236    def forward(237        self,238        input_ids: Optional[torch.LongTensor] = None,239        attention_mask: Optional[torch.Tensor] = None,240        position_ids: Optional[torch.LongTensor] = None,241        past_key_values: Optional[Tuple[torch.FloatTensor]] = None,242        inputs_embeds: Optional[torch.FloatTensor] = None,243        use_cache: Optional[bool] = None,244        output_attentions: Optional[bool] = None,245        output_hidden_states: Optional[bool] = None,246        return_dict: Optional[bool] = None,247    ) -> Union[Tuple, BaseModelOutputWithPast]:248        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions249        output_hidden_states = (250            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states251        )252        return_dict = return_dict if return_dict is not None else self.config.use_return_dict253 254        if input_ids is not None and inputs_embeds is not None:255            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")256        elif input_ids is not None:257            hidden_states = self.tok_embeddings(input_ids)258        elif inputs_embeds is not None:259            hidden_states = inputs_embeds260        else:261            raise ValueError("You must specify either input_ids or inputs_embeds")262 263        T = hidden_states.shape[1]264        cos = self.cos[:T]265        sin = self.sin[:T]266 267        all_hidden_states = () if output_hidden_states else None268        269        for layer in self.layers:270            if output_hidden_states:271                all_hidden_states += (hidden_states,)272 273            if self.gradient_checkpointing and self.training:274                hidden_states = self._gradient_checkpointing_func(275                    layer.__call__,276                    hidden_states,277                    cos,278                    sin,279                    attention_mask,280                )281            else:282                hidden_states = layer(283                    hidden_states,284                    cos,285                    sin,286                    attention_mask,287                )288 289        hidden_states = self.norm(hidden_states)290 291        if output_hidden_states:292            all_hidden_states += (hidden_states,)293 294        if not return_dict:295            return (hidden_states,)296 297        return BaseModelOutputWithPast(298            last_hidden_state=hidden_states,299            hidden_states=all_hidden_states,300            past_key_values=None,301        )302 303 304class TinyForCausalLM(TinyPreTrainedModel):305    _tied_weights_keys = ["tok_embeddings.weight"]306 307    def __init__(self, config: TinyConfig):308        super().__init__(config)309        self.padding_idx = config.pad_token_id310        self.tok_embeddings = nn.Embedding(311            config.vocab_size, config.hidden_size, self.padding_idx312        )313        self.layers = nn.ModuleList(314            [TransformerBlock(config) for _ in range(config.num_hidden_layers)]315        )316        self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)317        318        self.output = nn.Linear(config.hidden_size, config.vocab_size, bias=False)319 320        cos, sin = precompute_freqs_cis(321            dim=config.hidden_size // config.num_attention_heads,322            end=config.max_position_embeddings * 2,323            theta=config.rope_theta,324        )325        self.register_buffer("cos", cos, persistent=False)326        self.register_buffer("sin", sin, persistent=False)327 328        if config.tie_word_embeddings:329            self.output.weight = self.tok_embeddings.weight330 331        self.gradient_checkpointing = False332        self.post_init()333 334    def get_input_embeddings(self):335        return self.tok_embeddings336 337    def set_input_embeddings(self, value):338        self.tok_embeddings = value339        if self.config.tie_word_embeddings:340            self.output.weight = value.weight341 342    def get_output_embeddings(self):343        return self.output344 345    def set_output_embeddings(self, new_embeddings):346        self.output = new_embeddings347 348    def tie_weights(self):349        if self.config.tie_word_embeddings:350            self.tok_embeddings.weight = self.output.weight351 352    def forward(353        self,354        input_ids: Optional[torch.LongTensor] = None,355        attention_mask: Optional[torch.Tensor] = None,356        position_ids: Optional[torch.LongTensor] = None,357        past_key_values: Optional[Tuple[torch.FloatTensor]] = None,358        inputs_embeds: Optional[torch.FloatTensor] = None,359        labels: Optional[torch.LongTensor] = None,360        use_cache: Optional[bool] = None,361        output_attentions: Optional[bool] = None,362        output_hidden_states: Optional[bool] = None,363        return_dict: Optional[bool] = None,364        **kwargs,365    ) -> Union[Tuple, CausalLMOutputWithPast]:366        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions367        output_hidden_states = (368            output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states369        )370        return_dict = return_dict if return_dict is not None else self.config.use_return_dict371 372        if input_ids is not None and inputs_embeds is not None:373            raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")374        elif input_ids is not None:375            hidden_states = self.tok_embeddings(input_ids)376        elif inputs_embeds is not None:377            hidden_states = inputs_embeds378        else:379            raise ValueError("You must specify either input_ids or inputs_embeds")380 381        T = hidden_states.shape[1]382        cos = self.cos[:T]383        sin = self.sin[:T]384 385        all_hidden_states = () if output_hidden_states else None386        387        for layer in self.layers:388            if output_hidden_states:389                all_hidden_states += (hidden_states,)390 391            if self.gradient_checkpointing and self.training:392                hidden_states = self._gradient_checkpointing_func(393                    layer.__call__,394                    hidden_states,395                    cos,396                    sin,397                    attention_mask,398                )399            else:400                hidden_states = layer(401                    hidden_states,402                    cos,403                    sin,404                    attention_mask,405                )406 407        hidden_states = self.norm(hidden_states)408        logits = self.output(hidden_states)409 410        loss = None411        if labels is not None:412            shift_logits = logits[..., :-1, :].contiguous()413            shift_labels = labels[..., 1:].contiguous()414            loss_fct = nn.CrossEntropyLoss()415            loss = loss_fct(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1))416 417        if not return_dict:418            output = (logits,)419            if output_hidden_states:420                output = output + (all_hidden_states,)421            return (loss,) + output if loss is not None else output422 423        return CausalLMOutputWithPast(424            loss=loss,425            logits=logits,426            past_key_values=None,427            hidden_states=all_hidden_states,428            attentions=None,429        )430 431    def prepare_inputs_for_generation(432        self,433        input_ids: torch.LongTensor,434        past_key_values: Optional[Tuple[torch.FloatTensor]] = None,435        attention_mask: Optional[torch.Tensor] = None,436        inputs_embeds: Optional[torch.FloatTensor] = None,437        **kwargs,438    ) -> dict:439        if inputs_embeds is not None and past_key_values is None:440            model_inputs = {"inputs_embeds": inputs_embeds}441        else:442            model_inputs = {"input_ids": input_ids}443 444        model_inputs["attention_mask"] = attention_mask445        model_inputs["past_key_values"] = past_key_values446        return model_inputs