Supra-Mini com Difusão Isométrica de Hadamard e Anti-Resíduo de Gram-Schmidt Dinâmico

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:

  1. O custo de "apagar" o token de entrada no barramento residual (Residual Memory Burden).
  2. O desperdício paramétrico de projeções de posto completo em cada subcamada incremental.

1. Fundamentos da Arquitetura

[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]

Principais Inovações Teóricas:

  1. Atualizações de Baixo Rank com Difusão Isométrica de Hadamard:
    • Cada subcamada projeta apenas em $64$ dimensões ativas e preenche com $64$ zeros.
    • A matriz ortonormalizada de Sylvester-Hadamard ($H^T H = I$, $|H_{ij}| = 1/\sqrt{128}$) espalha a energia de forma equiprovável sobre todas as 128 coordenadas sem alterar a norma $\ell_2$.
    • As permutações ${\Pi_1, \dots, \Pi_{12}}$ garantem que os subespaços de atualização das 12 subcamadas tenham distância Grassmanniana máxima, permitindo cobrir todo o espaço $\mathbb{R}^{128}$ no acumulador.
  2. Projetor de Gram-Schmidt com Válvula de Escape Dinâmica:
    • Modela algebricamente a subtração da componente colinear a $x_0$, desonerando as camadas ocultas de aprenderem interferência destrutiva.
    • O gate $g(h) \in (0, 1)$ impede o colapso de repetição: quando a sintaxe exige repetição legítima de tokens (ex: --, hifens de lista - , quebras de linha \n), o gate desce para $\approx 0.18 - 0.25$, preservando $x_0$.

2. Checkpoints Disponíveis no Repositório

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).

3. Como Carregar e Usar em Python / JAX

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}")

4. Hardware e Throughput

  • Acelerador: Google Cloud TPU v5e-1 core.
  • Velocidade Média: $186.000$ a $246.000$ tokens/segundo.
  • Tempo Total de Computação: $< 7$ minutos para executar 100 milhões de tokens de treinamento e avaliação.
Downloads last month
-
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Datasets used to train jpllm/Research