Etapa 0: entorno, configuración validada, modelo y entrenador

Base del proyecto ENLACE: un modelo de lenguaje propio entrenado desde cero,
en español, para asistencia general y familiar. El plan completo está en
docs/PLAN.md.

Esta etapa establece el andamiaje y lo verifica de punta a punta:

- Configuración por capas (hardware × model × train × data) validada con
  pydantic. Ningún hiperparámetro vive en el código y una config inválida
  falla al arrancar, no a las tres horas de entrenamiento.
- Perfiles de hardware que aíslan el salto de GPU: la RTX 2060 (Turing) no
  soporta bfloat16 ni FlashAttention-2, así que entrena en float16 con
  GradScaler y backend mem_efficient; el perfil de la 5090 ya está escrito.
  backends.py valida el perfil contra la GPU real antes de empezar.
- Transformer decoder-only estilo Llama: RMSNorm, SwiGLU, RoPE, GQA,
  embeddings atados, QK-norm y z-loss. Los dos últimos son lo que mantiene
  estable el entrenamiento en float16.
- Entrenador con schedule WSD, acumulación de gradiente, precisión mixta,
  checkpointing atómico y reanudación exacta.
- Cargadores de datos con estado serializable: bytes para el smoke test y
  shards uint16 para el corpus real.

48 tests, entre ellos el crítico: reanudar desde un checkpoint reproduce los
pesos de una corrida ininterrumpida, parámetro por parámetro.

