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