Research / README.md
jpllm's picture
Adiciona arquitetura Hadamard-GS, config.json, checkpoints (10M a 90M) e documentação técnica
ecf9a39 verified
|
Raw History Blame Contribute Delete
6.5 kB
---
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.