File size: 5,193 Bytes
3b6d28f
821bd63
 
3b6d28f
 
 
 
 
821bd63
7c606de
821bd63
 
 
7c606de
 
821bd63
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7c606de
 
 
 
821bd63
 
 
 
 
bc6d608
821bd63
 
 
 
7c606de
821bd63
 
3b6d28f
 
 
7c606de
bc6d608
 
 
 
3b6d28f
821bd63
 
 
 
 
 
 
 
 
 
 
 
7c606de
821bd63
7c606de
 
 
 
 
 
 
 
 
 
 
 
 
bc6d608
7c606de
 
 
821bd63
 
 
 
7c606de
 
821bd63
7c606de
 
821bd63
 
 
7c606de
 
821bd63
7c606de
821bd63
7c606de
821bd63
7c606de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
# 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