--- license: apache-2.0 tags: - jax - flax - research - orthogonal-embeddings - walsh-hadamard - gram-schmidt - efficient-transformers datasets: - roneneldan/TinyStories - HuggingFaceFW/fineweb-edu - HuggingFaceTB/smollm-corpus - HuggingFaceTB/dclm-edu --- # 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 ```python 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.