Logotipo do Mini GPT
Java 21 · Maven 3.9 · zero bibliotecas de ML

Um GPT pequeno o bastante
para você ler inteiro

Um modelo de linguagem em nível de caractere, do zero em Java puro. Sem PyTorch, sem TensorFlow, sem DJL: toda a álgebra linear é double[][] que cabe na tela. Treze arquivos, cerca de 3.200 linhas — e três níveis, do bigrama que só conta pares até um Transformer com atenção causal.

Um treino, do começo ao fim
Mini GPT — train --model transformer
$ java -jar target/mini-gpt-java.jar train --model transformer --steps 2000
== Treino: modelo 'transformer' ==
Corpus: data/corpus.txt (1043271 caracteres, vocabulario V=96)
Contexto=64 | Treino: 938943 tokens | Validacao: 104328 tokens
passo loss_treino loss_val tempo
1 4.5762 4.5701 0.71s
100 2.8410 2.8395 75.94s
300 2.2137 2.2088 226.15s
500 2.0431 2.0396 376.02s
--- amostra (100 caracteres) ---
o mento de sua parte de comprar a estava da porta, e o menta de casa de dia
1000 1.8502 1.8571 751.38s
--- amostra (100 caracteres) ---
a noite chegava devagar, e o menino nao sabia mais o que dizia na porta
1500 1.7410 1.7566 1126.77s
--- amostra (100 caracteres) ---
o velho abriu a loja com cuidado e olhou para a rua, como fazia todos os dias
2000 1.6588 1.6902 1502.44s
--- amostra (100 caracteres) ---
as criancas corriam pela rua estreita e chamavam umas as outras pelos nomes
Vocabulario salvo em target/transformer.vocab

Repare no que acontece entre o passo 500 e o 2000: o que o modelo ganha é estatística. Primeiro sílabas, depois palavras, depois concordância. Nenhuma regra de português foi escrita em lugar nenhum: só a conta de prever o próximo caractere, repetida duas mil vezes. A sessão acima ilustra o formato que Trainer imprime e segue a meta do projeto (perda de validação abaixo de 1,7); os números da sua máquina dependem do seu corpus.

JDK 21+ e Maven 3.9+Uma única dependência: JUnit 6Roda em CPU comum, sem GPUGradientes conferidos por diferenças finitas
Nível 1 — Bigrama: contagem pura, sem gradienteNível 2 — MLP: backpropagation explícita, sem autodiffNível 3 — Transformer: autodiff e atenção
A ideia central

Um modelo de linguagem faz uma coisa só

Ele responde a uma pergunta, repetidamente: dado o texto até aqui, qual é o próximo caractere? Contagem, gradiente e atenção são três maneiras cada vez melhores de responder a essa única pergunta.

p(x₁, x₂, …, x_T)  =  ∏  p(x_t | x₁ … x_{t−1})
x_t
o caractere na posição t — aqui, uma letra, um espaço ou uma quebra de linha, não uma palavra
p(x_t | …)
a probabilidade daquele caractere dado tudo o que veio antes
o produto sobre todas as posições: a probabilidade do texto inteiro é o produto das probabilidades de cada passo

Essa igualdade é só a regra do produto da probabilidade, sem nenhuma hipótese escondida. O que muda de um nível para o outro é quanto do passado cada modelo consegue de fato usar: o bigrama usa um caractere, o MLP usa oito, e o Transformer usa sessenta e quatro, decidindo sozinho a quais deles prestar atenção.

E como se mede se a resposta foi boa

Prever uma distribuição é apostar. A métrica precisa premiar quem deu probabilidade alta ao caractere que de fato veio, e punir quem deu quase nenhuma. A entropia cruzada faz exatamente isso, e é a mesma nos três níveis, o que permite compará-los.

L  =  −  (1/N)  Σ  ln p(x_t real)
L
a entropia cruzada: a perda que o treino minimiza
N
quantas posições foram avaliadas
ln p
o logaritmo natural da probabilidade que o modelo deu ao caractere que de fato apareceu

Leia assim: se o modelo dá probabilidade alta ao caractere certo, −ln p fica perto de zero. Se dá probabilidade quase nula, o valor explode. A perda é a surpresa média do modelo, medida em nats por caractere. Um modelo que chuta uniformemente entre V caracteres tem perda ln V: com V = 96, isso é 4,56. Todo nível deste projeto começa exatamente aí, e desce.

Três níveis

Uma escada de três degraus

Cada degrau resolve uma limitação concreta do degrau anterior, e cada um tem que bater a perda do anterior para justificar a própria complexidade. É por isso que o projeto começa por um modelo que qualquer pessoa entende em cinco minutos.

1Bigrama
model/BigramModel.java
Contar pares. Só isso.

