# modeling_easyformer.py import math, torch, torch.nn as nn, torch.nn.functional as F from transformers import PreTrainedModel, GenerationMixin try: from .configuration_easyformer import EasyFormerConfig except ImportError: from configuration_easyformer import EasyFormerConfig # ----------------------------- СЛОИ ------------------------------- class RMSNorm(nn.Module): def __init__(self, d, eps=1e-5): super().__init__() self.w = nn.Parameter(torch.ones(d)) self.eps = eps def forward(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.w class EasyFormerAttention(nn.Module): def __init__(self, cfg): super().__init__() self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model, bias=False) self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False) self.drop = nn.Dropout(cfg.dropout) self.register_buffer("mask", torch.tril(torch.ones(cfg.ctx, cfg.ctx)).bool()) def forward(self, x): B, T, C = x.shape q, k, v = self.qkv(x).chunk(3, dim=-1) att = (q @ k.transpose(-2, -1)) / math.sqrt(C) att = att.masked_fill(~self.mask[:T, :T], float("-inf")) att = self.drop(F.softmax(att, dim=-1)) return self.proj(att @ v) class EasyFormerFFN(nn.Module): def __init__(self, cfg): super().__init__() self.fc1 = nn.Linear(cfg.d_model, 2 * cfg.d_model) self.fc2 = nn.Linear(2 * cfg.d_model, cfg.d_model) self.drop = nn.Dropout(cfg.dropout) def forward(self, x): return self.drop(self.fc2(F.relu(self.fc1(x)))) class EasyFormerBlock(nn.Module): def __init__(self, cfg): super().__init__() self.ln1 = RMSNorm(cfg.d_model) self.attn = EasyFormerAttention(cfg) self.ln2 = RMSNorm(cfg.d_model) self.ffn = EasyFormerFFN(cfg) def forward(self, x): x = x + self.attn(self.ln1(x)) x = x + self.ffn(self.ln2(x)) return x # ------------------------- HF-ОБЁРТКА ----------------------------- class EasyFormerPreTrainedModel(PreTrainedModel): config_class = EasyFormerConfig base_model_prefix = "easyformer" supports_gradient_checkpointing = False _no_split_modules = ["EasyFormerBlock"] class EasyFormerLMHeadModel(EasyFormerPreTrainedModel, GenerationMixin): config_class = EasyFormerConfig base_model_prefix = "easyformer" _tied_weights_keys = ["lm_head.weight"] all_tied_weights_keys = {"lm_head.weight": "tok_emb.weight"} _supports_cache_class = False _supports_flash_attn_2 = False _supports_sdpa = False main_input_name = "input_ids" def __init__(self, config): super().__init__(config) self.cfg = config self.tok_emb = nn.Embedding(config.vocab_size, config.d_model) self.pos_emb = nn.Embedding(config.ctx, config.d_model) self.drop = nn.Dropout(config.dropout) self.blocks = nn.ModuleList( [EasyFormerBlock(config) for _ in range(config.n_layer)] ) self.ln_f = RMSNorm(config.d_model) self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) self.lm_head.weight = self.tok_emb.weight self.post_init() # --- HF API --- def get_input_embeddings(self): return self.tok_emb def set_input_embeddings(self, value): self.tok_emb = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def tie_weights(self, recompute_mapping=False, **kwargs): self.lm_head.weight = self.tok_emb.weight # --- forward --- def forward(self, input_ids, attention_mask=None, labels=None, **kwargs): B, T = input_ids.shape pos = torch.arange(T, device=input_ids.device) x = self.drop(self.tok_emb(input_ids) + self.pos_emb(pos)) for b in self.blocks: x = b(x) logits = self.lm_head(self.ln_f(x)) loss = None if labels is not None: loss = F.cross_entropy( logits.view(-1, self.cfg.vocab_size), labels.view(-1), ignore_index=-100, ) return {"loss": loss, "logits": logits} if loss is not None else {"logits": logits} # --- generation --- def prepare_inputs_for_generation(self, input_ids, **kwargs): return {"input_ids": input_ids} @torch.no_grad() def generate(self, input_ids, max_new_tokens=40, temperature=0.6, top_k=20, do_sample=True, **kwargs): self.eval() for _ in range(max_new_tokens): idx_cond = input_ids[:, -self.cfg.ctx:] logits = self(idx_cond)["logits"][:, -1, :] / max(temperature, 1e-5) if top_k: v, _ = torch.topk(logits, top_k) logits[logits < v[:, [-1]]] = -float("inf") probs = F.softmax(logits, dim=-1) next_id = torch.multinomial(probs, 1) if do_sample else probs.argmax(-1, keepdim=True) input_ids = torch.cat([input_ids, next_id], dim=1) return input_ids