Krishna0812/Tiny_Stories
064
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