Team Ai
Modelpublic

E6E831728/fixed-minimal-binary-code

sourceHugging Faceapache-2.0updated 8d agoView on Hugging Face
0likes275downloads
model_n_embed_16_binary_n_layer_32.py372 linesDownload Raw Back to root
1import math
2
3import torch
4import torch.nn as nn
5from torch.nn import functional as F
6
7from transformers import PreTrainedModel, PretrainedConfig
8from transformers.generation import GenerationMixin
9from transformers.modeling_outputs import CausalLMOutput, CausalLMOutputWithCrossAttentions
10
11class BVVConfig(PretrainedConfig):
12    model_type = "model_n_embed_16_binary_n_layer_32"
13
14    def __init__(
15        self,
16        vocab_size=65536,
17        n_embed=16,
18        d_model=1024,
19        n_head=32,
20        n_layer=32,
21        block_size=1024,
22        dropout=0.00,
23        layer_norm_eps=1e-5,
24        initializer_range=0.02,
25        pad_token_id=57344,
26        pad_id=57344,  # legacy alias
27        bos_token_id=None,
28        eos_token_id=None,
29        tie_word_embeddings=False,
30        use_cache=False,
31        **kwargs,
32    ):
33        if pad_token_id is None:
34            pad_token_id = 57344 if pad_id is None else pad_id
35
36        super().__init__(
37            pad_token_id=pad_token_id,
38            bos_token_id=bos_token_id,
39            eos_token_id=eos_token_id,
40            tie_word_embeddings=tie_word_embeddings,
41            use_cache=use_cache,
42            **kwargs,
43        )
44
45        if d_model % n_embed != 0:
46            raise ValueError(f"d_model ({d_model}) must be divisible by n_embed ({n_embed})")
47        if d_model % n_head != 0:
48            raise ValueError(f"d_model ({d_model}) must be divisible by n_head ({n_head})")
49        if (d_model // n_head) % 2 != 0:
50            raise ValueError("head_dim must be even for rotary embeddings")
51
52        self.vocab_size = vocab_size
53        self.block_size = block_size
54        self.max_position_embeddings = block_size
55
56        self.n_embed = n_embed
57        self.d_model = d_model
58        self.n_head = n_head
59        self.n_layer = n_layer
60
61        self.dropout = dropout
62        self.layer_norm_eps = layer_norm_eps
63        self.initializer_range = initializer_range
64
65        self.scale = d_model // n_embed
66
67        # backward compatibility
68        self.pad_id = pad_token_id
69
70
71def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
72    freqs = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
73    t = torch.arange(end, device=freqs.device)
74    freqs = torch.outer(t, freqs).float()
75    freqs_cis = torch.polar(torch.ones_like(freqs), freqs)  # complex64
76    return freqs_cis
77
78
79def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
80    ndim = x.ndim
81    assert 0 <= 1 < ndim
82    assert freqs_cis.shape == (x.shape[1], x.shape[-1])
83    shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
84    return freqs_cis.view(*shape)
85
86
87def apply_rotary_emb(
88    xq: torch.Tensor,
89    xk: torch.Tensor,
90    freqs_cis: torch.Tensor,
91):
92    xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
93    xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
94    freqs_cis = reshape_for_broadcast(freqs_cis, xq_)
95    xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3)
96    xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
97    return xq_out.type_as(xq), xk_out.type_as(xk)
98
99
100class MultiHeadSelfAttention(nn.Module):
101    def __init__(self, d_model, n_head, dropout=0.0):
102        super().__init__()
103        assert d_model % n_head == 0
104
105        self.d_model = d_model
106        self.n_head = n_head
107        self.head_dim = d_model // n_head
108
109        assert self.head_dim % 2 == 0, "head_dim must be even for rotary embeddings"
110
111        self.q_proj = nn.Linear(d_model, d_model, bias=False)
112        self.k_proj = nn.Linear(d_model, d_model, bias=False)
113        self.v_proj = nn.Linear(d_model, d_model, bias=False)
114        self.o_proj = nn.Linear(d_model, d_model, bias=False)
115
116        self.dropout = nn.Dropout(dropout)
117
118    def forward(self, x, freqs_cis, mask=None):
119        B, T, C = x.shape
120
121        q = self.q_proj(x).view(B, T, self.n_head, self.head_dim)
122        k = self.k_proj(x).view(B, T, self.n_head, self.head_dim)
123        v = self.v_proj(x).view(B, T, self.n_head, self.head_dim)
124
125        q, k = apply_rotary_emb(q, k, freqs_cis=freqs_cis)
126
127        q = q.transpose(1, 2)  # (B, n_head, T, head_dim)
128        k = k.transpose(1, 2)
129        v = v.transpose(1, 2)
130
131        attn_scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
132
133        if mask is not None:
134            attn_scores = attn_scores + mask
135
136        attn_probs = F.softmax(attn_scores.float(), dim=-1).type_as(q)
137        attn_probs = self.dropout(attn_probs)
138
139        out = torch.matmul(attn_probs, v)
140        out = out.transpose(1, 2).contiguous().view(B, T, C)
141
142        return self.o_proj(out)
143
144
145class TransformerMLP(nn.Module):
146    def __init__(self, d_model, dropout=0.0):
147        super().__init__()
148        self.net = nn.Sequential(
149            nn.Linear(d_model, 4 * d_model),
150            nn.GELU(),
151            nn.Linear(4 * d_model, d_model),
152            nn.Dropout(dropout),
153        )
154
155    def forward(self, x):
156        return self.net(x)
157
158
159class TransformerBlock(nn.Module):
160    def __init__(self, d_model, n_head, dropout=0.0, layer_norm_eps=1e-5):
161        super().__init__()
162        self.self_attn = MultiHeadSelfAttention(d_model, n_head, dropout=dropout)
163        self.mlp = TransformerMLP(d_model, dropout=dropout)
164        self.input_layernorm = nn.LayerNorm(d_model, eps=layer_norm_eps)
165        self.post_attention_layernorm = nn.LayerNorm(d_model, eps=layer_norm_eps)
166
167    def forward(self, x, freqs_cis, mask=None):
168        x = x + self.self_attn(self.input_layernorm(x), freqs_cis, mask)
169        x = x + self.mlp(self.post_attention_layernorm(x))
170        return x
171
172
173class BVVForCausalLM(PreTrainedModel, GenerationMixin):
174    config_class = BVVConfig
175    main_input_name = "input_ids"
176
177    def __init__(self, config: BVVConfig):
178        super().__init__(config)
179
180        self.token_embeddings = nn.Embedding(
181            config.vocab_size,
182            config.n_embed,
183            padding_idx=config.pad_token_id,
184        )
185        self.scale = config.scale
186
187        self.transformer_layers = nn.ModuleList([
188            TransformerBlock(
189                config.d_model,
190                n_head=config.n_head,
191                dropout=config.dropout,
192                layer_norm_eps=config.layer_norm_eps,
193            )
194            for _ in range(config.n_layer)
195        ])
196
197        self.final_layernorm = nn.LayerNorm(config.d_model, eps=config.layer_norm_eps)
198        self.lm_head = nn.Linear(config.d_model, config.vocab_size)
199
200        self.register_buffer(
201            "freqs_cis",
202            precompute_freqs_cis(
203                config.d_model // config.n_head,
204                config.block_size,
205            ),
206            persistent=False,
207        )
208
209        self.post_init()
210
211    def _init_weights(self, module):
212        if isinstance(module, nn.Linear):
213            nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
214            if module.bias is not None:
215                nn.init.zeros_(module.bias)
216
217        elif isinstance(module, nn.Embedding):
218            nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
219            if module.padding_idx is not None:
220                module.weight.data[module.padding_idx].zero_()
221
222    def get_input_embeddings(self):
223        return self.token_embeddings
224
225    def set_input_embeddings(self, value):
226        self.token_embeddings = value
227
228    def get_output_embeddings(self):
229        return self.lm_head
230
231    def set_output_embeddings(self, new_embeddings):
232        self.lm_head = new_embeddings
233
234    def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs):
235        if input_ids.shape[1] > self.config.block_size:
236            input_ids = input_ids[:, -self.config.block_size:]
237            if attention_mask is not None:
238                attention_mask = attention_mask[:, -self.config.block_size:]
239
240        return {
241            "input_ids": input_ids,
242            "attention_mask": attention_mask,
243        }
244
245    def forward(
246        self,
247        input_ids=None,
248        attention_mask=None,
249        labels=None,
250        targets=None,
251        return_dict=None,
252        output_logits=True,
253        **kwargs,
254    ):
255        if input_ids is None:
256            raise ValueError("input_ids must be provided")
257    
258        if labels is not None and targets is not None:
259            raise ValueError("Use either labels or targets, not both.")
260    
261        return_dict = return_dict if return_dict is not None else self.config.use_return_dict
262    
263        B, T = input_ids.shape
264        if T > self.config.block_size:
265            raise ValueError(f"Sequence length {T} exceeds block_size {self.config.block_size}")
266    
267        token_emb = self.token_embeddings(input_ids)
268        x = token_emb.repeat(1, 1, self.scale)
269    
270        freqs_cis = self.freqs_cis[:T]
271        if not torch.is_complex(freqs_cis):
272            freqs_cis = torch.view_as_complex(freqs_cis.contiguous())
273        freqs_cis = freqs_cis.to(x.device)
274    
275        mask = None
276        mask_value = torch.finfo(x.dtype).min
277    
278        if T > 1:
279            mask = torch.full((1, 1, T, T), mask_value, device=x.device, dtype=x.dtype)
280            mask = torch.triu(mask, diagonal=1)
281    
282        if attention_mask is not None:
283            if attention_mask.shape != (B, T):
284                raise ValueError(f"attention_mask must have shape {(B, T)}, got {tuple(attention_mask.shape)}")
285            pad_mask = torch.zeros((B, 1, 1, T), device=x.device, dtype=x.dtype)
286            pad_mask = pad_mask.masked_fill(attention_mask[:, None, None, :].eq(0), mask_value)
287            mask = pad_mask if mask is None else mask + pad_mask
288    
289        for layer in self.transformer_layers:
290            x = layer(x, freqs_cis, mask)
291    
292        x = self.final_layernorm(x)
293        logits = self.lm_head(x)
294    
295        loss = None
296    
297        if labels is not None:
298            shift_logits = logits[:, :-1, :].contiguous()
299            shift_labels = labels[:, 1:].contiguous()
300    
301            if attention_mask is not None:
302                shift_labels = shift_labels.masked_fill(attention_mask[:, 1:].eq(0), -100)
303    
304            if self.config.pad_token_id is not None:
305                shift_labels = shift_labels.masked_fill(shift_labels == self.config.pad_token_id, -100)
306    
307            loss = F.cross_entropy(
308                shift_logits.float().view(-1, shift_logits.size(-1)),
309                shift_labels.view(-1),
310                ignore_index=-100,
311            )
312    
313        elif targets is not None:
314            legacy_targets = targets.contiguous()
315    
316            if attention_mask is not None:
317                legacy_targets = legacy_targets.masked_fill(attention_mask.eq(0), -100)
318    
319            if self.config.pad_token_id is not None:
320                legacy_targets = legacy_targets.masked_fill(legacy_targets == self.config.pad_token_id, -100)
321    
322            loss = F.cross_entropy(
323                logits.float().view(-1, logits.size(-1)),
324                legacy_targets.view(-1),
325                ignore_index=-100,
326            )
327    
328        if not return_dict:
329            if output_logits:
330                output = (logits,)
331                return ((loss,) + output) if loss is not None else output
332            return (loss,) if loss is not None else tuple()
333        
334        if output_logits:
335            return CausalLMOutput(loss=loss, logits=logits)
336        return CausalLMOutput(loss=loss, logits=None)
337
338    def generate(self, input_ids, max_new_tokens, attention_mask=None, do_sample=False):
339        was_training = self.training
340        self.eval()
341    
342        if attention_mask is None:
343            attention_mask = torch.ones_like(input_ids, dtype=torch.long)
344    
345        with torch.no_grad():
346            for _ in range(max_new_tokens):
347                input_ids_cond = input_ids[:, -self.config.block_size:]
348                attention_mask_cond = attention_mask[:, -self.config.block_size:]
349    
350                outputs = self(
351                    input_ids=input_ids_cond,
352                    attention_mask=attention_mask_cond,
353                    return_dict=True
354                )
355                logits = outputs.logits[:, -1, :]
356    
357                if do_sample:
358                    probs = F.softmax(logits, dim=-1)
359                    next_token = torch.multinomial(probs, num_samples=1)
360                else:
361                    next_token = torch.argmax(logits, dim=-1, keepdim=True)
362    
363                input_ids = torch.cat([input_ids, next_token], dim=1)
364                attention_mask = torch.cat(
365                    [attention_mask, torch.ones_like(next_token, dtype=attention_mask.dtype)],
366                    dim=1
367                )
368    
369        if was_training:
370            self.train()
371    
372        return input_ids