|
Download README.md from jpllm/Research: direct link, hf CLI and curl.
- Browser
- Download file 6.5 kB
-
https://huggingface.co/jpllm/Research/resolve/main/README.md
- Command line
-
hf download hf://jpllm/Research/README.md
-
curl -L -o README.md https://huggingface.co/jpllm/Research/resolve/main/README.md
6.5 kB
metadata
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:
- O custo de "apagar" o token de entrada no barramento residual (Residual Memory Burden).
- 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:
- 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.
- 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.