HuggingFaceFW/fineweb-edu
Viewer • Updated • 3.5B • 407k • 1.31k
Este repositório documenta a implementação, teoria e checkpoints de uma arquitetura Transformer compacta (812.800 parâmetros) projetada para mitigar dois gargalos teóricos fundamentais em modelos de linguagem:
[Token Entrada x_0] (128d)
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Subcamada l (Atenção ou FFN): │
│ 1. Extração não-linear em subespaço local de baixo rank (64d) │
│ 2. Padding com 64 zeros (64 -> 128) │
│ 3. Rotação pseudoaleatória fixa Π_l (Quebra simetria diádica) │
│ 4. Difusão isométrica máxima: Multiplicação por Hadamard (128) │
│ 5. Soma linear no barramento residual: x = x + Δx_l │
└─────────────────────────────────────────────────────────────────┘
│ (12 subcamadas com subespaços mutualmente incoerentes)
▼
[Fluxo Residual Acumulado de Posto Completo 128d]
│
▼
┌─────────────────────────────────────────────────────────────────┐
│ Cabeça de Decodificação: │
│ Filtro Dinâmico de Gram-Schmidt com Gate Dependente de h: │
│ g(h) = sigmoid(W_gate @ h_final + b) │
│ h_orth = h_final - g(h) ⊙ proj_{x_0}(h_final) │
└─────────────────────────────────────────────────────────────────┘
│
▼
[Tied Softmax Head: W_embed^T @ h_orth]
--, hifens de lista - , quebras de linha \n), o gate desce para $\approx 0.18 - 0.25$, preservando $x_0$.Os pesos estão serializados no formato Flax (.msgpack) no diretório checkpoints/:
| Checkpoint | Tokens Acumulados | Corpus de Treino | Função de Perda | Observações |
|---|---|---|---|---|
supra_mini_tied_gs_params.msgpack |
10M | TinyStories | Hard Cross-Entropy | Loss: 2.7381. Throughput: 246k tok/s. |
supra_mini_20m_general_mix.msgpack |
30M | Mix Web (50% FineWeb, 30% Cosmo, 10% DCLM, 10% Tiny) | Hard Cross-Entropy | Loss: 4.7582. Domínio de sintaxe e markdown. |
supra_mini_70m_general_mix.msgpack |
80M | Mix Web Geral | Hard CE (Anneal $5e-4 \to 5e-5$) | Loss estabilizado; limpeza de atratores espúrios. |
supra_mini_75m_self_distill.msgpack |
85M | Mix Web Geral | $0.65 \text{ OneHot} + 0.35 \text{ Top-5 Soft}$ | Queda abrupta de entropia ($-0.85$ nats). |
supra_mini_90m_self_distill.msgpack |
100M | Mix Web Geral | $0.65 \text{ OneHot} + 0.35 \text{ Top-5 Soft}$ | Ponto fixo assintótico comprovado ($\Delta H \approx -0.04$ nats). |
from pathlib import Path
import jax
import jax.numpy as jnp
import flax
from transformers import AutoTokenizer
from modeling_supra_gs import SupraMiniWithGate, generate_hadamard_matrix, generate_sublayer_permutations
# 1. Tokenizer e Matrizes Estruturadas
tokenizer = AutoTokenizer.from_pretrained("SupraLabs/Supra-Mini-v6-1M")
H_128 = generate_hadamard_matrix(128)
perms = generate_sublayer_permutations(num_layers=6, dim=128, seed=2026)
# 2. Inicialização do Modelo
model = SupraMiniWithGate(
vocab_size=tokenizer.vocab_size,
d_model=128,
num_layers=6,
num_heads=4,
d_head=16,
max_len=256
)
# 3. Carregar Pesos do Checkpoint (Exemplo: 90M)
dummy_input = jnp.zeros((1, 256), dtype=jnp.int32)
template_params = model.init(jax.random.PRNGKey(0), dummy_input, perms, H_128)["params"]
ckpt_file = "checkpoints/supra_mini_90m_self_distill.msgpack"
with open(ckpt_file, "rb") as f:
params = flax.serialization.from_bytes(template_params, f.read())
# 4. Inferência
prompt = "Photosynthesis is a biological process where plants use sunlight to"
tokens = tokenizer.encode(prompt)
tok_tensor = jnp.zeros((1, 256), dtype=jnp.int32).at[0, :len(tokens)].set(jnp.array(tokens))
logits, gate_h, proj_x_in, _ = model.apply({"params": params}, tok_tensor, perms, H_128)
print(f"Logits shape: {logits.shape} | Gate médio: {jnp.mean(gate_h[0, :len(tokens)]):.4f}")