O próximo caractere depende apenas do anterior. Estimamos a probabilidade contando quantas vezes cada par apareceu no corpus e dividindo. Não há gradiente, não há pesos, não há laço de treino.

  • Tokenização em nível de caractere
  • Estimativa de máxima verossimilhança
  • Suavização de Laplace: por que p = 0 é catastrófico
  • Entropia cruzada como métrica
  • Amostragem com temperatura e top-k
Contexto1 caractere
Treinoinstantâneo
Perda esperada~2–3 nats
Ler o nível 1
2MLP
model/MlpModel.java
Oito caracteres, e cada gradiente escrito explicitamente.

Uma rede rasa que olha uma janela fixa de oito caracteres. Embedding, concatenação, camada linear, tanh, camada de saída. Todos os gradientes escritos explicitamente, na ordem da regra da cadeia.

  • Embeddings: caracteres como vetores aprendidos
  • Backpropagation camada a camada
  • O gradiente conjunto de softmax + entropia cruzada
  • A derivada da tanh e a saturação
  • Scatter-add e mini-lotes
Contexto8 caracteres
Treinopoucos segundos
Perda esperadaabaixo do bigrama
Ler o nível 2
3Transformer
model/TransformerModel.java
Atenção causal, e um autodiff que deriva sozinho.

Um GPT em miniatura: embeddings de token e de posição, self-attention causal multi-head, LayerNorm, conexões residuais, feed-forward e pesos amarrados. Aqui ninguém escreve gradiente: o grafo do Tensor faz isso.

  • Atenção: consultas, chaves e valores
  • A máscara causal e por que ela é indispensável
  • Embeddings de posição
  • LayerNorm e conexões residuais
  • Diferenciação automática em modo reverso
Contexto64 caracteres
Treino~25 min (2000 passos)
Perda alvo< 1,7
Ler o nível 3

As perdas abaixo pressupõem um corpus real de cerca de 1 MB em português. Com o placeholder curto que vem no repositório, os três níveis parecem melhores do que são: o modelo decora em vez de generalizar, e é isso que a perda de validação existe para denunciar.

O caminho completo

Do arquivo de texto ao texto gerado

São seis etapas, e nenhuma delas é mágica. Cada caixa abaixo é um arquivo do repositório que você pode abrir hoje.

Do arquivo de texto ao texto gerado, em seis etapascorpus.txtidstokenizadorjanelas(x, y)logitsmodelopsoftmaxtextoamostrageme o texto gerado volta para o contexto: é isso que significa autoregressivo

Corpus

data/corpus.txt

Um arquivo .txt em UTF-8. Cerca de 1 MB de português em domínio público é o alvo. Quanto mais consistente o texto, mais legível a saída.

Tokenizador

data/CharTokenizer.java

Cada caractere distinto vira um número. Os caracteres são ordenados por code point, então o mesmo corpus gera sempre os mesmos ids. É uma bijeção, com decode(encode(t)) == t.

Janelas e lotes

data/Dataset.java

O corpus linear vira pares (x, y): y é x deslocado de uma posição. Os primeiros 90% são treino, o resto é validação, cortados por posição e nunca embaralhados.

Modelo