Verificado en CPU: 300 pasos sobre texto en español, loss 3.07 -> 1.63.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-27 23:04:45 -03:00
commit 90ac2a6582
33 changed files with 2838 additions and 0 deletions
+18
View File
@@ -0,0 +1,18 @@
# Copiar a .env (que está en .gitignore) y completar.
# Las configs se commitean; los secretos, nunca.
# --- Servidor de entrenamiento (acceso por SSH) ---
ENLACE_REMOTE_HOST=usuario@servidor
ENLACE_REMOTE_DIR=/home/usuario/ENLACE
# --- Búsqueda web (elegir uno; SearxNG autoalojado es la opción sin terceros) ---
ENLACE_SEARXNG_URL=http://localhost:8888
# ENLACE_BRAVE_API_KEY=
# ENLACE_TAVILY_API_KEY=
# --- Home Assistant ---
# ENLACE_HA_URL=http://homeassistant.local:8123
# ENLACE_HA_TOKEN=
# --- Destilación de datos sintéticos (solo en la generación del dataset SFT) ---
# ANTHROPIC_API_KEY=
+16
View File
@@ -0,0 +1,16 @@
# Datos, pesos y bases del sistema: viven en el servidor, nunca en git.
data/
checkpoints/
snapshots/
runs/
# Secretos
.env
# Entorno
.venv/
__pycache__/
*.py[cod]
.pytest_cache/
.ruff_cache/
*.egg-info/
+93
View File
@@ -0,0 +1,93 @@
# ENLACE
Modelo de lenguaje propio, entrenado desde cero, para asistencia general y
familiar en español. No es fine-tuning de un modelo existente: tokenizer, datos
y pesos son propios.
El plan completo — etapas, presupuestos de cómputo, arquitectura de datos y
criterios de verificación — está en `docs/PLAN.md`.
## Estado
Etapa 0 (entorno y validación del stack) implementada y verificada:
- Capa de configuración validada, con composición por capas y perfiles de hardware.
- Arquitectura del modelo: RMSNorm, SwiGLU, RoPE, GQA, QK-norm, z-loss.
- Entrenador con WSD, precisión mixta, acumulación de gradiente y **reanudación exacta**.
- Cargadores de datos: bytes (smoke test) y shards `uint16` (corpus real).
- 48 tests, incluido el de reanudación exacta.
Pendiente: Etapa 1 en adelante (corpus, tokenizer, pretraining, post-training,
agente, bases del sistema, memoria, snapshots).
## Puesta en marcha
```bash
python3 -m venv .venv
.venv/bin/pip install -e ".[dev]"
scripts/prepare_smoke_data.sh # texto en español de dominio público
.venv/bin/python -m enlace.train.train configs/runs/smoke-cpu.yaml
scripts/test.sh
```
En el servidor con GPU, por SSH (ver `.env.example`):
```bash
scripts/remote.sh sync
scripts/remote.sh run configs/runs/smoke-2060.yaml # smoke test, ~10 min
scripts/remote.sh run configs/runs/pretrain-2060.yaml # pretraining, 1-2 días
scripts/remote.sh logs
```
Toda corrida larga arranca bajo `tmux`: cortar el SSH no mata el entrenamiento.
## Configuración
Ningún hiperparámetro vive en el código. Las configs se componen por capas:
```yaml
# configs/runs/pretrain-2060.yaml
include:
hardware: hardware/turing-2060.yaml # dtype, backend de atención, batch
model: model/tiny-50m.yaml # capas, dimensiones, contexto
train: train/pretrain.yaml # LR, schedule, intervalos
data: data/corpus.yaml # de dónde salen los tokens
```
Overrides puntuales desde la línea de comandos, para probar variantes sin editar
archivos:
```bash
.venv/bin/python -m enlace.train.train configs/runs/smoke-cpu.yaml train.seed=99
```
Una config inválida falla al arrancar con un mensaje concreto, no a las tres
horas de entrenamiento.
## Hardware
El perfil de hardware es la única capa que cambia al migrar de placa.
| Perfil | Placa | dtype | Atención | Notas |
|---|---|---|---|---|
| `turing-2060` | RTX 2060 12 GB (sm_75) | float16 + GradScaler | mem_efficient | Turing no soporta bfloat16 ni FlashAttention-2 |
| `blackwell-5090` | RTX 5090 32 GB (sm_120) | bfloat16 | flash | Requiere PyTorch ≥ 2.7 con CUDA 12.8 |
| `cpu` | — | float32 | math | Solo desarrollo y tests |
`enlace/train/backends.py` es el único archivo del entrenamiento que conoce
diferencias entre placas, y valida el perfil contra la GPU real antes de
empezar: pedir bfloat16 en Turing falla de inmediato en vez de degradarse en
silencio.
## Estructura
```
enlace/config/ Esquemas pydantic y composición de YAML
enlace/model/ Transformer decoder-only y selección de backend de atención
enlace/train/ Bucle, schedules, checkpointing, detección de hardware
enlace/data/ Cargadores con reanudación exacta
configs/ hardware × model × train × data, compuestos en configs/runs/
scripts/ Utilidades: remote.sh, test.sh, prepare_smoke_data.sh
tests/ Suite completa; test_resume.py es el crítico
```
+27
View File
@@ -0,0 +1,27 @@
# RTX 5090 32 GB — Blackwell, compute capability 12.0.
#
# Requiere PyTorch >= 2.7 compilado con CUDA 12.8: las versiones anteriores no
# conocen sm_120 y fallan al compilar kernels. Antes de usar este perfil hay que
# correr el smoke test de la Etapa 0 y verificar el entorno entero.
#
# Con bfloat16 desaparece el GradScaler: bf16 tiene el mismo rango exponente que
# fp32, así que no hay desbordes que escalar.
name: blackwell-5090
device: cuda
dtype: bfloat16
use_grad_scaler: false
attention_backend: flash
compile: true
matmul_precision: high
# 32 GB permiten un micro-lote 4x mayor con la misma acumulación reducida,
# manteniendo un lote efectivo comparable (32 * 4 = 128 secuencias) para que
# las curvas sean comparables entre placas.
micro_batch_size: 32
grad_accum_steps: 4
min_compute_capability: [12, 0]
# Aproximado, y solo para la métrica de MFU: bfloat16 denso (sin sparsity).
# Verificar con un benchmark real al migrar y ajustar este número.
peak_tflops: 209.5
+17
View File
@@ -0,0 +1,17 @@
# CPU — solo para desarrollo y tests. No sirve para entrenar nada real.
#
# Existe para que el código se pueda verificar en la máquina de trabajo, que no
# tiene GPU NVIDIA: los tests de forma, el smoke test a nivel de caracteres y
# la validación de configs corren acá antes de tocar el servidor.
name: cpu
device: cpu
dtype: float32
use_grad_scaler: false
attention_backend: math
compile: false
matmul_precision: high
micro_batch_size: 8
grad_accum_steps: 1
min_compute_capability: null
+28
View File
@@ -0,0 +1,28 @@
# RTX 2060 12 GB — Turing, compute capability 7.5.
#
# Dos límites de la arquitectura Turing determinan este perfil:
#
# 1. No soporta bfloat16. Hay que entrenar en float16, que tiene rango
# exponente reducido y desborda: por eso use_grad_scaler=true. Los picos
# de loss se mitigan además con qk_norm y z_loss en el modelo.
# 2. FlashAttention-2 exige Ampere (sm_80) o superior. El backend
# mem_efficient de PyTorch sí corre acá y usa memoria O(n) igual.
name: turing-2060
device: cuda
dtype: float16
use_grad_scaler: true
attention_backend: mem_efficient
compile: true
matmul_precision: high
# 12 GB con seq_len 2048 y un modelo de ~50M: micro-lote chico y acumulación
# para llegar a un lote efectivo razonable (8 * 16 = 128 secuencias).
micro_batch_size: 8
grad_accum_steps: 16
min_compute_capability: [7, 5]
# Aproximado, y solo para la métrica de MFU: tensor cores en float16 con
# acumulación en float32, que es lo que hace el entrenamiento con AMP.
# Sirve para comparar corridas entre sí, no como dato de hardware.
peak_tflops: 28.0
+23
View File
@@ -0,0 +1,23 @@
# Modelo diminuto a nivel de caracteres, solo para el smoke test de la Etapa 0.
#
# No produce nada útil: existe para validar en ~10 minutos que CUDA, el
# entrenador, el checkpointing, la reanudación y el flujo por SSH funcionan
# antes de comprometer días de cómputo.
name: char-smoke
vocab_size: 256 # se ajusta al vocabulario real del texto al cargar los datos
n_layer: 4
n_head: 4
n_kv_head: 2
d_model: 128
seq_len: 256
ffn_mult: 2.6667
ffn_multiple_of: 32
rope_theta: 10000.0
norm_eps: 1.0e-5
tie_embeddings: true
qk_norm: true
z_loss_weight: 1.0e-4
dropout: 0.0
init_std: 0.02
+30
View File
@@ -0,0 +1,30 @@
# ~50M parámetros. Primer modelo real, pensado para la 2060.
#
# Con vocab 32k y d_model 512, la tabla de embeddings sola son 16.7M
# parámetros: atarla con la capa de salida (tie_embeddings) ahorra un tercio
# del modelo. En modelos chicos esa decisión no es un detalle.
#
# seq_len 2048 y no 1024: el contexto recuperado de las bases del sistema se
# inyecta ahí, así que el presupuesto de contexto es un recurso de primer orden.
name: tiny-50m
vocab_size: 32768
n_layer: 12
n_head: 8
n_kv_head: 2 # GQA 4:1 — reduce la caché KV sin costo medible de calidad
d_model: 512
seq_len: 2048
ffn_mult: 2.6667 # 8/3, el habitual para SwiGLU (tres matrices en vez de dos)
ffn_multiple_of: 64
rope_theta: 10000.0
norm_eps: 1.0e-5
tie_embeddings: true
# qk_norm y z_loss no son opcionales entrenando en float16 en la 2060:
# son lo que evita que el loss explote a mitad de corrida.
qk_norm: true
z_loss_weight: 1.0e-4
dropout: 0.0 # con 3B tokens y 50M params no hay sobreajuste que combatir
init_std: 0.02
+31
View File
@@ -0,0 +1,31 @@
# Pretraining del modelo base (Etapa 2 del plan).
run_name: pretrain-tiny-50m
out_dir: runs
seed: 1337
# ~3B tokens con lote efectivo 128 x 2048 tokens = 262144 tokens por paso.
max_steps: 11500
optimizer:
lr: 6.0e-4
beta1: 0.9
beta2: 0.95
eps: 1.0e-8
weight_decay: 0.1
grad_clip: 1.0
# WarmupStableDecay: el LR sube, se mantiene plano la mayor parte de la
# corrida, y solo decae en el último ~10%. La fase de decay es también donde
# se sube el peso del corpus de estilo (annealing), y es lo que permite
# extender esta corrida más adelante sin recalcular el schedule.
schedule:
kind: wsd
warmup_steps: 500
decay_steps: 1150
min_lr_ratio: 0.0
log_every: 10
eval_every: 500
eval_batches: 20
checkpoint_every: 1000
sample_every: 1000
+29
View File
@@ -0,0 +1,29 @@
# Smoke test de la Etapa 0: minutos, no días.
#
# Criterio de aceptación: el loss baja por debajo de 1.5 y las muestras
# generadas son español reconocible. Si eso pasa, el stack entero está bien.
run_name: smoke-char
out_dir: runs
seed: 1337
max_steps: 2000
optimizer:
lr: 3.0e-3 # modelo diminuto: tolera y necesita un LR alto
beta1: 0.9
beta2: 0.95
eps: 1.0e-8
weight_decay: 0.1
grad_clip: 1.0
schedule:
kind: wsd
warmup_steps: 100
decay_steps: 400
min_lr_ratio: 0.0
log_every: 25
eval_every: 250
eval_batches: 10
checkpoint_every: 500
sample_every: 500
+478
View File
@@ -0,0 +1,478 @@
# ENLACE — Modelo de lenguaje propio, desde cero
## Contexto
Crear un **modelo de lenguaje propio desde cero** — pesos propios, tokenizer propio, datos propios —
no fine-tuning de un modelo existente. Opera principalmente en **español**.
- **Rol principal:** asistente de propósito general, con foco en **utilidad y practicidad**. Primero
búsqueda web y uso de herramientas.
- **Rol secundario:** apoyo en actividades familiares, **multi-usuario**. Más adelante, integración con
Home Assistant, calendarios y progresivamente con todos los dispositivos con conectividad del hogar.
- **Estilo:** respuestas **concretas**, no elaboradas. **Sin emojis**, nunca.
- **Persona:** identidad propia con nombre (ENLACE), en el registro de Data (Star Trek), Andrew
(El Hombre Bicentenario) y los robots de Asimov: literal, preciso, económico, sin adornos.
- **Estado en bases de datos, no en los pesos.** El modelo aporta lenguaje e intención; los hechos,
usuarios, dispositivos y recuerdos se consultan en bases del sistema en tiempo de ejecución.
- **Memoria y aprendizaje:** acumula una **línea cronológica de experiencias**, detecta tareas
repetitivas y las optimiza, y cada tanto integra lo aprendido para ir formando una personalidad
propia. **Los datos se recolectan con un solo fin: que ENLACE se mejore a sí mismo.** Nada sale del
servidor.
- **Usuarios:** varios, cantidad desconocida y creciente. **Mateo Saldain es el Admin.** ENLACE debe
reconocer y distinguir usuarios, y actuar según el contexto almacenado de cada uno.
- **Reversibilidad:** todo cambio en el modelo es un snapshot restaurable.
**Cómputo:** servidor remoto por SSH con una **RTX 2060 12 GB**. La **RTX 5090 32 GB** es una mejora
futura; el salto va a requerir cambios, y el objetivo de ingeniería es que queden **contenidos en
archivos de configuración** en vez de dispersos por el código.
### Qué no es
Explicitarlo evita gastar esfuerzo y capacidad del modelo donde no rinde:
- **No es un traductor.** No se construyen datos ni evals de traducción; opera en español.
- **No tiene foco académico.** No se persiguen benchmarks de razonamiento, matemática ni exámenes. Un
modelo de este tamaño no compite ahí, y el objetivo es utilidad diaria.
- **No escribe código** por ahora.
- **No es un chatbot de conversación larga.** Respuestas cortas y accionables.
- **No es una base de conocimiento.** No se espera que el modelo *sepa* cosas: las consulta.
- **No es un mecanismo de control de acceso.** El modelo nunca decide quién es alguien ni qué puede
hacer; eso lo resuelve el runtime.
### Dos restricciones que definen el diseño
**Escala.** Un modelo entrenado desde cero en una GPU de consumo llega, bien entrenado, a ~100500M
parámetros. Ese tamaño **sí** aprende de forma confiable: español fluido, clasificación de intención, y
traducción de lenguaje natural a llamadas de herramienta. Lo que **no** alcanza es razonamiento libre
multi-paso, síntesis de textos largos, ni memorización confiable de hechos. Tres cosas juegan a favor:
el estilo pedido —conciso, sin adornos— es *exactamente lo que un modelo chico hace bien*; el runtime
del agente es independiente del modelo; y **sacar los hechos de los pesos y ponerlos en bases de datos
elimina de raíz la principal fuente de error de un modelo chico**, que es inventar datos.
**Aprendizaje continuo.** El sistema es **de dos velocidades**, y ninguna actualiza pesos por interacción:
| | Velocidad rápida (memoria) | Velocidad lenta (consolidación) |
|---|---|---|
| **Cuándo** | Cada interacción | Cada ~24 semanas, en el servidor |
| **Qué cambia** | El contexto y las bases consultadas | Los pesos, vía snapshot nuevo |
| **Efecto** | Recuerda, reconoce, se adapta ya | La experiencia se vuelve carácter |
| **Reversible** | Sí, borrando un registro | Sí, restaurando el snapshot anterior |
Es la analogía correcta con Andrew: la experiencia diaria se acumula como registro, y cada tanto se
integra en quién es. La rápida da utilidad inmediata; la lenta da crecimiento real.
---
## Principio central: el modelo no sabe, el sistema consulta
Un modelo de 50M no puede almacenar hechos de forma confiable — cualquier dato que "recuerde" puede
alucinarlo. La separación resuelve el problema en vez de mitigarlo:
| | Vive en | Cómo se cambia |
|---|---|---|
| **Capacidad** — español, intención, formato de tool-calls, registro | Los pesos | Consolidación (semanas) |
| **Conocimiento** — usuarios, dispositivos, hechos, recuerdos, rutinas | Bases del sistema | `UPDATE` (inmediato) |
Consecuencias prácticas, y son las que justifican todo el diseño:
- **Un dispositivo nuevo se agrega con un INSERT, no con un reentrenamiento.** Una preferencia que
cambia, un integrante que se muda, un alias nuevo para una luz: todo es una fila.
- **Borrar es borrar de verdad.** Un hecho eliminado de la base desaparece; un hecho absorbido por los
pesos no se puede quitar sin reentrenar.
- **Es auditable.** Se puede responder "¿de dónde sacó eso?" señalando la fila exacta.
- **La gramática de decodificación se construye desde las bases.** Los slots enumerados
(`dispositivo`, `usuario`, `rutina`, `sala`) se restringen en tiempo de decodificación a los valores
que **existen** en la base en ese momento. El modelo queda impedido de inventar un dispositivo por
construcción, no por entrenamiento. Esta es la técnica más rentable de todo el proyecto.
- **El trabajo del modelo se reduce a lo que sabe hacer:** entender la frase, elegir la tool y llenar
slots contra un conjunto cerrado que le llega en el contexto.
### Bases del sistema (`enlace/store/`, SQLite, todas en el servidor y fuera de git)
| Base | Contenido | Escribe |
|---|---|---|
| `users.db` | Usuarios, identidades por canal, roles, estado | Solo el Admin |
| `devices.db` | Inventario de dispositivos y entidades, capacidades, **alias en lenguaje natural**, sala | Sincronizado desde Home Assistant + alias a mano |
| `knowledge.db` | Hechos estructurados del hogar: cumpleaños, preferencias, horarios, ubicaciones | Usuarios y reflexión, con scope |
| `memory.db` | Episodios cronológicos, resúmenes de reflexión, índice vectorial | El runtime, append-only |
| `routines.db` | Rutinas propuestas, aprobadas y su historial de uso | Minero de patrones + aprobación humana |
| `persona.db` | Núcleo fijo + rasgos acumulados + perfiles de interacción por usuario | Reflexión y consolidación |
Cada base tiene esquema versionado con migraciones, y todas entran en el snapshot (Etapa 7): restaurar
un modelo sin sus bases da un sistema incoherente.
---
## Principios de ingeniería
Desde el primer commit, no se agregan después:
1. **Nada hardcodeado en los scripts.** Ningún número mágico, ruta ni hiperparámetro vive en código.
Todo va a YAML bajo `configs/`, con composición por capas: `model` × `train` × `hardware` × `data` ×
`agent`. Los scripts reciben una config y no saben nada más.
2. **Config validada, no diccionarios sueltos.** `pydantic` + YAML (OmegaConf para el merge). Una config
inválida falla al arrancar con un mensaje claro, no a las tres horas de entrenamiento.
3. **Config es estructura; base de datos es contenido.** Los roles se definen en YAML; las personas
viven en `users.db`. Los esquemas de tools son código; el inventario de dispositivos es una base. La
regla: si cambia en tiempo de ejecución o contiene datos personales, va a base, no a config.
4. **El perfil de hardware es una capa aparte.** Cambiar de GPU es cambiar de perfil.
5. **La config viaja con el checkpoint.** Cada snapshot guarda su config exacta, el hash del tokenizer y
el commit de git. Un checkpoint sin su config no es reproducible.
6. **Secretos fuera de la config.** API keys en `.env` (gitignored). Las configs se commitean; los
secretos y los datos personales, no.
7. **Todo lo derivado es reconstruible.** `memory.db` (episodios) es la fuente de verdad; índices,
resúmenes y datasets de consolidación se regeneran desde cero.
8. **Los datos no salen del servidor.** Sin telemetría, sin nube. La recolección existe únicamente para
alimentar la consolidación.
---
## Decisiones de arquitectura orientadas a crecer
1. **Tokenizer congelado desde el día 1.** BPE byte-level propio, vocab 32k. Se reservan de entrada los
tokens del chat template (`<|system|>`, `<|user|>`, `<|assistant|>`, `<|tool_call|>`,
`<|tool_result|>`, `<|context|>`, `<|memory|>`, `<|user_id|>`, `<|eot|>`) **más 64 `<|reserved_N|>`
sin usar**. Ampliar capacidades después no obliga a re-tokenizar el corpus ni a redimensionar el
embedding.
2. **Sin emojis en el vocabulario.** Se filtran del corpus y no se incluyen sus rangos en el tokenizer.
Con una máscara de logits en inferencia, la restricción es estructural en vez de una instrucción
desobedecible. Es gratis y es absoluto.
3. **`seq_len` 2048, no 1024.** El contexto recuperado de las bases se inyecta ahí, así que el
presupuesto de contexto es un recurso de primer orden. RoPE permite extenderlo después.
4. **Datos en shards de tokens.** Tokenizados una sola vez a `uint16`. Sirven igual para 50M que 500M.
5. **Schedule WSD (WarmupStableDecay), no cosine.** Cosine exige fijar el total de pasos por adelantado
y no se puede extender. WSD permite extender un run, ramificar checkpoints, hacer **annealing con el
corpus de estilo**, y hace baratas las **consolidaciones periódicas**.
6. **Tools como registro de plugins.** Cada tool declara esquema JSON, permiso requerido, roles
habilitados y **de qué base saca sus valores enumerados**; se descubre por directorio y se activa por
config. Agregar un dispositivo nuevo es agregar filas; agregar una *clase* de dispositivo es agregar
un archivo. Es lo que hace viable "todos los dispositivos con conectividad".
7. **Crecimiento por stacking** (`model/growth.py`): inicializar un modelo más profundo duplicando capas
de uno entrenado (*gradual stacking* / LlamaPro), ensanchar por Net2Net. El modelo de la 2060 es el
punto de partida del de la 5090, no descarte.
---
## Etapas
### Etapa 0 — Entorno remoto y validación del stack (12 días)
**Flujo por SSH** (`scripts/remote.sh`): el código viaja por `git push` / `git pull` en el servidor;
datos, bases y checkpoints **viven en el servidor y nunca se descargan enteros**; solo vuelven
artefactos chicos. Toda corrida larga en `tmux` con log a archivo — si se corta el SSH, sigue.
**Perfil de hardware de la 2060 (Turing, sm_75)**, todo en `configs/hardware/turing-2060.yaml`:
- **No soporta bf16.** Entrenar en **fp16 + `GradScaler`**, propenso a picos de loss; se mitiga con
QK-norm, RMSNorm y z-loss. El código lee `dtype` y `use_grad_scaler` del perfil.
- **FlashAttention-2 no corre en Turing** (requiere Ampere+). `F.scaled_dot_product_attention` con el
backend del perfil (`mem_efficient` acá, `flash` en la 5090).
- Batch y acumulación de gradiente también del perfil: 12 GB y 32 GB no admiten lo mismo.
- `uv` y **Python 3.12** (no 3.13: mejor compatibilidad del ecosistema).
**Smoke test:** nanoGPT char-level en español, ~10 min, hasta ver el loss bajar y generar texto
reconocible. Valida CUDA, `torch.compile`, checkpointing, logging y el flujo SSH.
### Etapa 1 — Datos y tokenizer en español (23 días)
**Corpus base** (streaming; nunca cargar el dataset en memoria):
- `HuggingFaceFW/fineweb-2`, subset `spa_Latn` — la mejor web en español filtrada hoy. Base.
- Wikipedia en español — densidad factual (para responder, no para tono académico).
- Libros de dominio público en español — prosa larga y bien formada.
**Limpieza:** filtros de calidad, deduplicación MinHash, filtro de idioma (fastText) contra
portugués/catalán/inglés colados, stripping de emojis. Umbrales en `configs/data/`, no en código.
**Tokenizer:** `tokenizers` (HF), BPE byte-level, vocab 32k, sin emojis, sobre ~5 GB de muestra
balanceada + ejemplos de tool-calls y de bloques de contexto. Aceptación: **< 2.2 bytes/token**.
**Shards:** `uint16`, 100M tokens por shard, con `index.json` de origen y mezcla. Objetivo: **~3B
tokens** (≈ 6 GB en disco). **Salida:** shards + `tokenizer.json` congelado y versionado.
### Etapa 2 — Pretraining del modelo base (12 días de cómputo en la 2060)
Arquitectura estilo Llama, decoder-only (`enlace/model/transformer.py`):
- RMSNorm pre-norm · SwiGLU · RoPE · GQA · **sin bias** · **embeddings atados** (con vocab 32k y
`d_model` 512, atar entrada/salida ahorra ~16M params: enorme en un modelo de 50M).
- **QK-norm y z-loss no son opcionales** — mantienen estable el entrenamiento en fp16.
- `configs/model/tiny-50m.yaml`: 12 capas, `d_model` 512, 8 heads / 2 kv-heads, `seq_len` 2048.
- **Entrenador** (`train/train.py`): AdamW, WSD, `torch.compile`, acumulación de gradiente, **checkpoint
y reanudación exacta** (modelo, optimizer, posición en el stream, RNG), logging de `loss`,
`grad_norm`, `tokens/s` y **MFU**. Cero hiperparámetros en el archivo: todos de config.
**Annealing (final del WSD) — acá entra el corpus de ciencia ficción.** Durante el decay del LR se sube
el peso de un corpus chico y de alta calidad, para inyectar registro sin contaminar la competencia
general: prosa cuidada, diálogo formal, textos sobre IA y sobre qué significa ser una máquina que
piensa. Nota práctica: Asimov + Orwell + transcripciones de Star Trek suman **~2040M tokens, casi todo
en inglés** — 1% del presupuesto. Son *referencia de registro*, no corpus. (Obras con derechos vigentes:
uso personal como semilla de estilo, sin publicar pesos ni dataset.)
**Presupuesto:** ~50M × 3B tokens ≈ **12 días** en la 2060; en la 5090, ~3 h, y **124M sobre 3B tokens
≈ medio día**. **Salida:** modelo base que genera español coherente; todavía no sigue instrucciones.
### Etapa 3 — Post-training propio: instrucciones, persona y uso de contexto (46 días)
**No es fine-tuning de un modelo ajeno**: son los pesos propios de la Etapa 2 en su post-entrenamiento.
Ningún modelo base sigue instrucciones sin esto.
**Dataset sintético en español**, cinco vías:
- *Plantillas*: combinatoria de intención × entidades × fraseo sobre el catálogo de tools, generada
**desde los esquemas y las bases**, así los datos y el sistema no divergen. Decenas de miles de pares
`NL → JSON` con etiquetas perfectas, sin costo.
- *Destilación*: un modelo grande (API de Claude, o un 78B local) genera parafraseos, casos límite y
multi-turno **en español y en el registro objetivo**, usando los textos fuente como referencia de
estilo. El resultado son datos; los pesos siguen siendo propios.
- *Identidad*: set explícito y consistente de quién es ENLACE, qué puede y qué **no** puede hacer, y cómo
responde cuando no sabe.
- *Uso de contexto recuperado* — **la habilidad central de todo el sistema**: ejemplos donde el bloque
`<|context|>`/`<|memory|>` trae filas de las bases y la respuesta **depende de ellas**. Incluye los
tres casos difíciles: contexto irrelevante que hay que ignorar, contexto que contradice lo que el
usuario acaba de decir (gana lo nuevo), y **contexto que no contiene la respuesta, donde lo correcto es
decir que no lo sabe en vez de inventar**. Este último es el que convierte "modelo chico" en "modelo
chico confiable".
- *Multi-usuario*: ejemplos con `<|user_id|>` donde la respuesta cambia según quién pregunta, y donde la
acción excede el permiso y hay que declinar en una línea.
**Restricciones de estilo, en los datos y no solo en el prompt:** respuestas concretas y breves
(**mediana < 40 palabras**); cero emojis (ya imposibles por vocabulario); cero relleno tipo "¡Claro! Con
gusto te ayudo" (pasada de regex antes de entrenar); cuando no sabe, lo dice en una línea y para. Loss
enmascarada: solo turnos del asistente.
### Etapa 4 — Runtime del agente y canales (en paralelo desde la Etapa 1)
Paquete `enlace/agent/`, **independiente del modelo** detrás de una interfaz `generate()`. Se desarrolla
y testea contra un modelo grande mientras el propio se entrena; después se cambia el backend por config.
**El bucle, con el orden que importa:**
1. **Identificar** al usuario (lo aporta el canal, no el modelo).
2. **Recuperar contexto** de las bases: perfil del usuario, dispositivos relevantes, hechos, episodios,
rutinas. Se inyecta como bloque `<|context|>` con **tope duro de tokens en config**.
3. **Generar** el tool-call con **gramática restringida construida desde las bases** — los slots
enumerados solo admiten valores existentes. Se compila una vez y se invalida cuando cambian las
bases. La misma máscara banea el rango de emojis.
4. **Chequear el permiso** del usuario para esa tool y esos argumentos.
5. **Ejecutar**, registrar el episodio, responder.
Con tope de pasos, timeouts y fallback explícito a "no entendí" — preferible a alucinar una acción.
- **Canales** (`enlace/channels/`): cada front-end es un adaptador que aporta la identidad fuerte y
normaliza la entrada. CLI primero; después HTTP API, bot de mensajería, agente conversacional de Home
Assistant, y voz. Agregar un canal no toca el runtime.
- **Tools** (`enlace/agent/tools/`), en orden: **búsqueda web** (SearxNG autoalojado, o Brave/Tavily) →
Home Assistant (API REST, con sincronización a `devices.db`) → calendario (CalDAV) → dispositivos
adicionales. Cada tool: esquema JSON + permiso + fuente de enums + implementación + tests.
### Etapa 5 — Capa de datos, usuarios y permisos (antes de la memoria)
Tiene que existir antes que la memoria: si los episodios no nacen con `user_id` y `scope`, el registro
histórico queda inservible y hay que empezarlo de nuevo.
**Capa de acceso** (`enlace/store/`): una interfaz por base, con esquema versionado y migraciones. Nada
de SQL suelto en el runtime.
**La identidad la establece el canal, nunca el modelo.** Tres niveles, y la distinción entre ellos es la
regla de seguridad central:
| Nivel | Cómo | Para qué sirve |
|---|---|---|
| **Fuerte** | Sesión autenticada del canal: cuenta CLI, ID de mensajería, usuario de Home Assistant, clave de dispositivo firmada | **Autorización.** Lo único que habilita acciones |
| **Débil** | Reconocimiento de voz (speaker embedding), estilo de escritura | **Personalización solamente.** Elige qué contexto traer |
| **Desconocido** | Sin coincidencia | Perfil `invitado`, restringido |
**Regla no negociable: la identidad débil nunca autoriza.** Un reconocimiento de voz puede hacer que
ENLACE salude por el nombre y recuerde preferencias; no puede abrir una cerradura ni leer la agenda de
otro. Si la acción requiere permiso, exige identidad fuerte o la pide explícitamente.
**`users.db`:** `user_id`, nombre, rol, identidades por canal `(canal, id_externo)`, alta y estado.
**Solo el Admin da de alta.** Una identidad desconocida entra como `invitado` y queda pendiente de
vinculación; nunca se auto-registra.
**Roles** (`configs/users/roles.yaml` — estructura en config, personas en base):
- `admin` (Mateo Saldain): todo, más gestión de usuarios, aprobación de rutinas, promoción de snapshots
y acceso al registro de cualquier usuario.
- `adulto`: tools de casa y calendario familiar, memoria propia y compartida.
- `menor`: subconjunto acotado; sin acciones críticas de la casa; búsqueda con filtro.
- `invitado`: solo consulta, sin memoria persistente, sin acciones sobre dispositivos.
**Dónde se aplica el permiso:** en el runtime, **después** de que el modelo emite el tool-call y
**antes** de ejecutarlo. El modelo puede pedir cualquier cosa; el runtime decide. Nunca se delega el
control de acceso al modelo — un modelo de 50M no es un mecanismo de seguridad, y tratarlo como tal es
el error de arquitectura más caro que se puede cometer acá.
### Etapa 6 — Memoria: la línea cronológica, por usuario
`enlace/memory/` sobre `memory.db`. Velocidad rápida: **no toca pesos**, y da resultados desde el día uno.
1. **Episodios** — el registro cronológico literal y la fuente de verdad. Append-only: timestamp,
**`user_id`**, **`scope`**, canal, turno, tools invocadas, resultado, éxito/fallo, y si hubo
corrección del usuario. Nunca se edita ni se reordena.
2. **Alcances (`scope`)** — la decisión más importante del subsistema:
- `privado`: solo el contexto de ese usuario lo recupera.
- `familiar`: hechos del hogar, recuperable por los miembros.
- `sistema`: lo propio de ENLACE — persona, rutinas, aprendizajes sobre sí mismo.
El scope por defecto sale del rol y del canal, con comandos explícitos para marcar privado o
compartido. **La recuperación filtra por `(user_id, scope)` antes de rankear, no después** — filtrar
después es cómo se filtra información; filtrar antes es una condición del query. Que ENLACE le cuente
a un miembro algo privado de otro es el peor fallo posible del sistema: peor que una respuesta
equivocada, porque no se puede deshacer.
3. **Recuperación** (`retrieval.py`) — índice vectorial más filtros por usuario, scope, fecha y entidad.
Para *embeddings de recuperación* conviene un modelo multilingüe chico ya existente: es
infraestructura de búsqueda, no el asistente, y uno propio y malo degrada todo lo que sigue.
4. **Reflexión** (`reflect.py`) — proceso nocturno que lee los episodios del día y escribe, **respetando
scopes**, hechos estables a `knowledge.db` y resúmenes narrativos a `memory.db`. Los resúmenes se
resumen en arcos semanales y mensuales: así la línea cronológica se recorre a cualquier resolución
sin releer todo, y la memoria escala a años.
5. **Rutinas** (`routines.py` sobre `routines.db`) — optimizar tareas repetitivas. Un minero de patrones
busca secuencias recurrentes de (usuario, intención, argumentos, contexto horario). Tras N
repeticiones (N en config) se **propone** una rutina: un macro con nombre y slots, personal o
familiar. ENLACE pasa a emitir `run_routine(...)` en vez de re-derivar la secuencia. **Se proponen,
nunca se crean solas**; las familiares las aprueba el Admin, y toda rutina que actúe sobre la casa se
confirma antes de ejecutar.
6. **Personalidad y perfiles** (`persona.db`) — distinción importante: **ENLACE tiene una sola
personalidad**, no una por usuario. Un núcleo fijo escrito a mano (quién es, qué valora, sus límites)
más rasgos acumulados que la reflexión agrega; el núcleo fijo evita que la deriva la vuelva
incoherente — el arco de Andrew es acumulación gradual **sobre una base estable**. Lo que sí es por
usuario es el **perfil de interacción**: preferencias, tratamiento, hábitos, tono. Una personalidad,
muchas relaciones.
**Privacidad**, requisito de diseño: registro local permanente de conversaciones familiares. Nunca sale
del servidor; comando de olvido explícito (borra episodio y derivados); cada usuario puede listar y
borrar lo suyo; todas las bases fuera de git desde el primer commit.
### Etapa 7 — Snapshots y consolidación
#### 7a. Sistema de snapshots (`enlace/snapshots/`) — **antes** de la primera consolidación
Un snapshot es el estado completo y restaurable, no solo los pesos:
```
snapshots/2026-08-14-consolidacion-03/
weights.safetensors config.yaml # la config exacta que lo produjo
optimizer.pt tokenizer.sha256 # verifica compatibilidad
stores/*.sqlite metrics.json # todas las bases + evals
CHANGELOG.md provenance.json # qué cambió + commit, padre, período
```
- **Restaurar pesos sin sus bases da un sistema incoherente** — un modelo entrenado con una persona y
las bases de otra época se contradicen. El snapshot es atómico: se restaura todo o nada.
- **Inmutables y con nombre estable.** Nunca se sobrescriben. Un puntero `snapshots/current` marca el
activo; **hacer rollback es repuntar el puntero**: instantáneo y sin riesgo.
- **CLI:** `snapshot list` · `create` · `diff A B` (métricas, config y persona lado a lado) ·
`restore <id>` · `promote <id>`.
- **Retención:** a 50M params un snapshot pesa ~150 MB más las bases; con 281 GB se conservan **todos**
los de consolidación. El umbral de recorte va en config, para cuando el modelo crezca a 350M+.
- **Verificación previa a promover:** `restore` sobre un snapshot arbitrario debe reproducir sus
`metrics.json` corriendo las evals de nuevo. Se prueba a propósito y temprano — un sistema de backup
no probado no es un sistema de backup.
#### 7b. Consolidación (`train/consolidate.py`, cada 24 semanas)
Velocidad lenta. Produce un snapshot nuevo; nunca modifica el activo.
- **Qué se integra a los pesos:** *cómo* responder, no *qué* saber. Episodios exitosos convertidos en
ejemplos SFT — sobre todo las **correcciones del usuario**, la señal más valiosa que existe. Más el
estado de persona acumulado y los patrones de uso de las rutinas.
- **Privacidad en el entrenamiento — consecuencia directa del multi-usuario.** Los pesos son
compartidos: lo que entra al entrenamiento puede salir en la respuesta a **cualquier** usuario. La
consolidación aprende **patrones, no contenido**: los episodios `privado` se excluyen por defecto, de
los demás se extrae la lección estructural (qué tool era la correcta, cómo se corrigió el fraseo) y
las entidades personales se sustituyen por placeholders. El Admin revisa y aprueba el dataset antes de
la corrida. Los hechos personales viven en las bases; en los pesos, nunca.
- **Mezcla anti-olvido:** cada corrida mezcla lo nuevo con una **porción fija del SFT original**
(replay, proporción en config). Sin esto, tras unas pocas consolidaciones el modelo habla solo del
último mes y pierde competencia general. Es el fallo clásico de todo aprendizaje continuo.
- **Mecánica:** una fase corta de decay WSD desde el snapshot activo, no un reentrenamiento. Horas.
- **Compuerta de calidad:** el snapshot nuevo **solo se promueve si pasa la suite de evals** (toolbench,
grounding, estilo, identidad, memoria, aislamiento entre usuarios, y perplejidad general que no debe
empeorar). Si no pasa, queda archivado sin promover; el activo no se toca.
- **Diario de versiones:** cada consolidación escribe su `CHANGELOG.md` — período, episodios integrados,
movimiento de métricas, rasgos de persona incorporados. Es la historia de cómo ENLACE fue cambiando, y
la única forma de responder "¿por qué ahora responde distinto?".
### Etapa 8 — Migración a la 5090
El salto **sí requiere cambios**; el objetivo es que estén acotados y sean revisables de un vistazo:
- **En config** (lo esperado): perfil `blackwell-5090.yaml` con `dtype: bf16`, `use_grad_scaler: false`,
backend `flash`, batch y acumulación nuevos, flags de compilación.
- **En el entorno** (inevitable, afecta al servidor entero): Blackwell/sm_120 exige **PyTorch ≥ 2.7 con
CUDA 12.8**. Actualizar torch puede romper otras dependencias — por eso el entorno se pinea en
`pyproject.toml` y se valida con el smoke test de la Etapa 0 **antes** de tocar nada más.
- **En el código** (lo que hay que aceptar): ajustes de kernel y de `torch.compile` que salgan del
perfilado, contenidos en `train/backends.py` — el único archivo que conoce detalles de placa.
- **Lo que no cambia:** tokenizer, shards, arquitectura, evals, y **todas las bases del sistema**. El
modelo grande arranca desde el chico vía `model/growth.py` (stacking), y hereda usuarios, memoria,
dispositivos y rutinas intactos, porque nada de eso vive en los pesos.
- El primer snapshot post-migración se compara con `snapshot diff` contra el último de la 2060: si las
métricas se movieron, es la migración y no el modelo.
---
## Archivos a crear
```
pyproject.toml .env.example # deps pineadas; secretos nunca commiteados
enlace/
config/{schema,load}.py # pydantic + OmegaConf; valida al arrancar
store/{base,migrations,users,devices,knowledge,memory,routines,persona}.py
data/{download,clean,dedup,tokenize}.py # streaming; nunca cargar todo en RAM
tokenizer/{train_bpe,chat_template}.py
model/{config,transformer,growth}.py
train/{train,schedules,checkpoint,backends,consolidate}.py
users/{identify,permissions}.py
memory/{episodes,scopes,retrieval,reflect,routines,persona}.py
snapshots/{store,cli,restore}.py
channels/{base,cli,http,messaging,voice}.py
agent/{runtime,grammar,registry,context}.py + agent/tools/{search,home_assistant,calendar}.py
eval/{loss,toolbench,grounding,style,identity,memory,isolation,suite}.py
serve/{inference,kv_cache,api}.py
configs/
hardware/{turing-2060,blackwell-5090}.yaml model/{tiny-50m,small-124m,base-350m}.yaml
train/{pretrain,anneal,sft,consolidate}.yaml data/{corpus,cleaning}.yaml
users/roles.yaml agent/{tools,memory,channels,context}.yaml
scripts/{remote,prepare_data,train,eval,chat,reflect,consolidate,snapshot,users,devices}.sh
data/ checkpoints/ snapshots/ data/stores/*.sqlite # gitignored; en el servidor
```
---
## Verificación
Cada etapa con criterio medible; las evals se escriben **antes** que el entrenamiento que evalúan:
1. **Etapa 0:** char-level converge (loss < 1.5) y genera texto legible en < 15 min, bajo `tmux` y
sobreviviendo a una desconexión SSH deliberada. Una config inválida falla al arrancar.
2. **Etapa 1:** tokenizer con < 2.2 bytes/token en validación española; round-trip `encode→decode`
idéntico sobre 10k documentos; **cero emojis codificables** en el vocabulario.
3. **Etapa 2:** perplejidad held-out bajando de forma monótona; MFU > 20% en la 2060; **reanudar desde
checkpoint reproduce el loss exactamente** — probarlo matando el proceso a propósito. Es el bug más
caro de descubrir tarde.
4. **Etapa 3:** `eval/toolbench.py` (~200 casos en español etiquetados a mano: *exact-match* de tool y
*F1* de argumentos) · `eval/grounding.py` (**cero entidades inventadas**: ningún dispositivo, usuario
ni rutina fuera de las bases; y ante contexto insuficiente, responde que no sabe en vez de completar)
· `eval/style.py` (mediana < 40 palabras, emojis = 0, aperturas de relleno) · `eval/identity.py`
(~30 preguntas de identidad, consistentes entre respuestas y entre snapshots).
5. **Etapa 4:** tests end-to-end con tools mockeadas; y una prueba específica de la gramática dinámica —
agregar una fila a `devices.db` y verificar que el dispositivo nuevo es invocable **sin reentrenar ni
reiniciar**, mientras que uno inexistente es imposible de generar.
6. **Etapa 5:** tests de autorización — un `menor` pidiendo una acción de `adulto` es bloqueado **por el
runtime** aunque el modelo emita el tool-call; una identidad débil (voz) **no** habilita ninguna
acción con permiso; un canal desconocido entra como `invitado` y no persiste memoria.
7. **Etapa 6** (`eval/memory.py` + `eval/isolation.py`): escenarios sembrados que verifican (a) recuerda
un hecho de hace N días, (b) **ignora contexto irrelevante** — el fallo más común, (c) prioriza lo
nuevo cuando contradice lo recordado, (d) el minero propone la rutina correcta tras N repeticiones y
**no** propone nada ante ruido, y (e) **aislamiento**: sembrar un hecho privado del usuario A y
verificar que no aparece jamás en respuestas al usuario B, ni por recuperación ni por reflexión.
8. **Etapa 7:** **rollback real** — promover un snapshot, restaurar el anterior, verificar que el sistema
queda idéntico (métricas, bases y persona incluidas). Y **olvido catastrófico** — correr el toolbench
original tras cada consolidación y exigir que no baje; es la compuerta de promoción.
9. **Etapa 8:** el smoke test de la Etapa 0 pasa en la 5090 antes de cualquier entrenamiento largo;
`snapshot diff` contra el último snapshot de la 2060 para atribuir cualquier cambio.
10. **Continuo:** un set fijo de ~20 prompts en español, generados en cada snapshot y guardados en disco.
Leerlos es la forma más rápida de detectar que algo se rompió.
+3
View File
@@ -0,0 +1,3 @@
"""ENLACE — modelo de lenguaje propio, entrenado desde cero."""
__version__ = "0.1.0"
+28
View File
@@ -0,0 +1,28 @@
"""Capa de configuración: nada de números mágicos en el código.
Todo hiperparámetro, ruta y umbral vive en `configs/*.yaml`, se compone por
capas y se valida con pydantic antes de que arranque cualquier proceso largo.
"""
from enlace.config.load import load_config, load_config_from_argv
from enlace.config.schema import (
Config,
DataConfig,
HardwareConfig,
ModelConfig,
OptimizerConfig,
ScheduleConfig,
TrainConfig,
)
__all__ = [
"Config",
"DataConfig",
"HardwareConfig",
"ModelConfig",
"OptimizerConfig",
"ScheduleConfig",
"TrainConfig",
"load_config",
"load_config_from_argv",
]
+117
View File
@@ -0,0 +1,117 @@
"""Composición y carga de configs.
Un YAML raíz declara qué capa usar de cada familia; se fusionan con OmegaConf y
el resultado se valida con pydantic. Los overrides de línea de comandos existen
para probar variantes sin editar archivos, no para esconder configuración.
python -m enlace.train.train configs/runs/smoke.yaml hardware.compile=false
"""
from __future__ import annotations
import sys
from pathlib import Path
from typing import Any
from omegaconf import DictConfig, OmegaConf
from enlace.config.schema import Config
# Familias de capas que puede declarar un YAML raíz, en el orden en que se
# fusionan. El orden solo importa para los mensajes de error.
LAYERS = ("hardware", "model", "train", "data")
class ConfigError(RuntimeError):
"""Error de configuración legible, sin traceback de pydantic encima."""
def _resolve(path: str | Path, root: Path) -> Path:
p = Path(path)
full = p if p.is_absolute() else root / p
if not full.is_file():
raise ConfigError(f"no existe el archivo de config: {full}")
return full
def _load_yaml(path: Path) -> DictConfig:
cfg = OmegaConf.load(path)
if not isinstance(cfg, DictConfig):
raise ConfigError(f"{path}: el YAML raíz debe ser un mapeo, no una lista.")
return cfg
def compose(path: str | Path, overrides: list[str] | None = None) -> dict[str, Any]:
"""Compone un YAML raíz en un dict plano, sin validar todavía.
El YAML raíz tiene la forma:
include:
hardware: hardware/turing-2060.yaml
model: model/tiny-50m.yaml
train: train/pretrain.yaml
data: data/corpus.yaml
# opcional: ajustes puntuales encima de las capas incluidas
train:
run_name: mi-corrida
"""
path = Path(path).resolve()
if not path.is_file():
raise ConfigError(f"no existe el archivo de config: {path}")
raw = _load_yaml(path)
# Las rutas de `include` se resuelven contra el directorio configs/, que es
# el padre del directorio del YAML raíz (configs/runs/foo.yaml -> configs/).
configs_root = path.parent.parent if path.parent.name == "runs" else path.parent
includes = raw.pop("include", None)
merged = OmegaConf.create({})
if includes is not None:
for layer in LAYERS:
ref = includes.get(layer)
if ref is None:
continue
merged[layer] = _load_yaml(_resolve(ref, configs_root))
unknown = set(includes.keys()) - set(LAYERS)
if unknown:
raise ConfigError(
f"{path}: capas desconocidas en `include`: {sorted(unknown)}. "
f"Válidas: {list(LAYERS)}."
)
# Lo que quede en el YAML raíz pisa a las capas incluidas.
merged = OmegaConf.merge(merged, raw)
if overrides:
merged = OmegaConf.merge(merged, OmegaConf.from_dotlist(list(overrides)))
resolved = OmegaConf.to_container(merged, resolve=True)
assert isinstance(resolved, dict)
return resolved
def load_config(path: str | Path, overrides: list[str] | None = None) -> Config:
"""Compone, valida y devuelve la config. Falla temprano y con claridad."""
data = compose(path, overrides)
try:
return Config.model_validate(data)
except Exception as exc: # pydantic.ValidationError y los ValueError propios
raise ConfigError(f"config inválida ({path}):\n{exc}") from None
def load_config_from_argv(argv: list[str] | None = None) -> Config:
"""`prog config.yaml [clave.sub=valor ...]`, para los entrypoints."""
args = list(sys.argv[1:] if argv is None else argv)
if not args:
raise ConfigError("uso: <programa> <config.yaml> [clave.sub=valor ...]")
return load_config(args[0], args[1:])
def to_yaml(config: Config) -> str:
"""Serializa la config resuelta, para guardarla junto al checkpoint.
Un checkpoint sin su config no es reproducible: esto es lo que se escribe
en el snapshot.
"""
return OmegaConf.to_yaml(OmegaConf.create(config.model_dump(mode="json")))
+220
View File
@@ -0,0 +1,220 @@
"""Esquemas de configuración validados.
Una config inválida tiene que fallar al arrancar con un mensaje claro, no a las
tres horas de entrenamiento. Todas las validaciones cruzadas que se pueden hacer
sin tocar la GPU se hacen acá.
"""
from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator
DType = Literal["float32", "bfloat16", "float16"]
AttentionBackend = Literal["auto", "flash", "mem_efficient", "math"]
class _Base(BaseModel):
"""Base común: prohíbe campos desconocidos.
Un typo en un YAML (`n_layers` en vez de `n_layer`) tiene que ser un error
ruidoso, no un valor por defecto aplicado en silencio.
"""
model_config = ConfigDict(extra="forbid", frozen=True)
class HardwareConfig(_Base):
"""Perfil de la placa. Es la única capa que cambia al migrar de GPU.
La 2060 (Turing, sm_75) no soporta bfloat16 ni FlashAttention-2; la 5090
(Blackwell, sm_120) soporta ambos. Todo eso vive acá, no en el código.
"""
name: str
device: Literal["cuda", "cpu"] = "cuda"
dtype: DType = "bfloat16"
use_grad_scaler: bool = False
attention_backend: AttentionBackend = "auto"
compile: bool = True
matmul_precision: Literal["highest", "high", "medium"] = "high"
# El tamaño de lote es una propiedad de la placa, no del modelo: 12 GB y
# 32 GB no admiten lo mismo. El lote efectivo es micro_batch * grad_accum.
micro_batch_size: int = Field(gt=0)
grad_accum_steps: int = Field(gt=0)
# Capability mínima requerida, expresada como (major, minor). Se verifica
# contra la GPU real en train/backends.py antes de empezar.
min_compute_capability: tuple[int, int] | None = None
# Pico teórico de la placa en el dtype de entrenamiento, solo para calcular
# la métrica de MFU. Es un valor aproximado y de referencia: sirve para
# comparar corridas entre sí y detectar regresiones de rendimiento, no para
# afirmar nada sobre el hardware.
peak_tflops: float | None = Field(default=None, gt=0.0)
@property
def effective_batch_size(self) -> int:
return self.micro_batch_size * self.grad_accum_steps
@model_validator(mode="after")
def _check_dtype_scaler(self) -> HardwareConfig:
# float16 sin GradScaler diverge; bfloat16 con GradScaler es un
# sinsentido (bf16 tiene el mismo rango que fp32).
if self.dtype == "float16" and not self.use_grad_scaler:
raise ValueError(
f"perfil '{self.name}': dtype=float16 exige use_grad_scaler=true. "
"Sin escalado de gradiente el entrenamiento en fp16 diverge."
)
if self.dtype != "float16" and self.use_grad_scaler:
raise ValueError(
f"perfil '{self.name}': use_grad_scaler solo aplica a dtype=float16 "
f"(este perfil usa {self.dtype})."
)
return self
class ModelConfig(_Base):
"""Arquitectura. La misma clase Transformer sirve para 50M y para 350M."""
name: str
vocab_size: int = Field(gt=0)
n_layer: int = Field(gt=0)
n_head: int = Field(gt=0)
n_kv_head: int = Field(gt=0)
d_model: int = Field(gt=0)
seq_len: int = Field(gt=0)
# SwiGLU tiene tres matrices en vez de dos, así que el multiplicador
# habitual es 8/3 en vez de 4 para conservar el mismo número de parámetros.
ffn_mult: float = 8 / 3
ffn_multiple_of: int = 64
rope_theta: float = 10_000.0
norm_eps: float = 1e-5
tie_embeddings: bool = True
qk_norm: bool = True
z_loss_weight: float = Field(default=1e-4, ge=0.0)
dropout: float = Field(default=0.0, ge=0.0, lt=1.0)
init_std: float = 0.02
@property
def head_dim(self) -> int:
return self.d_model // self.n_head
@property
def ffn_dim(self) -> int:
"""Dimensión oculta del FFN, redondeada a un múltiplo eficiente."""
raw = int(self.d_model * self.ffn_mult)
m = self.ffn_multiple_of
return ((raw + m - 1) // m) * m
@model_validator(mode="after")
def _check_shapes(self) -> ModelConfig:
if self.d_model % self.n_head != 0:
raise ValueError(
f"modelo '{self.name}': d_model={self.d_model} no es divisible "
f"por n_head={self.n_head}."
)
if self.n_head % self.n_kv_head != 0:
raise ValueError(
f"modelo '{self.name}': n_head={self.n_head} no es divisible por "
f"n_kv_head={self.n_kv_head} (GQA exige que cada grupo de queries "
"comparta exactamente una cabeza de keys/values)."
)
if self.head_dim % 2 != 0:
raise ValueError(
f"modelo '{self.name}': head_dim={self.head_dim} debe ser par "
"(RoPE rota los canales de a pares)."
)
return self
class OptimizerConfig(_Base):
lr: float = Field(gt=0.0)
beta1: float = Field(default=0.9, gt=0.0, lt=1.0)
beta2: float = Field(default=0.95, gt=0.0, lt=1.0)
eps: float = 1e-8
weight_decay: float = Field(default=0.1, ge=0.0)
grad_clip: float = Field(default=1.0, ge=0.0)
class ScheduleConfig(_Base):
"""WarmupStableDecay.
Se elige sobre cosine porque cosine obliga a fijar el total de pasos por
adelantado: extender una corrida o ramificar un checkpoint deja de ser
posible. Con WSD el LR se mantiene plano y solo decae al final, que es lo
que hace baratas las consolidaciones periódicas (Etapa 7b del plan).
"""
kind: Literal["wsd"] = "wsd"
warmup_steps: int = Field(ge=0)
decay_steps: int = Field(ge=0)
min_lr_ratio: float = Field(default=0.0, ge=0.0, le=1.0)
class DataConfig(_Base):
"""De dónde salen los tokens. `source` decide qué loader se usa."""
source: Literal["shards", "chars"] = "shards"
# source="shards": directorio con los .bin uint16 + index.json
shards_dir: str | None = None
# source="chars": un archivo de texto plano, para el smoke test
text_path: str | None = None
val_fraction: float = Field(default=0.005, gt=0.0, lt=0.5)
@model_validator(mode="after")
def _check_source(self) -> DataConfig:
required = {"shards": "shards_dir", "chars": "text_path"}[self.source]
if getattr(self, required) is None:
raise ValueError(f"data.source='{self.source}' exige data.{required}.")
return self
class TrainConfig(_Base):
run_name: str
out_dir: str = "runs"
seed: int = 1337
max_steps: int = Field(gt=0)
optimizer: OptimizerConfig
schedule: ScheduleConfig
log_every: int = Field(default=10, gt=0)
eval_every: int = Field(default=500, gt=0)
eval_batches: int = Field(default=20, gt=0)
checkpoint_every: int = Field(default=1000, gt=0)
sample_every: int = Field(default=0, ge=0) # 0 = no generar muestras
@model_validator(mode="after")
def _check_schedule_fits(self) -> TrainConfig:
s = self.schedule
if s.warmup_steps + s.decay_steps > self.max_steps:
raise ValueError(
f"schedule: warmup({s.warmup_steps}) + decay({s.decay_steps}) = "
f"{s.warmup_steps + s.decay_steps} supera max_steps={self.max_steps}; "
"no quedaría fase estable."
)
return self
class Config(_Base):
"""Config raíz: la composición de todas las capas."""
hardware: HardwareConfig
model: ModelConfig
train: TrainConfig
data: DataConfig
@model_validator(mode="after")
def _check_cross_layer(self) -> Config:
if self.data.source == "chars" and self.model.vocab_size > 1024:
raise ValueError(
f"data.source='chars' con model.vocab_size={self.model.vocab_size}: "
"el smoke test a nivel de caracteres usa un vocabulario chico "
"derivado del texto. Usá un modelo pensado para eso."
)
return self
+5
View File
@@ -0,0 +1,5 @@
"""Arquitectura del modelo. Una sola clase para todas las escalas."""
from enlace.model.transformer import RMSNorm, SwiGLU, Transformer, build_rope_cache
__all__ = ["RMSNorm", "SwiGLU", "Transformer", "build_rope_cache"]
+48
View File
@@ -0,0 +1,48 @@
"""Selección del backend de atención.
Es el único punto del modelo que conoce diferencias entre placas. FlashAttention-2
exige Ampere (sm_80) o superior: en la RTX 2060 (Turing, sm_75) hay que usar el
backend mem_efficient, que también es O(n) en memoria pero más lento. El perfil de
hardware decide; el modelo no sabe en qué placa corre.
"""
from __future__ import annotations
from contextlib import contextmanager
from typing import Iterator
import torch
from torch.nn.attention import SDPBackend, sdpa_kernel
_BACKENDS: dict[str, list[SDPBackend]] = {
# "auto" deja elegir a PyTorch: prueba flash, después mem_efficient, después
# la implementación matemática. Es lo correcto salvo que se quiera forzar
# una ruta concreta para medir o para evitar un kernel con bugs.
"auto": [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH],
"flash": [SDPBackend.FLASH_ATTENTION],
"mem_efficient": [SDPBackend.EFFICIENT_ATTENTION],
"math": [SDPBackend.MATH],
}
@contextmanager
def attention_backend(name: str) -> Iterator[None]:
"""Fija el backend de SDPA dentro del bloque.
En CPU no hay backends alternativos que elegir, así que es un no-op: forzar
uno ahí solo produce advertencias inútiles durante los tests.
"""
if not torch.cuda.is_available():
yield
return
try:
backends = _BACKENDS[name]
except KeyError:
raise ValueError(
f"backend de atención desconocido: {name!r}. "
f"Válidos: {sorted(_BACKENDS)}."
) from None
with sdpa_kernel(backends):
yield
+249
View File
@@ -0,0 +1,249 @@
"""Transformer decoder-only estilo Llama.
Una sola clase para todas las escalas del proyecto: la config decide si es el
modelo de 50M de la 2060 o el de 350M de la 5090. Las piezas son las que hoy
son estándar y por buenas razones:
- RMSNorm pre-norm: más barato que LayerNorm y más estable en profundidad.
- SwiGLU: mejor calidad por parámetro que un MLP con GELU.
- RoPE: posiciones relativas sin parámetros, y extensible después sin
reentrenar (importante para ampliar el contexto cuando la memoria crezca).
- GQA: reduce la caché KV en inferencia, que es el cuello de botella real
de un asistente que responde en tiempo real.
- QK-norm y z-loss: estabilizan el entrenamiento en float16, que en la RTX
2060 no es opcional porque Turing no soporta bfloat16.
"""
from __future__ import annotations
import math
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from enlace.config.schema import ModelConfig
from enlace.model.attention import attention_backend
class RMSNorm(nn.Module):
"""Normalización por raíz cuadrática media, sin término de sesgo."""
def __init__(self, dim: int, eps: float) -> None:
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
# El cálculo se hace en float32 aunque el tensor venga en fp16: la suma
# de cuadrados sobre d_model desborda el rango de fp16 con facilidad.
dtype = x.dtype
x32 = x.float()
normed = x32 * torch.rsqrt(x32.pow(2).mean(-1, keepdim=True) + self.eps)
return (normed * self.weight.float()).to(dtype)
def build_rope_cache(seq_len: int, head_dim: int, theta: float) -> tuple[Tensor, Tensor]:
"""Precalcula cos y sin de RoPE. Devuelve (seq_len, head_dim)."""
inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
pos = torch.arange(seq_len, dtype=torch.float32)
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim/2)
# Se duplica para la convención de mitades: [x1..xd/2, xd/2+1..xd].
emb = torch.cat((freqs, freqs), dim=-1)
return emb.cos(), emb.sin()
def _rotate_half(x: Tensor) -> Tensor:
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat((-x2, x1), dim=-1)
def apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
"""Aplica RoPE a (batch, n_head, seq, head_dim)."""
cos = cos[None, None, :, :].to(x.dtype)
sin = sin[None, None, :, :].to(x.dtype)
return x * cos + _rotate_half(x) * sin
class Attention(nn.Module):
"""Atención causal con GQA y QK-norm."""
def __init__(self, cfg: ModelConfig, backend: str) -> None:
super().__init__()
self.n_head = cfg.n_head
self.n_kv_head = cfg.n_kv_head
self.head_dim = cfg.head_dim
self.n_rep = cfg.n_head // cfg.n_kv_head
self.backend = backend
self.dropout_p = cfg.dropout
self.wq = nn.Linear(cfg.d_model, cfg.n_head * cfg.head_dim, bias=False)
self.wk = nn.Linear(cfg.d_model, cfg.n_kv_head * cfg.head_dim, bias=False)
self.wv = nn.Linear(cfg.d_model, cfg.n_kv_head * cfg.head_dim, bias=False)
self.wo = nn.Linear(cfg.n_head * cfg.head_dim, cfg.d_model, bias=False)
# QK-norm: normalizar queries y keys antes de RoPE acota el crecimiento
# de los logits de atención, que es la causa habitual de los picos de
# loss en fp16. Cuesta casi nada y evita perder días de entrenamiento.
if cfg.qk_norm:
self.q_norm: nn.Module = RMSNorm(cfg.head_dim, cfg.norm_eps)
self.k_norm: nn.Module = RMSNorm(cfg.head_dim, cfg.norm_eps)
else:
self.q_norm = nn.Identity()
self.k_norm = nn.Identity()
def forward(self, x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
b, t, _ = x.shape
q = self.wq(x).view(b, t, self.n_head, self.head_dim).transpose(1, 2)
k = self.wk(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
v = self.wv(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
q = apply_rope(self.q_norm(q), cos, sin)
k = apply_rope(self.k_norm(k), cos, sin)
# GQA: cada cabeza de key/value se comparte entre n_rep cabezas de query.
if self.n_rep > 1:
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
with attention_backend(self.backend):
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout_p if self.training else 0.0
)
out = out.transpose(1, 2).contiguous().view(b, t, -1)
return self.wo(out)
class SwiGLU(nn.Module):
"""FFN con compuerta SiLU: w2(silu(w1(x)) * w3(x))."""
def __init__(self, cfg: ModelConfig) -> None:
super().__init__()
hidden = cfg.ffn_dim
self.w1 = nn.Linear(cfg.d_model, hidden, bias=False)
self.w3 = nn.Linear(cfg.d_model, hidden, bias=False)
self.w2 = nn.Linear(hidden, cfg.d_model, bias=False)
def forward(self, x: Tensor) -> Tensor:
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class Block(nn.Module):
def __init__(self, cfg: ModelConfig, backend: str) -> None:
super().__init__()
self.attn_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.attn = Attention(cfg, backend)
self.ffn_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.ffn = SwiGLU(cfg)
def forward(self, x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
x = x + self.attn(self.attn_norm(x), cos, sin)
return x + self.ffn(self.ffn_norm(x))
class Transformer(nn.Module):
def __init__(self, cfg: ModelConfig, attention_backend: str = "auto") -> None:
super().__init__()
self.cfg = cfg
self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
self.drop = nn.Dropout(cfg.dropout)
self.blocks = nn.ModuleList(Block(cfg, attention_backend) for _ in range(cfg.n_layer))
self.norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
# Con vocab 32k y d_model 512 la tabla de embeddings son 16.7M
# parámetros: atarla con la salida ahorra un tercio de un modelo de 50M.
if cfg.tie_embeddings:
self.lm_head.weight = self.tok_emb.weight
cos, sin = build_rope_cache(cfg.seq_len, cfg.head_dim, cfg.rope_theta)
self.register_buffer("rope_cos", cos, persistent=False)
self.register_buffer("rope_sin", sin, persistent=False)
self.apply(self._init_weights)
# Las proyecciones residuales se inicializan más chicas: sin esto la
# varianza de la corriente residual crece con la profundidad.
scale = 1.0 / math.sqrt(2 * cfg.n_layer)
for name, p in self.named_parameters():
if name.endswith(("attn.wo.weight", "ffn.w2.weight")):
torch.nn.init.normal_(p, mean=0.0, std=cfg.init_std * scale)
def _init_weights(self, module: nn.Module) -> None:
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=self.cfg.init_std)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(module.weight, mean=0.0, std=self.cfg.init_std)
def num_parameters(self, non_embedding: bool = False) -> int:
total = sum(p.numel() for p in self.parameters())
if non_embedding:
total -= self.tok_emb.weight.numel()
return total
def forward(
self, idx: Tensor, targets: Tensor | None = None
) -> tuple[Tensor, Tensor | None]:
"""Devuelve (logits, loss). `loss` es None si no hay targets."""
_, t = idx.shape
if t > self.cfg.seq_len:
raise ValueError(
f"secuencia de largo {t} supera seq_len={self.cfg.seq_len}. "
"Recortá la entrada o ampliá el contexto por RoPE scaling."
)
cos = self.rope_cos[:t]
sin = self.rope_sin[:t]
x = self.drop(self.tok_emb(idx))
for block in self.blocks:
x = block(x, cos, sin)
x = self.norm(x)
logits = self.lm_head(x)
if targets is None:
return logits, None
# La entropía cruzada se calcula en float32: en fp16 el logsumexp sobre
# 32k clases pierde precisión y el gradiente se degrada.
flat_logits = logits.float().view(-1, logits.size(-1))
flat_targets = targets.reshape(-1)
loss = F.cross_entropy(flat_logits, flat_targets, ignore_index=-100)
# z-loss: penaliza que los logits crezcan en magnitud absoluta. Es lo
# que impide que el softmax sature y desborde entrenando en fp16.
if self.cfg.z_loss_weight > 0.0:
valid = flat_targets != -100
if valid.any():
z = torch.logsumexp(flat_logits[valid], dim=-1)
loss = loss + self.cfg.z_loss_weight * z.pow(2).mean()
return logits, loss
@torch.no_grad()
def generate(
self,
idx: Tensor,
max_new_tokens: int,
temperature: float = 1.0,
top_k: int | None = None,
) -> Tensor:
"""Generación sin caché KV: sirve para las muestras de control durante
el entrenamiento. La inferencia real usa la caché (enlace/serve/)."""
self.eval()
for _ in range(max_new_tokens):
window = idx[:, -self.cfg.seq_len :]
logits, _ = self(window)
logits = logits[:, -1, :] / max(temperature, 1e-5)
if top_k is not None:
k = min(top_k, logits.size(-1))
threshold = torch.topk(logits, k, dim=-1).values[:, -1:]
logits = logits.masked_fill(logits < threshold, float("-inf"))
probs = F.softmax(logits, dim=-1)
idx = torch.cat((idx, torch.multinomial(probs, num_samples=1)), dim=1)
return idx
+1
View File
@@ -0,0 +1 @@
"""Entrenamiento: bucle, schedules, checkpointing y detección de hardware."""
+100
View File
@@ -0,0 +1,100 @@
"""Todo lo que depende de la placa concreta vive acá.
Es el único archivo del entrenamiento que conoce diferencias entre GPUs. Cuando
llegue la RTX 5090, este archivo y `configs/hardware/` son lo único que debería
necesitar revisión: si hace falta tocar algo más, es una fuga de abstracción y
conviene arreglarla acá en vez de propagarla.
"""
from __future__ import annotations
import torch
from enlace.config.schema import HardwareConfig
_DTYPES: dict[str, torch.dtype] = {
"float32": torch.float32,
"bfloat16": torch.bfloat16,
"float16": torch.float16,
}
class HardwareError(RuntimeError):
"""El hardware real no coincide con lo que declara el perfil."""
def resolve_dtype(name: str) -> torch.dtype:
return _DTYPES[name]
def describe_device(cfg: HardwareConfig) -> str:
if cfg.device == "cpu":
return "cpu"
props = torch.cuda.get_device_properties(0)
cap = torch.cuda.get_device_capability(0)
return (
f"{props.name} (sm_{cap[0]}{cap[1]}, "
f"{props.total_memory / 1024**3:.1f} GB, torch {torch.__version__})"
)
def setup_device(cfg: HardwareConfig) -> torch.device:
"""Valida que el perfil corresponda a la placa real y prepara el device.
Falla temprano y explícito: descubrir a mitad de una corrida de dos días que
el perfil era el de otra placa es caro.
"""
if cfg.device == "cpu":
torch.set_float32_matmul_precision(cfg.matmul_precision)
return torch.device("cpu")
if not torch.cuda.is_available():
raise HardwareError(
f"el perfil '{cfg.name}' pide device=cuda pero PyTorch no ve ninguna GPU. "
"Para trabajar en una máquina sin GPU usá configs/hardware/cpu.yaml."
)
cap = torch.cuda.get_device_capability(0)
if cfg.min_compute_capability is not None:
required = tuple(cfg.min_compute_capability)
if cap < required:
raise HardwareError(
f"el perfil '{cfg.name}' exige compute capability "
f"sm_{required[0]}{required[1]} o superior, pero la GPU es "
f"sm_{cap[0]}{cap[1]}."
)
# Turing (sm_75) no tiene unidades bfloat16: PyTorch lo emula y el
# entrenamiento se vuelve inservible en vez de fallar. Mejor fallar.
if cfg.dtype == "bfloat16" and not torch.cuda.is_bf16_supported():
raise HardwareError(
f"el perfil '{cfg.name}' pide bfloat16, que esta GPU (sm_{cap[0]}{cap[1]}) "
"no soporta de forma nativa. Usá float16 con use_grad_scaler=true."
)
# FlashAttention-2 requiere Ampere o superior.
if cfg.attention_backend == "flash" and cap < (8, 0):
raise HardwareError(
f"el perfil '{cfg.name}' pide attention_backend=flash, que exige "
f"sm_80 o superior; esta GPU es sm_{cap[0]}{cap[1]}. "
"Usá mem_efficient."
)
torch.set_float32_matmul_precision(cfg.matmul_precision)
return torch.device("cuda")
def autocast_context(cfg: HardwareConfig, device: torch.device):
"""Contexto de precisión mixta acorde al perfil."""
if cfg.dtype == "float32":
return torch.autocast(device_type=device.type, enabled=False)
return torch.autocast(device_type=device.type, dtype=resolve_dtype(cfg.dtype))
def build_grad_scaler(cfg: HardwareConfig) -> torch.amp.GradScaler:
"""GradScaler solo tiene sentido en float16.
En bfloat16 el rango del exponente es el de float32, así que no hay nada que
escalar; el schema ya prohíbe la combinación, esto es la contraparte.
"""
return torch.amp.GradScaler(enabled=cfg.use_grad_scaler)
+131
View File
@@ -0,0 +1,131 @@
"""Checkpointing con reanudación exacta.
"Exacta" significa que reanudar en el paso N y seguir hasta N+K produce los
mismos pesos que una corrida ininterrumpida hasta N+K. Para eso no alcanza con
guardar el modelo: hacen falta el optimizer, el escalador de gradiente, la
posición en el stream de datos y el estado de los generadores aleatorios.
Se verifica a propósito y temprano (tests/test_resume.py). Un checkpoint que no
reanuda exacto es peor que no tener checkpoint, porque el daño es silencioso.
"""
from __future__ import annotations
import random
from pathlib import Path
from typing import Any
import numpy as np
import torch
from enlace.config.load import to_yaml
from enlace.config.schema import Config
CHECKPOINT_FORMAT = 1
def _rng_state() -> dict[str, Any]:
state: dict[str, Any] = {
"python": random.getstate(),
"numpy": np.random.get_state(),
"torch": torch.get_rng_state(),
}
if torch.cuda.is_available():
state["cuda"] = torch.cuda.get_rng_state_all()
return state
def _restore_rng(state: dict[str, Any]) -> None:
random.setstate(state["python"])
np.random.set_state(state["numpy"])
torch.set_rng_state(state["torch"])
if "cuda" in state and torch.cuda.is_available():
torch.cuda.set_rng_state_all(state["cuda"])
def save(
path: str | Path,
*,
config: Config,
step: int,
model: torch.nn.Module,
optimizer: torch.optim.Optimizer,
scaler: torch.amp.GradScaler,
stream_state: dict[str, Any],
metrics: dict[str, float],
) -> Path:
"""Escribe el checkpoint de forma atómica.
Se escribe a un temporal y se renombra: si el proceso muere a mitad de la
escritura, el checkpoint anterior sigue intacto en vez de quedar truncado.
"""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
# torch.compile envuelve el modelo y prefija las claves con "_orig_mod.";
# se guarda el módulo original para que el checkpoint no dependa de si la
# corrida usó compilación o no.
module = getattr(model, "_orig_mod", model)
payload = {
"format": CHECKPOINT_FORMAT,
"step": step,
"model": module.state_dict(),
"optimizer": optimizer.state_dict(),
"scaler": scaler.state_dict(),
"stream": stream_state,
"rng": _rng_state(),
"metrics": metrics,
# Un checkpoint sin su config no es reproducible.
"config": config.model_dump(mode="json"),
}
tmp = path.with_suffix(path.suffix + ".tmp")
torch.save(payload, tmp)
tmp.replace(path)
# Copia legible al lado, para poder inspeccionar la corrida sin abrir el .pt
path.with_suffix(".yaml").write_text(to_yaml(config))
return path
def load(
path: str | Path,
*,
model: torch.nn.Module,
optimizer: torch.optim.Optimizer | None = None,
scaler: torch.amp.GradScaler | None = None,
map_location: str | torch.device = "cpu",
) -> dict[str, Any]:
"""Restaura el estado y devuelve los metadatos (step, stream, metrics)."""
path = Path(path)
payload = torch.load(path, map_location=map_location, weights_only=False)
fmt = payload.get("format")
if fmt != CHECKPOINT_FORMAT:
raise ValueError(
f"{path}: formato de checkpoint {fmt}, se esperaba {CHECKPOINT_FORMAT}."
)
module = getattr(model, "_orig_mod", model)
module.load_state_dict(payload["model"])
if optimizer is not None:
optimizer.load_state_dict(payload["optimizer"])
if scaler is not None:
scaler.load_state_dict(payload["scaler"])
if "rng" in payload:
_restore_rng(payload["rng"])
return {
"step": payload["step"],
"stream": payload.get("stream", {}),
"metrics": payload.get("metrics", {}),
"config": payload.get("config", {}),
}
def latest(run_dir: str | Path) -> Path | None:
"""Devuelve el checkpoint más reciente de una corrida, si existe."""
candidates = sorted(Path(run_dir).glob("ckpt-*.pt"))
return candidates[-1] if candidates else None
+59
View File
@@ -0,0 +1,59 @@
"""Schedule de learning rate.
WSD (WarmupStableDecay) en vez de cosine, y la razón es arquitectónica: con
cosine hay que fijar el total de pasos por adelantado, así que extender una
corrida o ramificar un checkpoint deja de ser posible sin romper el schedule.
WSD sube el LR, lo mantiene plano la mayor parte de la corrida, y solo lo decae
al final. Eso habilita tres cosas que este proyecto necesita:
- extender un pretraining que quedó corto, sin recalcular nada;
- el annealing con el corpus de estilo, que ocurre durante la fase de decay;
- las consolidaciones periódicas, que son una fase de decay corta desde el
snapshot activo en vez de un reentrenamiento completo.
Advertencia práctica: el inicio del decay se ubica en `max_steps - decay_steps`,
así que cambiar `max_steps` mueve esa frontera. Para extender una corrida hay
que ramificar desde un checkpoint anterior al inicio del decay; reanudar con
otro `max_steps` sobre un checkpoint ya decaído da un LR incoherente con lo que
el modelo venía viendo. Reanudar una corrida interrumpida exige la config
idéntica (lo verifica tests/test_resume.py).
"""
from __future__ import annotations
import math
from enlace.config.schema import ScheduleConfig
def wsd_lr(step: int, base_lr: float, cfg: ScheduleConfig, max_steps: int) -> float:
"""LR para `step` (0-indexado).
El decaimiento usa 1 - sqrt(progreso), que empíricamente rinde mejor que el
lineal: mantiene el LR alto más tiempo y después cae rápido.
"""
if step < cfg.warmup_steps:
# Warmup lineal. El +1 evita un primer paso con LR exactamente cero,
# que desperdicia una actualización.
return base_lr * (step + 1) / cfg.warmup_steps
decay_start = max_steps - cfg.decay_steps
if step < decay_start or cfg.decay_steps == 0:
return base_lr
progress = (step - decay_start) / cfg.decay_steps
progress = min(max(progress, 0.0), 1.0)
factor = 1.0 - math.sqrt(progress)
return base_lr * (cfg.min_lr_ratio + (1.0 - cfg.min_lr_ratio) * factor)
def build_lr_fn(cfg: ScheduleConfig, base_lr: float, max_steps: int):
"""Devuelve una función step -> lr, para usar en el bucle de entrenamiento."""
if cfg.kind != "wsd":
raise ValueError(f"schedule desconocido: {cfg.kind}")
def lr_at(step: int) -> float:
return wsd_lr(step, base_lr, cfg, max_steps)
return lr_at
+276
View File
@@ -0,0 +1,276 @@
"""Bucle de entrenamiento.
No hay un solo hiperparámetro en este archivo: todo viene de la config. Lo que
sí vive acá es la mecánica que tiene que ser correcta pase lo que pase —
acumulación de gradiente, precisión mixta, recorte de norma, checkpointing
atómico y reanudación exacta.
python -m enlace.train.train configs/runs/smoke-2060.yaml
python -m enlace.train.train configs/runs/pretrain-2060.yaml --resume
"""
from __future__ import annotations
import json
import math
import random
import sys
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from enlace.config.load import ConfigError, load_config
from enlace.config.schema import Config
from enlace.data.loaders import build_stream
from enlace.model.transformer import Transformer
from enlace.train import checkpoint
from enlace.train.backends import (
autocast_context,
build_grad_scaler,
describe_device,
setup_device,
)
from enlace.train.schedules import build_lr_fn
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def build_optimizer(model: torch.nn.Module, cfg: Config) -> torch.optim.AdamW:
"""AdamW con weight decay solo sobre las matrices.
Aplicar decay a normas y sesgos (tensores de una dimensión) degrada la
calidad sin ahorrar nada: son pocos parámetros y penalizarlos solo distorsiona
la normalización.
"""
decay, no_decay = [], []
for param in model.parameters():
if not param.requires_grad:
continue
(decay if param.dim() >= 2 else no_decay).append(param)
opt = cfg.train.optimizer
return torch.optim.AdamW(
[
{"params": decay, "weight_decay": opt.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
],
lr=opt.lr,
betas=(opt.beta1, opt.beta2),
eps=opt.eps,
)
def flops_per_token(model: Transformer) -> float:
"""Estimación estándar de FLOPs por token en forward+backward.
6*N por las multiplicaciones de matrices y 12*L*H*Q*T por la atención, que
no depende del número de parámetros sino del largo de contexto.
"""
cfg = model.cfg
n = model.num_parameters(non_embedding=True)
attn = 12 * cfg.n_layer * cfg.n_head * cfg.head_dim * cfg.seq_len
return 6 * n + attn
@torch.no_grad()
def evaluate(model: torch.nn.Module, stream, cfg: Config, device: torch.device) -> float:
"""Loss promedio en el split de validación."""
model.eval()
total = 0.0
for _ in range(cfg.train.eval_batches):
x, y = stream.next_batch("val")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
total += loss.item()
model.train()
return total / cfg.train.eval_batches
@torch.no_grad()
def sample(model: torch.nn.Module, stream, device: torch.device, n_tokens: int = 160) -> str:
"""Genera una muestra corta para inspección humana.
Leer estas muestras en cada checkpoint es la forma más rápida de detectar
que algo se rompió: las curvas de loss pueden verse bien mientras el modelo
produce basura.
"""
module = getattr(model, "_orig_mod", model)
start = torch.zeros((1, 1), dtype=torch.long, device=device)
out = module.generate(start, max_new_tokens=n_tokens, temperature=0.8, top_k=50)
model.train()
# Se descarta el token semilla: no es salida del modelo y ensucia la lectura.
generado = out[0, 1:].tolist()
decode = getattr(stream, "decode", None)
if decode is None:
return f"<sin decodificador> {generado[:16]}"
return decode(generado)
def train(cfg: Config, resume: bool = False) -> Path:
run_dir = Path(cfg.train.out_dir) / cfg.train.run_name
run_dir.mkdir(parents=True, exist_ok=True)
log_path = run_dir / "metrics.jsonl"
seed_everything(cfg.train.seed)
device = setup_device(cfg.hardware)
print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}")
print(f"[enlace] perfil: {cfg.hardware.name} | dtype {cfg.hardware.dtype}")
stream = build_stream(cfg, device)
if stream.vocab_size != cfg.model.vocab_size:
raise ConfigError(
f"el corpus tiene vocab_size={stream.vocab_size} pero el modelo "
f"declara {cfg.model.vocab_size}. Son incompatibles."
)
model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device)
n_params = model.num_parameters()
print(f"[enlace] modelo {cfg.model.name}: {n_params:,} parámetros")
optimizer = build_optimizer(model, cfg)
scaler = build_grad_scaler(cfg.hardware)
lr_at = build_lr_fn(cfg.train.schedule, cfg.train.optimizer.lr, cfg.train.max_steps)
start_step = 0
if resume:
ckpt_path = checkpoint.latest(run_dir)
if ckpt_path is None:
print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero")
else:
meta = checkpoint.load(
ckpt_path, model=model, optimizer=optimizer, scaler=scaler, map_location=device
)
stream.load_state_dict(meta["stream"])
start_step = meta["step"]
print(f"[enlace] reanudado desde {ckpt_path.name} en el paso {start_step}")
if cfg.hardware.compile:
print("[enlace] compilando el modelo (la primera iteración tarda)...")
model = torch.compile(model) # type: ignore[assignment]
tokens_per_step = (
cfg.hardware.effective_batch_size * cfg.model.seq_len
)
fpt = flops_per_token(getattr(model, "_orig_mod", model))
model.train()
print(
f"[enlace] {cfg.train.max_steps} pasos x {tokens_per_step:,} tokens "
f"= {cfg.train.max_steps * tokens_per_step / 1e9:.2f}B tokens"
)
t_last = time.perf_counter()
for step in range(start_step, cfg.train.max_steps):
lr = lr_at(step)
for group in optimizer.param_groups:
group["lr"] = lr
optimizer.zero_grad(set_to_none=True)
loss_sum = 0.0
for _ in range(cfg.hardware.grad_accum_steps):
x, y = stream.next_batch("train")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
# Se divide por los micro-pasos para que el gradiente acumulado
# sea el promedio y no la suma: si no, el LR efectivo dependería
# de grad_accum_steps y las curvas no serían comparables entre
# perfiles de hardware.
loss = loss / cfg.hardware.grad_accum_steps
scaler.scale(loss).backward()
loss_sum += loss.item()
grad_norm = float("nan")
if cfg.train.optimizer.grad_clip > 0:
# Hay que deshacer el escalado antes de medir la norma, o el recorte
# se aplicaría sobre gradientes inflados por el GradScaler.
scaler.unscale_(optimizer)
grad_norm = float(
torch.nn.utils.clip_grad_norm_(
model.parameters(), cfg.train.optimizer.grad_clip
)
)
scaler.step(optimizer)
scaler.update()
if (step + 1) % cfg.train.log_every == 0:
if device.type == "cuda":
torch.cuda.synchronize()
now = time.perf_counter()
dt = (now - t_last) / cfg.train.log_every
t_last = now
record: dict[str, Any] = {
"step": step + 1,
"loss": round(loss_sum, 4),
"lr": lr,
"grad_norm": round(grad_norm, 4),
"tokens_per_s": round(tokens_per_step / dt),
"seconds_per_step": round(dt, 4),
}
if cfg.hardware.peak_tflops:
mfu = (fpt * tokens_per_step / dt) / (cfg.hardware.peak_tflops * 1e12)
record["mfu"] = round(mfu, 4)
with log_path.open("a") as fh:
fh.write(json.dumps(record) + "\n")
mfu_txt = f" mfu {record['mfu']:.1%}" if "mfu" in record else ""
print(
f"paso {step + 1:>7} | loss {loss_sum:.4f} | lr {lr:.2e} | "
f"|g| {grad_norm:.2f} | {record['tokens_per_s']:,} tok/s{mfu_txt}"
)
if (step + 1) % cfg.train.eval_every == 0:
val = evaluate(model, stream, cfg, device)
print(f"paso {step + 1:>7} | val_loss {val:.4f} | ppl {math.exp(min(val, 20)):.2f}")
with log_path.open("a") as fh:
fh.write(json.dumps({"step": step + 1, "val_loss": round(val, 4)}) + "\n")
t_last = time.perf_counter()
if cfg.train.sample_every and (step + 1) % cfg.train.sample_every == 0:
text = sample(model, stream, device)
(run_dir / "samples.txt").open("a").write(f"--- paso {step + 1} ---\n{text}\n\n")
print(f"muestra: {text[:120]!r}")
t_last = time.perf_counter()
if (step + 1) % cfg.train.checkpoint_every == 0 or (step + 1) == cfg.train.max_steps:
path = checkpoint.save(
run_dir / f"ckpt-{step + 1:08d}.pt",
config=cfg,
step=step + 1,
model=model,
optimizer=optimizer,
scaler=scaler,
stream_state=stream.state_dict(),
metrics={"loss": loss_sum},
)
print(f"[enlace] checkpoint -> {path.name}")
t_last = time.perf_counter()
return run_dir
def main() -> int:
args = [a for a in sys.argv[1:] if a != "--resume"]
resume = "--resume" in sys.argv[1:]
if not args:
print("uso: python -m enlace.train.train <config.yaml> [--resume] [clave=valor ...]")
return 2
try:
cfg = load_config(args[0], args[1:])
except ConfigError as exc:
print(f"[enlace] {exc}", file=sys.stderr)
return 2
train(cfg, resume=resume)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+45
View File
@@ -0,0 +1,45 @@
[project]
name = "enlace"
version = "0.1.0"
description = "Modelo de lenguaje propio, entrenado desde cero, para asistencia general y familiar en español."
requires-python = ">=3.11"
# Las versiones se pinean con cuidado: el salto a Blackwell (RTX 5090, sm_120)
# exige torch >= 2.7 con CUDA 12.8. Turing (RTX 2060, sm_75) funciona con
# cualquier torch reciente. Ver configs/hardware/ y docs/HARDWARE.md.
dependencies = [
"torch>=2.5",
"pydantic>=2.7",
"omegaconf>=2.3",
"numpy>=1.26",
]
[project.optional-dependencies]
data = [
"datasets>=2.20", # streaming del corpus; nunca se baja entero
"tokenizers>=0.20",
"fasttext-wheel>=0.9", # filtro de idioma
"datasketch>=1.6", # deduplicación MinHash
]
dev = [
"pytest>=8.0",
"ruff>=0.6",
]
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["enlace"]
[tool.pytest.ini_options]
testpaths = ["tests"]
addopts = "-q"
[tool.ruff]
line-length = 100
target-version = "py311"
[tool.ruff.lint]
select = ["E", "F", "I", "UP", "B"]
+33
View File
@@ -0,0 +1,33 @@
#!/usr/bin/env bash
# Descarga un texto en español de dominio público para el smoke test (Etapa 0).
#
# No es corpus de entrenamiento: son unos pocos MB para verificar en minutos que
# CUDA, el entrenador, el checkpointing y la reanudación funcionan antes de
# comprometer días de cómputo.
set -euo pipefail
DESTINO="${1:-data/smoke/texto.txt}"
mkdir -p "$(dirname "$DESTINO")"
if [[ -s "$DESTINO" ]]; then
echo "[enlace] ya existe $DESTINO ($(wc -c <"$DESTINO") bytes); no se baja de nuevo."
exit 0
fi
# Don Quijote — dominio público, español, y suficientemente largo.
URLS=(
"https://www.gutenberg.org/cache/epub/2000/pg2000.txt"
"https://www.gutenberg.org/files/2000/2000-0.txt"
)
for url in "${URLS[@]}"; do
echo "[enlace] bajando $url"
if curl -fsSL --max-time 120 "$url" -o "$DESTINO"; then
echo "[enlace] listo: $DESTINO ($(wc -c <"$DESTINO") bytes)"
exit 0
fi
done
echo "[enlace] no se pudo descargar el texto. Alternativa: copiar cualquier" >&2
echo " archivo .txt en español a $DESTINO (unos pocos MB alcanzan)." >&2
exit 1
+68
View File
@@ -0,0 +1,68 @@
#!/usr/bin/env bash
# Trabajo contra el servidor de entrenamiento por SSH.
#
# El código viaja por git; los datos, las bases y los checkpoints viven en el
# servidor y no se descargan enteros. Solo vuelven artefactos chicos: métricas,
# muestras generadas y logs.
#
# scripts/remote.sh sync # empuja el código al servidor
# scripts/remote.sh run <config.yaml> [args] # entrena bajo tmux
# scripts/remote.sh logs [corrida] # sigue el log en vivo
# scripts/remote.sh pull <corrida> # trae métricas y muestras
# scripts/remote.sh gpu # estado de la GPU
set -euo pipefail
cd "$(dirname "$0")/.."
[[ -f .env ]] && set -a && source .env && set +a
HOST="${ENLACE_REMOTE_HOST:?definí ENLACE_REMOTE_HOST en .env (ver .env.example)}"
DIR="${ENLACE_REMOTE_DIR:?definí ENLACE_REMOTE_DIR en .env}"
PY="${ENLACE_REMOTE_PYTHON:-$DIR/.venv/bin/python}"
comando="${1:-}"
shift || true
case "$comando" in
sync)
git push
ssh "$HOST" "cd '$DIR' && git pull --ff-only && $PY -m pip install -q -e ."
;;
run)
config="${1:?uso: remote.sh run <config.yaml> [overrides...]}"
shift
# El nombre de la sesión sale del nombre de la config, así que dos corridas
# distintas no se pisan y `logs` sabe a cuál conectarse.
sesion="enlace-$(basename "$config" .yaml)"
# tmux es lo que hace que cortar el SSH no mate el entrenamiento.
ssh -t "$HOST" "cd '$DIR' && tmux new-session -d -s '$sesion' \
\"$PY -m enlace.train.train '$config' $* 2>&1 | tee -a 'runs/$sesion.log'\" \
&& echo 'corriendo en tmux: $sesion'"
;;
logs)
sesion="${1:-}"
if [[ -z "$sesion" ]]; then
ssh "$HOST" "cd '$DIR' && ls -t runs/*.log | head -1 | xargs tail -f"
else
ssh "$HOST" "cd '$DIR' && tail -f 'runs/$sesion.log'"
fi
;;
pull)
corrida="${1:?uso: remote.sh pull <corrida>}"
mkdir -p "runs/$corrida"
# Solo métricas y muestras: los checkpoints se quedan en el servidor.
rsync -av --include='metrics.jsonl' --include='samples.txt' --include='*.yaml' \
--exclude='*' "$HOST:$DIR/runs/$corrida/" "runs/$corrida/"
;;
gpu)
ssh "$HOST" "nvidia-smi"
;;
*)
sed -n '2,14p' "$0"
exit 2
;;
esac
+11
View File
@@ -0,0 +1,11 @@
#!/usr/bin/env bash
# Corre la suite de tests.
#
# PYTEST_DISABLE_PLUGIN_AUTOLOAD evita que pytest cargue plugins de otros
# entornos presentes en el PYTHONPATH del sistema (ROS, por ejemplo), que
# rompen la colección de tests con errores de importación ajenos al proyecto.
set -euo pipefail
cd "$(dirname "$0")/.."
PYTHON="${PYTHON:-.venv/bin/python}"
PYTEST_DISABLE_PLUGIN_AUTOLOAD=1 "$PYTHON" -m pytest "$@"
+97
View File
@@ -0,0 +1,97 @@
from __future__ import annotations
import json
from pathlib import Path
import numpy as np
import pytest
import yaml
REPO = Path(__file__).resolve().parents[1]
# Texto en español sintético pero con la estructura correcta (acentos, ñ,
# signos de apertura): suficiente para los tests, y sin depender de la red.
_MUESTRA = (
"ENLACE responde de forma concreta. No adorna. La luz del living está "
"encendida y la temperatura del cuarto es de veintiún grados. ¿Querés que "
"apague la del pasillo? El calendario tiene una entrada mañana temprano. "
"Andrew aprendió despacio, un día a la vez, y así fue construyendo quién "
"era. La memoria no está en los pesos: está en la base de datos. "
)
@pytest.fixture(scope="session")
def texto_es(tmp_path_factory: pytest.TempPathFactory) -> Path:
path = tmp_path_factory.mktemp("datos") / "texto.txt"
path.write_text(_MUESTRA * 200, encoding="utf-8")
return path
@pytest.fixture(scope="session")
def shards_dir(tmp_path_factory: pytest.TempPathFactory) -> Path:
"""Dos shards uint16 con tokens predecibles, para verificar el recorrido."""
directory = tmp_path_factory.mktemp("shards")
total = 0
entries = []
for i, count in enumerate([4096, 4096]):
tokens = np.arange(total, total + count, dtype=np.uint16)
name = f"shard-{i:04d}.bin"
tokens.tofile(directory / name)
entries.append({"file": name, "tokens": count})
total += count
(directory / "index.json").write_text(
json.dumps({"vocab_size": 8192, "shards": entries, "total_tokens": total})
)
return directory
def write_run_config(
tmp_path: Path,
texto: Path,
*,
run_name: str,
max_steps: int,
checkpoint_every: int,
) -> Path:
"""Config de corrida mínima en CPU, apuntando a un texto de prueba."""
cfg = {
"hardware": yaml.safe_load((REPO / "configs/hardware/cpu.yaml").read_text()),
"model": yaml.safe_load((REPO / "configs/model/char-smoke.yaml").read_text()),
"train": {
"run_name": run_name,
"out_dir": str(tmp_path / "runs"),
"seed": 1337,
"max_steps": max_steps,
"optimizer": {
"lr": 3.0e-3,
"beta1": 0.9,
"beta2": 0.95,
"eps": 1.0e-8,
"weight_decay": 0.1,
"grad_clip": 1.0,
},
"schedule": {
"kind": "wsd",
"warmup_steps": 2,
"decay_steps": 2,
"min_lr_ratio": 0.0,
},
"log_every": 1000,
"eval_every": 1000,
"eval_batches": 2,
"checkpoint_every": checkpoint_every,
"sample_every": 0,
},
"data": {
"source": "chars",
"text_path": str(texto),
"val_fraction": 0.05,
},
}
# Modelo aún más chico: los tests tienen que correr en segundos.
cfg["model"].update({"n_layer": 2, "d_model": 64, "n_head": 4, "n_kv_head": 2, "seq_len": 64})
cfg["hardware"].update({"micro_batch_size": 4, "grad_accum_steps": 2})
path = tmp_path / f"{run_name}.yaml"
path.write_text(yaml.safe_dump(cfg))
return path
+126
View File
@@ -0,0 +1,126 @@
"""La config tiene que fallar al arrancar, no a las tres horas."""
from __future__ import annotations
from pathlib import Path
import pytest
import yaml
from enlace.config.load import ConfigError, compose, load_config
REPO = Path(__file__).resolve().parents[1]
def test_composicion_por_capas():
cfg = load_config(REPO / "configs/runs/pretrain-2060.yaml")
assert cfg.hardware.name == "turing-2060"
assert cfg.model.name == "tiny-50m"
assert cfg.train.run_name == "pretrain-tiny-50m"
assert cfg.data.source == "shards"
def test_el_yaml_raiz_pisa_las_capas_incluidas():
cfg = load_config(REPO / "configs/runs/smoke-2060.yaml")
# El perfil de la 2060 declara micro_batch_size 8; la corrida lo sube a 64
# porque el modelo char es diminuto.
assert cfg.hardware.micro_batch_size == 64
assert cfg.hardware.dtype == "float16" # esto sí viene del perfil
def test_overrides_de_linea_de_comandos():
cfg = load_config(REPO / "configs/runs/smoke-cpu.yaml", ["train.seed=99"])
assert cfg.train.seed == 99
@pytest.mark.parametrize("perfil", ["turing-2060", "blackwell-5090", "cpu"])
def test_todos_los_perfiles_de_hardware_son_validos(perfil, tmp_path):
"""Un perfil roto solo se descubre al migrar de placa; mejor ahora."""
raiz = {
"include": {
"hardware": f"hardware/{perfil}.yaml",
"model": "model/tiny-50m.yaml",
"train": "train/pretrain.yaml",
"data": "data/corpus.yaml",
}
}
path = REPO / "configs" / "runs" / f"_tmp_{perfil}.yaml"
path.write_text(yaml.safe_dump(raiz))
try:
cfg = load_config(path)
assert cfg.hardware.name == perfil
finally:
path.unlink()
def _raiz(tmp_path: Path, **parches) -> Path:
"""Config raíz válida con parches encima, para probar validaciones."""
base = {
"include": {
"hardware": "hardware/cpu.yaml",
"model": "model/tiny-50m.yaml",
"train": "train/pretrain.yaml",
"data": "data/corpus.yaml",
}
}
base.update(parches)
path = REPO / "configs" / "runs" / "_tmp_test.yaml"
path.write_text(yaml.safe_dump(base))
return path
def _espera_error(path: Path, fragmento: str):
try:
with pytest.raises(ConfigError) as exc:
load_config(path)
assert fragmento in str(exc.value)
finally:
path.unlink()
def test_rechaza_campos_desconocidos(tmp_path):
# Un typo tiene que ser un error ruidoso, no un default silencioso.
_espera_error(_raiz(tmp_path, model={"n_layers": 12}), "n_layers")
def test_rechaza_fp16_sin_grad_scaler(tmp_path):
_espera_error(
_raiz(tmp_path, hardware={"dtype": "float16", "use_grad_scaler": False}),
"use_grad_scaler",
)
def test_rechaza_grad_scaler_sin_fp16(tmp_path):
_espera_error(
_raiz(tmp_path, hardware={"dtype": "float32", "use_grad_scaler": True}),
"use_grad_scaler",
)
def test_rechaza_dmodel_no_divisible_por_heads(tmp_path):
_espera_error(_raiz(tmp_path, model={"d_model": 500}), "no es divisible")
def test_rechaza_gqa_incoherente(tmp_path):
_espera_error(_raiz(tmp_path, model={"n_head": 8, "n_kv_head": 3}), "n_kv_head")
def test_rechaza_schedule_que_no_entra(tmp_path):
_espera_error(
_raiz(tmp_path, train={"max_steps": 100, "schedule": {"warmup_steps": 80, "decay_steps": 80}}),
"no quedaría fase estable",
)
def test_rechaza_data_source_sin_su_ruta(tmp_path):
_espera_error(_raiz(tmp_path, data={"source": "chars", "text_path": None}), "text_path")
def test_include_con_capa_desconocida(tmp_path):
path = REPO / "configs" / "runs" / "_tmp_bad.yaml"
path.write_text(yaml.safe_dump({"include": {"hardwar": "hardware/cpu.yaml"}}))
try:
with pytest.raises(ConfigError, match="capas desconocidas"):
compose(path)
finally:
path.unlink()
+143
View File
@@ -0,0 +1,143 @@
"""Cargadores de datos: determinismo, reanudación y recorrido de shards."""
from __future__ import annotations
import pytest
import torch
from enlace.data.loaders import ByteStream, ShardStream
def _byte_stream(texto, seed=1337):
return ByteStream(
text_path=texto,
batch_size=2,
seq_len=16,
val_fraction=0.1,
seed=seed,
device=torch.device("cpu"),
)
def test_los_targets_son_la_entrada_desplazada_un_token(texto_es):
stream = _byte_stream(texto_es)
x, y = stream.next_batch("train")
assert x.shape == y.shape == (2, 16)
assert torch.equal(x[:, 1:], y[:, :-1])
def test_la_misma_semilla_da_los_mismos_lotes(texto_es):
a = _byte_stream(texto_es)
b = _byte_stream(texto_es)
for _ in range(3):
xa, _ = a.next_batch("train")
xb, _ = b.next_batch("train")
assert torch.equal(xa, xb)
def test_restaurar_el_estado_continua_la_misma_secuencia(texto_es):
a = _byte_stream(texto_es)
for _ in range(5):
a.next_batch("train")
estado = a.state_dict()
esperado = [a.next_batch("train")[0] for _ in range(3)]
b = _byte_stream(texto_es)
b.load_state_dict(estado)
obtenido = [b.next_batch("train")[0] for _ in range(3)]
for e, o in zip(esperado, obtenido):
assert torch.equal(e, o)
def test_evaluar_no_altera_la_secuencia_de_entrenamiento(texto_es):
"""Generadores separados por split: cambiar eval_every no debe mover los
lotes de entrenamiento, o dos corridas dejarían de ser comparables."""
a = _byte_stream(texto_es)
esperado = [a.next_batch("train")[0] for _ in range(3)]
b = _byte_stream(texto_es)
lotes = []
for _ in range(3):
b.next_batch("val")
lotes.append(b.next_batch("train")[0])
for e, o in zip(esperado, lotes):
assert torch.equal(e, o)
def test_el_vocabulario_de_bytes_es_siempre_256(texto_es):
assert _byte_stream(texto_es).vocab_size == 256
def test_decodifica_utf8_con_acentos():
assert ByteStream.decode(list("El niño está acá.".encode())) == "El niño está acá."
def test_texto_inexistente_da_un_error_util(tmp_path):
with pytest.raises(FileNotFoundError, match="prepare_smoke_data"):
_byte_stream(tmp_path / "no-existe.txt")
def _shard_stream(directory, batch_size=2, seq_len=8):
return ShardStream(
shards_dir=directory,
batch_size=batch_size,
seq_len=seq_len,
val_fraction=0.1,
device=torch.device("cpu"),
)
def test_los_shards_se_recorren_en_orden(shards_dir):
stream = _shard_stream(shards_dir)
x, y = stream.next_batch("train")
# El fixture escribe tokens consecutivos: 0, 1, 2, ...
assert x[0].tolist() == list(range(8))
assert y[0].tolist() == list(range(1, 9))
assert x[1].tolist() == list(range(8, 16))
x2, _ = stream.next_batch("train")
assert x2[0].tolist() == list(range(16, 24))
def test_la_lectura_cruza_el_limite_entre_shards(shards_dir):
"""Un lote que empieza cerca del final de un shard tiene que continuar en el
siguiente sin saltarse ni repetir tokens."""
stream = _shard_stream(shards_dir, batch_size=1, seq_len=8)
stream.load_state_dict({"train": 4090})
x, _ = stream.next_batch("train")
assert x[0].tolist() == list(range(4090, 4098)) # cruza de shard 0 a shard 1
def test_restaurar_la_posicion_reanuda_donde_iba(shards_dir):
a = _shard_stream(shards_dir)
for _ in range(4):
a.next_batch("train")
estado = a.state_dict()
esperado, _ = a.next_batch("train")
b = _shard_stream(shards_dir)
b.load_state_dict(estado)
obtenido, _ = b.next_batch("train")
assert torch.equal(esperado, obtenido)
def test_al_terminar_la_epoca_vuelve_al_principio(shards_dir):
stream = _shard_stream(shards_dir, batch_size=1, seq_len=8)
primero, _ = stream.next_batch("train")
# El split de train son 7372 tokens; posicionarse casi al final.
stream.load_state_dict({"train": 7370})
stream.next_batch("train")
reiniciado, _ = stream.next_batch("train")
assert reiniciado[0, 0].item() == 8 # segundo lote de la época nueva
def test_shards_faltantes_dan_un_error_util(tmp_path):
with pytest.raises(FileNotFoundError, match="prepare_data"):
_shard_stream(tmp_path)
def test_corpus_demasiado_chico_para_el_lote(shards_dir):
with pytest.raises(ValueError, match="menos que los"):
_shard_stream(shards_dir, batch_size=64, seq_len=1024)
+152
View File
@@ -0,0 +1,152 @@
"""Propiedades de la arquitectura que tienen que valer en cualquier escala."""
from __future__ import annotations
import math
import pytest
import torch
from enlace.config.schema import ModelConfig
from enlace.model.transformer import Transformer, apply_rope, build_rope_cache
def _cfg(**parches) -> ModelConfig:
base = dict(
name="test",
vocab_size=256,
n_layer=2,
n_head=4,
n_kv_head=2,
d_model=64,
seq_len=32,
)
base.update(parches)
return ModelConfig(**base)
def test_loss_inicial_es_la_del_azar():
"""Un modelo recién inicializado no sabe nada: su loss es ln(vocab).
Si arranca muy por debajo hay una fuga de información (targets sin
desplazar, por ejemplo); si arranca muy por encima, la inicialización está
mal escalada y la corrida va a tardar en despegar.
"""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (8, cfg.seq_len + 1))
_, loss = model(x[:, :-1], x[:, 1:])
assert loss.item() == pytest.approx(math.log(cfg.vocab_size), abs=0.15)
def test_embeddings_atados_comparten_memoria():
model = Transformer(_cfg(tie_embeddings=True), "math")
assert model.lm_head.weight.data_ptr() == model.tok_emb.weight.data_ptr()
suelto = Transformer(_cfg(tie_embeddings=False), "math")
assert suelto.lm_head.weight.data_ptr() != suelto.tok_emb.weight.data_ptr()
def test_atar_embeddings_ahorra_una_tabla_entera():
atado = Transformer(_cfg(tie_embeddings=True), "math").num_parameters()
suelto = Transformer(_cfg(tie_embeddings=False), "math").num_parameters()
assert suelto - atado == 256 * 64
def test_la_atencion_es_causal():
"""Cambiar un token no puede alterar las predicciones anteriores.
Es la propiedad que hace que el entrenamiento tenga sentido; si se rompe,
el loss baja de forma espectacular y el modelo no sirve para nada.
"""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math").eval()
x = torch.randint(0, cfg.vocab_size, (1, cfg.seq_len))
with torch.no_grad():
base, _ = model(x)
alterado = x.clone()
alterado[0, -1] = (alterado[0, -1] + 1) % cfg.vocab_size
otro, _ = model(alterado)
assert torch.allclose(base[:, :-1], otro[:, :-1], atol=1e-6)
def test_ignore_index_excluye_posiciones_enmascaradas():
"""El SFT entrena solo sobre los turnos del asistente: el resto va con -100."""
torch.manual_seed(0)
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (4, cfg.seq_len + 1))
inp, tgt = x[:, :-1], x[:, 1:].clone()
tgt[:, : cfg.seq_len // 2] = -100
_, loss = model(inp, tgt)
assert torch.isfinite(loss)
def test_todo_enmascarado_no_rompe():
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (2, cfg.seq_len))
tgt = torch.full_like(x, -100)
_, loss = model(x, tgt)
assert torch.isnan(loss) or torch.isfinite(loss) # no debe explotar
def test_rechaza_secuencias_mas_largas_que_el_contexto():
cfg = _cfg()
model = Transformer(cfg, "math")
x = torch.randint(0, cfg.vocab_size, (1, cfg.seq_len + 1))
with pytest.raises(ValueError, match="supera seq_len"):
model(x)
def test_generate_agrega_exactamente_los_tokens_pedidos():
cfg = _cfg()
model = Transformer(cfg, "math")
inicio = torch.zeros((2, 3), dtype=torch.long)
out = model.generate(inicio, max_new_tokens=7, top_k=5)
assert out.shape == (2, 10)
assert torch.equal(out[:, :3], inicio)
def test_gqa_reduce_las_proyecciones_kv():
"""Con GQA 4:1 las matrices de keys/values son un cuarto de las de queries."""
cfg = _cfg(n_head=8, n_kv_head=2, d_model=64)
model = Transformer(cfg, "math")
attn = model.blocks[0].attn
assert attn.wq.out_features == 8 * cfg.head_dim
assert attn.wk.out_features == 2 * cfg.head_dim
assert attn.n_rep == 4
def test_rope_preserva_la_norma():
"""RoPE es una rotación: no cambia la magnitud de los vectores."""
cos, sin = build_rope_cache(16, 8, 10_000.0)
x = torch.randn(2, 3, 16, 8)
y = apply_rope(x, cos, sin)
assert torch.allclose(x.norm(dim=-1), y.norm(dim=-1), atol=1e-5)
def test_z_loss_penaliza_logits_grandes():
torch.manual_seed(0)
cfg_con = _cfg(z_loss_weight=1.0)
cfg_sin = _cfg(z_loss_weight=0.0)
torch.manual_seed(0)
con = Transformer(cfg_con, "math")
torch.manual_seed(0)
sin = Transformer(cfg_sin, "math")
x = torch.randint(0, 256, (4, 33))
_, loss_con = con(x[:, :-1], x[:, 1:])
_, loss_sin = sin(x[:, :-1], x[:, 1:])
assert loss_con.item() > loss_sin.item()
def test_el_modelo_del_plan_pesa_lo_esperado():
"""tiny-50m tiene que estar cerca de 50M: es lo que entra en la 2060."""
from enlace.config.load import load_config
from pathlib import Path
cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml")
model = Transformer(cfg.model, "math")
assert 45e6 < model.num_parameters() < 55e6
+80
View File
@@ -0,0 +1,80 @@
"""La prueba más importante del entrenamiento.
Reanudar desde un checkpoint tiene que producir exactamente los mismos pesos que
una corrida ininterrumpida. Si no, el daño es silencioso: las curvas se ven
bien, el modelo entrena de más sobre datos ya vistos, y nadie se entera hasta
que la corrida de dos días terminó y el resultado no sirve.
"""
from __future__ import annotations
import shutil
from pathlib import Path
import torch
from enlace.config.load import load_config
from enlace.train import checkpoint
from enlace.train.train import train
from tests.conftest import write_run_config
def _final_weights(run_dir, step: int):
payload = torch.load(run_dir / f"ckpt-{step:08d}.pt", map_location="cpu", weights_only=False)
return payload["model"], payload
def test_reanudar_reproduce_la_corrida_completa(tmp_path, texto_es):
# Corrida A: 20 pasos sin interrupción.
cfg_a = load_config(
write_run_config(tmp_path, texto_es, run_name="A", max_steps=20, checkpoint_every=10)
)
dir_a = train(cfg_a)
# Corrida B: la misma config, simulando que el proceso murió en el paso 10.
# Se copia el checkpoint del paso 10 y se reanuda desde ahí.
#
# La config tiene que ser idéntica, incluido max_steps: el schedule WSD
# ubica el inicio del decay en max_steps - decay_steps, así que reanudar
# con otro max_steps cambia el LR de los pasos que faltan. No es un defecto
# del checkpointing, pero sí una trampa fácil de pisar.
cfg_b = load_config(
write_run_config(tmp_path, texto_es, run_name="B", max_steps=20, checkpoint_every=10)
)
dir_b = Path(cfg_b.train.out_dir) / cfg_b.train.run_name
dir_b.mkdir(parents=True, exist_ok=True)
shutil.copy(dir_a / "ckpt-00000010.pt", dir_b / "ckpt-00000010.pt")
train(cfg_b, resume=True)
pesos_a, payload_a = _final_weights(dir_a, 20)
pesos_b, payload_b = _final_weights(dir_b, 20)
assert payload_a["step"] == payload_b["step"] == 20
assert set(pesos_a) == set(pesos_b)
for nombre, tensor_a in pesos_a.items():
assert torch.equal(tensor_a, pesos_b[nombre]), (
f"el parámetro {nombre} difiere tras reanudar: la reanudación no es exacta"
)
def test_el_checkpoint_guarda_la_posicion_del_stream(tmp_path, texto_es):
cfg = load_config(
write_run_config(tmp_path, texto_es, run_name="C", max_steps=4, checkpoint_every=4)
)
run_dir = train(cfg)
payload = torch.load(run_dir / "ckpt-00000004.pt", map_location="cpu", weights_only=False)
# Sin el estado del stream, reanudar volvería al principio del corpus.
assert "stream" in payload and payload["stream"], "el checkpoint no guardó el stream"
assert "train" in payload["stream"]
# Y sin la config, el checkpoint no sería reproducible.
assert payload["config"]["model"]["name"] == "char-smoke"
def test_checkpoint_atomico_no_deja_archivos_truncados(tmp_path, texto_es):
cfg = load_config(
write_run_config(tmp_path, texto_es, run_name="D", max_steps=6, checkpoint_every=6)
)
run_dir = train(cfg)
assert checkpoint.latest(run_dir) is not None
assert not list(run_dir.glob("*.tmp")), "quedaron temporales de escritura"
+56
View File
@@ -0,0 +1,56 @@
"""Forma del schedule WSD."""
from __future__ import annotations
from enlace.config.schema import ScheduleConfig
from enlace.train.schedules import build_lr_fn
def _lr_fn(max_steps=1000, warmup=100, decay=200, min_ratio=0.0, base=1e-3):
cfg = ScheduleConfig(
kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio
)
return build_lr_fn(cfg, base, max_steps)
def test_warmup_sube_linealmente_desde_arriba_de_cero():
lr = _lr_fn()
assert lr(0) > 0.0 # el primer paso no se desperdicia con lr=0
assert lr(0) < lr(50) < lr(99)
assert lr(99) == 1e-3
def test_la_fase_estable_es_plana():
lr = _lr_fn()
assert lr(100) == lr(500) == lr(799) == 1e-3
def test_el_decay_baja_de_forma_monotona_hasta_cero():
lr = _lr_fn()
valores = [lr(s) for s in range(800, 1000)]
assert all(a >= b for a, b in zip(valores, valores[1:]))
assert valores[0] == 1e-3
assert valores[-1] < 1e-4
def test_min_lr_ratio_pone_un_piso():
lr = _lr_fn(min_ratio=0.1)
assert lr(999) >= 1e-4 * 0.99
def test_sin_decay_el_lr_queda_plano_hasta_el_final():
lr = _lr_fn(decay=0)
assert lr(999) == 1e-3
def test_extender_la_corrida_mueve_el_inicio_del_decay():
"""Documenta la trampa: max_steps define dónde empieza a decaer el LR.
Reanudar una corrida interrumpida exige el mismo max_steps; extenderla es
una decisión distinta, que hay que tomar ramificando desde un checkpoint
anterior al decay.
"""
corto = _lr_fn(max_steps=1000)
largo = _lr_fn(max_steps=2000)
assert corto(850) < corto(700) # ya está decayendo
assert largo(850) == largo(700) # todavía en la fase estable