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:
@@ -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
@@ -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/
|
||||||
@@ -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
|
||||||
|
```
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
# Warmup–Stable–Decay: 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
|
||||||
@@ -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
@@ -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 ~100–500M
|
||||||
|
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 ~2–4 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 (Warmup–Stable–Decay), 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 (1–2 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 (2–3 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 (1–2 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 **~20–40M 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 ≈ **1–2 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 (4–6 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 7–8B 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 2–4 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ó.
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
"""ENLACE — modelo de lenguaje propio, entrenado desde cero."""
|
||||||
|
|
||||||
|
__version__ = "0.1.0"
|
||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -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")))
|
||||||
@@ -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):
|
||||||
|
"""Warmup–Stable–Decay.
|
||||||
|
|
||||||
|
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
|
||||||
@@ -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"]
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Entrenamiento: bucle, schedules, checkpointing y detección de hardware."""
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""Schedule de learning rate.
|
||||||
|
|
||||||
|
WSD (Warmup–Stable–Decay) 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
|
||||||
@@ -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())
|
||||||
@@ -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"]
|
||||||
Executable
+33
@@ -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
|
||||||
@@ -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
|
||||||
Executable
+11
@@ -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 "$@"
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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"
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user