model/*.java

O nível que você escolheu transforma o contexto em logits: um número por caractere do vocabulário, ainda sem normalizar.

Treino

train/Trainer.java

Sorteia um mini-lote, roda forward e backward, e pede um passo ao AdamOptimizer. A cada 500 passos gera uma amostra, para você ver o texto evoluir.

Amostragem

generate/Sampler.java

Softmax nos logits, temperatura, top-k, sorteio. O caractere sorteado entra no contexto e o laço recomeça: é isso que quer dizer autoregressivo.

Veja acontecer

As contas, rodando no seu navegador

As quatro demonstrações abaixo são traduções fiéis do código Java do repositório, reescritas em TypeScript para rodarem aqui. Mexa nos controles: o objetivo é você sentir o que cada parâmetro faz antes de ler a fórmula.

1. Tokenizador em nível de caractereCharTokenizer.java

Escreva qualquer coisa. Cada caractere distinto do corpus recebeu um id na ordem do code point Unicode, e é esse número, não a letra, que entra na conta. Repare que o espaço e a quebra de linha também são caracteres, e que ã é um id como qualquer outro.

O121m28e21n29i25n29o301o30l27h24a17v37a171o301r33i25o30.3
vocabulário do corpus: V = 48caracteres digitados: 22fora do vocabulário: 0
É isto que stoi e itos fazem no Java. Nível de caractere é a escolha didática do projeto: o vocabulário fica pequeno (dezenas de ids em vez de dezenas de milhares) e o modelo tem que aprender a soletrar, o que torna o progresso visível a olho nu.
2. O nível 1 inteiro, treinado agoraBigramModel.java

Este botão faz o que fit() faz: uma passada pelo corpus contando pares. Depois gera texto sorteando da distribuição contada. É o modelo inteiro: no nível 1 não há mais nada. O seletor de corpus troca a linguagem que ele conta, e é aí que mora a lição.

Corpus
Massa mínima dada a todo par, inclusive aos que nunca ocorreram. Suba até 3 e veja o texto virar ruído: a suavização passou a pesar mais que as contagens.
Abaixo de 1 esfria e repete; acima de 1 esquenta e delira.
Mantém só os k mais prováveis e renormaliza. 0 desliga.
O menino aitum,cheínar.é a s beiravar, a chUDFD nojhentos oranis eEô Opo eOHarnUO SNcortobde pão s, uto ti orsacantguínva anãodo o m nta ado qumadavapenola te dle ca niuisiF co. nto vo enmeris, arararto ca pas, pondes doso fiui. poia o a penoba ntrtHtra A.lascídz caris diseidomum o re sva corcanãmonia cha,pázQFàzP:deia com qui
semente 20240824
vocabulário: V = 48perda no corpus: 2.213 natschute uniforme (ln V): 3.871pares nunca vistos: 1958 de 2304
Em português o texto tem a textura certa (a proporção de vogais, os acentos em lugares plausíveis, o tamanho das palavras) e quase nenhuma palavra real. Aumentar o corpus não resolve: com 75 vezes mais texto a perda cai cerca de 0,26 nats e as palavras continuam inventadas, porque o bigrama esquece tudo menos um caractere. Essa distância entre parecer e ser é o que o nível 2 vai atacar.

Troque para a linguagem de ordem 1 e o mesmo modelo passa a acertar tudo. São doze palavras em que cada letra interna determina sozinha a sua sucessora, ou seja, uma linguagem em que a hipótese do bigrama é verdadeira em vez de aproximada. O contador de palavras inválidas mostra isso acontecendo. Baixe o α até 0,01 e ele chega a zero; suba para 1 e as quimeras voltam (mãsal, nóAzul), porque a suavização dá massa às transições que a linguagem proíbe. É a mesma α que existe para impedir que um par nunca visto mande a perda para o infinito: aqui ela protege a métrica e estraga o texto, e dá para ver os dois efeitos no mesmo controle.

Repare no que ela não acerta: a ordem das palavras. Depois de um espaço, a linha da tabela é a mesma qualquer que tenha sido a palavra anterior, então o espaço apaga todo o contexto. Nenhum ajuste de α, τ ou top-k recupera isso, e é por isso que doze palavras é perto do teto: cada letra interna gasta uma sucessora exclusiva do alfabeto.

3. Temperatura e top-kSampler.java

A distribuição abaixo é real: são as probabilidades do próximo caractere depois de um a, contadas no mesmo corpus. Os dois controles não mudam o modelo; mudam apenas como se sorteia dele.

caractere anterior:
Leve para 0,2 e veja uma barra engolir as outras; leve para 2 e veja todas se igualarem.
As barras cortadas vão a zero e a massa delas é redistribuída entre as que ficaram.
33.2%
s9.6%
v8.3%
n7.8%
r7.2%
d5.6%
,3.7%
l2.9%
i2.7%
m2.7%
.2.2%
b1.9%

Os 12 caracteres mais prováveis depois de a, em ordem fixa. A barra mostra a probabilidade depois da temperatura e do top-k, normalizada pela maior.

Temperatura baixa concentra a massa no mais provável (texto conservador e repetitivo); temperatura alta achata a distribuição (texto criativo e ruidoso). O top-k corta a cauda: opções individualmente improváveis que, somadas, ainda roubam sorteios. No limite τ → 0 a amostragem vira argmax e o texto passa a se repetir para sempre.
4. A máscara causalTransformerModel.java

Cada linha é uma posição da sequência; cada coluna é uma posição que ela pode olhar. Arraste o controle para escolher a posição que está sendo prevista e veja o que ela enxerga.

omeninoo×××××××××××××m×××××e××××n×××i××n×o
pode olhar bloqueado (−∞) a própria posição

O que a posição 5 enxerga

Ela pode olhar 6 posições: o ␣ m e n i. Tudo o que vem depois está a −∞ e sai do softmax com peso exatamente zero.

Repare na posição 0: ela só pode olhar para si mesma. É o caso em que a atenção não tem nada a somar, e é por isso que o começo de qualquer geração é o pedaço mais frágil do texto.

E repare que o triângulo cheio é exatamente o desenho do logotipo deste site: é a figura que o projeto inteiro existe para explicar.

As células apagadas viram −∞ antes do softmax, o que as leva a peso exatamente zero depois dele. Sem essa máscara, o modelo veria o caractere que deveria prever: a perda de treino despencaria, o modelo não aprenderia nada útil e nada no programa reclamaria.
Como saber que está certo

Um gradiente errado treina em silêncio

Um sinal de menos trocado numa derivada não quebra nada: o programa roda, a perda cai um pouco e o modelo fica ruim sem dizer por quê. É o tipo de defeito que consome semanas. Por isso todo gradiente deste projeto é conferido contra a definição de derivada.

∂f/∂θ  ≈  ( f(θ + ε) − f(θ − ε) ) / (2ε)
f
a perda, como função de um único parâmetro
θ
o parâmetro sendo conferido — um número dentro de uma matriz de pesos
ε
um deslocamento minúsculo, da ordem de 1e-5

A diferença central é lenta demais para treinar (uma perturbação por parâmetro, duas avaliações cada), mas é independente do backward que se quer testar. Se o gradiente analítico e o numérico batem até a quinta casa, o backward está certo. Se não batem, o teste falha e diz qual tensor.

MlpGradCheckTest

Confere todos os gradientes explícitos do nível 2 (dC, dW1, db1, dW2, db2), com erro abaixo de 1e-5.

TransformerGradCheckTest

Confere o grafo inteiro do nível 3: atenção, LayerNorm, residuais, feed-forward e a cabeça de pesos amarrados, todos pelo mesmo critério.

TensorGradCheckTest

Confere mais de doze operações do autodiff isoladamente: matmul, softmaxRows, layerNorm, causalMask, crossEntropyRows e outras.

TrainingLearnsTest

Confere a afirmação que importa no fim: que Trainer mais AdamOptimizer de fato reduzem a perda do MLP e do Transformer.

$ mvn test
De onde vem a ideia

Um MINIX para modelos de linguagem

A aposta deste projeto já deu certo antes, em outra área da computação.

Em 1987, Andrew Tanenbaum publicou Operating Systems: Design and Implementation com um sistema operacional inteiro junto — o MINIX — e o código-fonte completo impresso no apêndice do livro. A aposta era que um aluno aprende mais lendo um sistema operacional inteiro do que lendo sobre um. Quatro anos depois, Linus Torvalds escreveu a primeira versão do Linux numa máquina rodando MINIX e a anunciou no grupo comp.os.minix.

Modelos de linguagem estão hoje onde os sistemas operacionais estavam então: todo mundo usa, quase ninguém leu um por dentro. Os que estão em produção têm bilhões de parâmetros e dependem de bibliotecas grandes demais para alguém ler de ponta a ponta. Este tem treze arquivos e cerca de 3.200 linhas, nenhuma dependência além do JUnit, e treina e gera texto. Dá para ler inteiro num fim de semana.

O que se herda do MINIX é método, não códigoNenhuma linha daqui vem do MINIX. O que este projeto copia dele são três decisões: ser completo, no sentido de treinar e gerar texto de verdade; ser pequeno de propósito; e manter a explicação no Javadoc da classe, ao lado da linha que ela explica.

A linhagem técnica

A matemática vem de dois artigos. Bengio et al. (2003), A Neural Probabilistic Language Model, é a arquitetura do nível 2 e o artigo que o Javadoc de MlpModel cita; foi ele que popularizou a ideia de embeddings aprendidos. Vaswani et al. (2017), Attention Is All You Need, é o nível 3. Os três degraus seguem a ordem em que essas ideias apareceram: primeiro a contagem, depois a rede rasa, depois a atenção.

As regras do jogo

O que este projeto se proíbe de fazer

Cada restrição abaixo existe para que o código continue legível por quem está aprendendo.

Nenhuma biblioteca de ML

Nada de DJL, ND4J, DL4J ou TensorFlow. Toda a álgebra linear é implementada aqui mesmo, com double[][], porque uma chamada de biblioteca é exatamente o ponto em que o aprendizado pararia.

Uma dependência, só nos testes

JUnit 6, e mais nada. O pom.xml cabe numa tela, e mvn test roda numa máquina recém-formatada.

A matemática mora no código

Cada classe traz, no Javadoc, a fórmula que ela implementa. O comentário e a linha de código ficam a centímetros um do outro, que é a única distância em que os dois se mantêm sincronizados.

CPU comum, sem GPU

O Tensor.matmul distribui as linhas do resultado entre os núcleos com um parallel stream. É o que mantém o Transformer dentro de meia hora sem recorrer a nada nativo.

Reprodutível por semente

Todo sorteio passa por um Random com semente explícita (--seed). Mesma semente, mesmo corpus, mesmo texto, o que torna um experimento comparável com o anterior.

Nada de peso mágico

Os modelos treináveis não persistem pesos: cada execução treina do zero. Custa tempo, e em troca não existe nenhum arquivo binário fazendo o trabalho que o código deveria mostrar.