File size: 6,640 Bytes
f4e6a29
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
"""Cortex-1 en PyTorch — implementación canónica equivalente al modelo NumPy.

Mapeo de pesos (NumPy -> PyTorch), verificado por test de equivalencia de
logits (max|Δ| < 1e-4 en float32):

    tok_emb.w                  -> model.embed_tokens.weight    (vocab, d)
    pos_emb.w                  -> model.embed_positions.weight (ctx, d)
    blocks.i.attn_norm.g       -> model.layers.i.input_layernorm.weight
    blocks.i.attn.wqkv         -> model.layers.i.self_attn.qkv.weight   (transpuesta)
    blocks.i.attn.wo           -> model.layers.i.self_attn.proj.weight  (transpuesta)
    blocks.i.mlp_norm.g        -> model.layers.i.post_attention_layernorm.weight
    blocks.i.mlp.w1            -> model.layers.i.mlp.gate_proj.weight   (transpuesta)
    blocks.i.mlp.w3            -> model.layers.i.mlp.up_proj.weight     (transpuesta)
    blocks.i.mlp.w2            -> model.layers.i.mlp.down_proj.weight   (transpuesta)
    final_norm.g               -> model.norm.weight
    (lm_head atado a tok_emb — tie_word_embeddings=True)

Arquitectura: GPT pre-norm, RMSNorm sin sesgos, atención causal multi-cabeza
con QKV fusionado y SwiGLU — sin dropout ni stochasticidad: eval == generate.
"""
from __future__ import annotations

import torch
import torch.nn as nn
import torch.nn.functional as F

from .configuration_cortex import CortexConfig

try:
    from transformers.modeling_utils import PreTrainedModel
    from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
    from transformers import GenerationMixin
except ImportError:  # transformers antiguo
    from transformers import PreTrainedModel, GenerationMixin
    from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast


class RMSNorm(nn.Module):
    def __init__(self, dim: int, eps: float):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))

    def forward(self, x):
        norm = x.float() * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps)
        return (norm * self.weight.float()).to(x.dtype)


class CortexAttention(nn.Module):
    """Atención causal con QKV fusionado (una GEMM), como el original NumPy."""

    def __init__(self, cfg: CortexConfig):
        super().__init__()
        self.n_heads = cfg.num_attention_heads
        self.d_head = cfg.hidden_size // cfg.num_attention_heads
        self.scale = self.d_head ** -0.5
        self.qkv = nn.Linear(cfg.hidden_size, 3 * cfg.hidden_size, bias=False)
        self.proj = nn.Linear(cfg.hidden_size, cfg.hidden_size, bias=False)

    def forward(self, x):
        B, T, d = x.shape
        qkv = self.qkv(x).view(B, T, 3, self.n_heads, self.d_head)
        q, k, v = qkv.unbind(dim=2)                       # (B, T, H, dh)
        q = q.transpose(1, 2)                             # (B, H, T, dh)
        k = k.transpose(1, 2)
        v = v.transpose(1, 2)
        att = (q @ k.transpose(-2, -1)) * self.scale      # (B, H, T, T)
        mask = torch.triu(torch.full((T, T), float("-inf"), device=x.device, dtype=att.dtype), 1)
        att = att + mask
        att = att.softmax(dim=-1)
        y = (att @ v).transpose(1, 2).reshape(B, T, d)    # reensambla cabezas
        return self.proj(y)


class CortexMLP(nn.Module):
    """SwiGLU: down(silu(gate(x)) * up(x)) — w1=gate, w3=up, w2=down."""

    def __init__(self, cfg: CortexConfig):
        super().__init__()
        self.gate_proj = nn.Linear(cfg.hidden_size, cfg.intermediate_size, bias=False)
        self.up_proj = nn.Linear(cfg.hidden_size, cfg.intermediate_size, bias=False)
        self.down_proj = nn.Linear(cfg.intermediate_size, cfg.hidden_size, bias=False)

    def forward(self, x):
        return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))


class CortexBlock(nn.Module):
    def __init__(self, cfg: CortexConfig):
        super().__init__()
        self.input_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
        self.self_attn = CortexAttention(cfg)
        self.post_attention_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
        self.mlp = CortexMLP(cfg)

    def forward(self, x):
        x = x + self.self_attn(self.input_layernorm(x))
        x = x + self.mlp(self.post_attention_layernorm(x))
        return x


class CortexModel(nn.Module):
    def __init__(self, cfg: CortexConfig):
        super().__init__()
        self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
        self.embed_positions = nn.Embedding(cfg.max_position_embeddings, cfg.hidden_size)
        self.layers = nn.ModuleList(CortexBlock(cfg) for _ in range(cfg.num_hidden_layers))
        self.norm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)

    def forward(self, input_ids):
        T = input_ids.shape[1]
        h = self.embed_tokens(input_ids) + self.embed_positions(torch.arange(T, device=input_ids.device))
        for layer in self.layers:
            h = layer(h)
        return self.norm(h)


class CortexForCausalLM(PreTrainedModel, GenerationMixin):
    config_class = CortexConfig
    _tied_weights_keys = ["lm_head.weight"]
    _dynamic_tied_weights_keys = ["lm_head.weight"]

    def __init__(self, cfg: CortexConfig):
        super().__init__(cfg)
        self.model = CortexModel(cfg)
        self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
        if cfg.tie_word_embeddings:
            self.lm_head.weight = self.model.embed_tokens.weight

    def forward(self, input_ids, attention_mask=None, labels=None,
                output_hidden_states=False, use_cache=False, **kwargs):
        h = self.model(input_ids)
        logits = self.lm_head(h)
        loss = None
        if labels is not None:
            shift_logits = logits[..., :-1, :].contiguous()
            shift_labels = labels[..., 1:].contiguous()
            loss = F.cross_entropy(
                shift_logits.view(-1, shift_logits.size(-1)),
                shift_labels.view(-1),
                ignore_index=-100,
            )
        return CausalLMOutputWithPast(
            loss=loss,
            logits=logits,
            hidden_states=(h,) if output_hidden_states else None,
        )

    @staticmethod
    def _init_weights(module):
        if isinstance(module, nn.Linear):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)
            if module.bias is not None:
                nn.init.zeros_(module.bias)
        elif isinstance(module, nn.Embedding):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)

    def prepare_inputs_for_generation(self, input_ids, **kwargs):
        return {"input_ids": input_ids}