Theo-500M-Base / src /model.py
plvictor's picture
Upload src/model.py with huggingface_hub
a9b9ea6 verified
Raw
History Blame Contribute Delete
8.75 kB
"""
model.py — Transformer Decoder-Only (estilo GPT) implementado do zero em PyTorch.
Arquitetura:
Embedding de tokens + Positional Encoding absoluto aprendido
→ N x (Causal Self-Attention + FFN com residual + LayerNorm pré-norma)
→ Linear head projetando de d_model → vocab_size
Não usa nenhum peso externo. Todos os pesos são inicializados aleatoriamente
e aprendidos do zero durante o treino.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from config import BabelConfig
# ── Bloco de Atenção Causal ────────────────────────────────────────────────────
class CausalSelfAttention(nn.Module):
"""
Multi-head self-attention com máscara causal (cada token só enxerga
os tokens anteriores e a si mesmo — nunca os futuros).
"""
def __init__(self, config: BabelConfig):
super().__init__()
assert config.d_model % config.n_heads == 0, (
f"d_model ({config.d_model}) deve ser divisível por n_heads ({config.n_heads})"
)
self.n_heads = config.n_heads
self.d_head = config.d_model // config.n_heads
self.d_model = config.d_model
self.dropout = config.dropout
self.qkv_proj = nn.Linear(config.d_model, 3 * config.d_model, bias=False)
self.out_proj = nn.Linear(config.d_model, config.d_model, bias=False)
self.resid_dropout = nn.Dropout(config.dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, T, C = x.shape
qkv = self.qkv_proj(x)
q, k, v = qkv.split(self.d_model, dim=2)
q = q.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
k = k.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
v = v.view(B, T, self.n_heads, self.d_head).transpose(1, 2)
# Flash Attention via PyTorch 2.x — is_causal=True substitui a máscara manual
dropout_p = self.dropout if self.training else 0.0
out = F.scaled_dot_product_attention(q, k, v, dropout_p=dropout_p, is_causal=True)
out = out.transpose(1, 2).contiguous().view(B, T, C)
return self.resid_dropout(self.out_proj(out))
# ── Feed-Forward Network ───────────────────────────────────────────────────────
class FeedForward(nn.Module):
"""
FFN com dois lineares e GELU. d_ff = 4 * d_model por convenção GPT.
"""
def __init__(self, config: BabelConfig):
super().__init__()
self.net = nn.Sequential(
nn.Linear(config.d_model, config.d_ff, bias=False),
nn.GELU(),
nn.Linear(config.d_ff, config.d_model, bias=False),
nn.Dropout(config.dropout),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
# ── Bloco Transformer ──────────────────────────────────────────────────────────
class TransformerBlock(nn.Module):
"""
Bloco padrão: pré-norma → atenção → residual → pré-norma → FFN → residual.
Pré-norma (LayerNorm antes das sub-camadas) é mais estável em treinos
do zero do que a pós-norma do paper original.
"""
def __init__(self, config: BabelConfig):
super().__init__()
self.ln1 = nn.LayerNorm(config.d_model)
self.attn = CausalSelfAttention(config)
self.ln2 = nn.LayerNorm(config.d_model)
self.ffn = FeedForward(config)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln1(x))
x = x + self.ffn(self.ln2(x))
return x
# ── Modelo Completo ────────────────────────────────────────────────────────────
class BabelTransformer(nn.Module):
"""
Transformer decoder-only treinado do zero.
Uso:
config = get_config("babel_25M")
model = BabelTransformer(config)
logits, loss = model(input_ids, targets)
"""
def __init__(self, config: BabelConfig):
super().__init__()
self.config = config
self.token_emb = nn.Embedding(config.vocab_size, config.d_model)
self.pos_emb = nn.Embedding(config.context_length, config.d_model)
self.emb_dropout = nn.Dropout(config.dropout)
self.blocks = nn.ModuleList(
[TransformerBlock(config) for _ in range(config.n_layers)]
)
self.ln_final = nn.LayerNorm(config.d_model)
# Projeção final: d_model → vocab_size (sem bias, como GPT-2)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
# Weight tying: token_emb e lm_head compartilham os mesmos pesos.
# Reduz parâmetros e melhora qualidade (Press & Wolf, 2017).
self.lm_head.weight = self.token_emb.weight
self._init_weights()
def _init_weights(self):
"""Inicialização normal com std 0.02, igual ao GPT-2."""
for module in self.modules():
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.LayerNorm):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def forward(
self,
input_ids: torch.Tensor, # (B, T)
targets: torch.Tensor | None = None, # (B, T) — None em inferência
) -> tuple[torch.Tensor, torch.Tensor | None]:
B, T = input_ids.shape
assert T <= self.config.context_length, (
f"Sequência de tamanho {T} excede context_length={self.config.context_length}"
)
positions = torch.arange(T, device=input_ids.device).unsqueeze(0) # (1, T)
x = self.emb_dropout(self.token_emb(input_ids) + self.pos_emb(positions))
for block in self.blocks:
if self.config.gradient_checkpointing and self.training:
x = checkpoint(block, x, use_reentrant=False)
else:
x = block(x)
x = self.ln_final(x)
logits = self.lm_head(x) # (B, T, vocab_size)
loss = None
if targets is not None:
# Cross-entropy: achata (B*T, V) vs (B*T,)
loss = F.cross_entropy(
logits.view(-1, self.config.vocab_size),
targets.view(-1),
ignore_index=-1,
)
return logits, loss
@torch.no_grad()
def generate(
self,
input_ids: torch.Tensor, # (1, T) — um único prompt
max_new_tokens: int = 200,
temperature: float = 0.8,
top_p: float = 0.9,
repetition_penalty: float = 1.3,
) -> torch.Tensor:
"""
Gera texto auto-regressivamente com top-p (nucleus) sampling.
Retorna tensor (1, T + max_new_tokens).
"""
self.eval()
for _ in range(max_new_tokens):
# Trunca se ultrapassar context_length
ctx = input_ids[:, -self.config.context_length:]
logits, _ = self(ctx)
logits = logits[:, -1, :] / temperature # apenas o último token
# Penaliza tokens já gerados para reduzir repetição
if repetition_penalty != 1.0:
for token_id in input_ids[0].tolist():
if logits[0, token_id] > 0:
logits[0, token_id] /= repetition_penalty
else:
logits[0, token_id] *= repetition_penalty
# Top-p sampling
sorted_logits, sorted_idx = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
# Remove tokens acima do threshold top_p
sorted_logits[cumulative_probs - F.softmax(sorted_logits, dim=-1) > top_p] = float("-inf")
probs = F.softmax(sorted_logits, dim=-1)
next_token_sorted = torch.multinomial(probs, num_samples=1)
next_token = sorted_idx.gather(-1, next_token_sorted)
input_ids = torch.cat([input_ids, next_token], dim=1)
return input_ids
def count_parameters(self) -> int:
return sum(p.numel() for p in self.parameters() if p.requires_grad)