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

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

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

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

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

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

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-27 23:04:45 -03:00
commit 90ac2a6582
33 changed files with 2838 additions and 0 deletions
+3
View File
@@ -0,0 +1,3 @@
"""ENLACE — modelo de lenguaje propio, entrenado desde cero."""
__version__ = "0.1.0"
+28
View File
@@ -0,0 +1,28 @@
"""Capa de configuración: nada de números mágicos en el código.
Todo hiperparámetro, ruta y umbral vive en `configs/*.yaml`, se compone por
capas y se valida con pydantic antes de que arranque cualquier proceso largo.
"""
from enlace.config.load import load_config, load_config_from_argv
from enlace.config.schema import (
Config,
DataConfig,
HardwareConfig,
ModelConfig,
OptimizerConfig,
ScheduleConfig,
TrainConfig,
)
__all__ = [
"Config",
"DataConfig",
"HardwareConfig",
"ModelConfig",
"OptimizerConfig",
"ScheduleConfig",
"TrainConfig",
"load_config",
"load_config_from_argv",
]
+117
View File
@@ -0,0 +1,117 @@
"""Composición y carga de configs.
Un YAML raíz declara qué capa usar de cada familia; se fusionan con OmegaConf y
el resultado se valida con pydantic. Los overrides de línea de comandos existen
para probar variantes sin editar archivos, no para esconder configuración.
python -m enlace.train.train configs/runs/smoke.yaml hardware.compile=false
"""
from __future__ import annotations
import sys
from pathlib import Path
from typing import Any
from omegaconf import DictConfig, OmegaConf
from enlace.config.schema import Config
# Familias de capas que puede declarar un YAML raíz, en el orden en que se
# fusionan. El orden solo importa para los mensajes de error.
LAYERS = ("hardware", "model", "train", "data")
class ConfigError(RuntimeError):
"""Error de configuración legible, sin traceback de pydantic encima."""
def _resolve(path: str | Path, root: Path) -> Path:
p = Path(path)
full = p if p.is_absolute() else root / p
if not full.is_file():
raise ConfigError(f"no existe el archivo de config: {full}")
return full
def _load_yaml(path: Path) -> DictConfig:
cfg = OmegaConf.load(path)
if not isinstance(cfg, DictConfig):
raise ConfigError(f"{path}: el YAML raíz debe ser un mapeo, no una lista.")
return cfg
def compose(path: str | Path, overrides: list[str] | None = None) -> dict[str, Any]:
"""Compone un YAML raíz en un dict plano, sin validar todavía.
El YAML raíz tiene la forma:
include:
hardware: hardware/turing-2060.yaml
model: model/tiny-50m.yaml
train: train/pretrain.yaml
data: data/corpus.yaml
# opcional: ajustes puntuales encima de las capas incluidas
train:
run_name: mi-corrida
"""
path = Path(path).resolve()
if not path.is_file():
raise ConfigError(f"no existe el archivo de config: {path}")
raw = _load_yaml(path)
# Las rutas de `include` se resuelven contra el directorio configs/, que es
# el padre del directorio del YAML raíz (configs/runs/foo.yaml -> configs/).
configs_root = path.parent.parent if path.parent.name == "runs" else path.parent
includes = raw.pop("include", None)
merged = OmegaConf.create({})
if includes is not None:
for layer in LAYERS:
ref = includes.get(layer)
if ref is None:
continue
merged[layer] = _load_yaml(_resolve(ref, configs_root))
unknown = set(includes.keys()) - set(LAYERS)
if unknown:
raise ConfigError(
f"{path}: capas desconocidas en `include`: {sorted(unknown)}. "
f"Válidas: {list(LAYERS)}."
)
# Lo que quede en el YAML raíz pisa a las capas incluidas.
merged = OmegaConf.merge(merged, raw)
if overrides:
merged = OmegaConf.merge(merged, OmegaConf.from_dotlist(list(overrides)))
resolved = OmegaConf.to_container(merged, resolve=True)
assert isinstance(resolved, dict)
return resolved
def load_config(path: str | Path, overrides: list[str] | None = None) -> Config:
"""Compone, valida y devuelve la config. Falla temprano y con claridad."""
data = compose(path, overrides)
try:
return Config.model_validate(data)
except Exception as exc: # pydantic.ValidationError y los ValueError propios
raise ConfigError(f"config inválida ({path}):\n{exc}") from None
def load_config_from_argv(argv: list[str] | None = None) -> Config:
"""`prog config.yaml [clave.sub=valor ...]`, para los entrypoints."""
args = list(sys.argv[1:] if argv is None else argv)
if not args:
raise ConfigError("uso: <programa> <config.yaml> [clave.sub=valor ...]")
return load_config(args[0], args[1:])
def to_yaml(config: Config) -> str:
"""Serializa la config resuelta, para guardarla junto al checkpoint.
Un checkpoint sin su config no es reproducible: esto es lo que se escribe
en el snapshot.
"""
return OmegaConf.to_yaml(OmegaConf.create(config.model_dump(mode="json")))
+220
View File
@@ -0,0 +1,220 @@
"""Esquemas de configuración validados.
Una config inválida tiene que fallar al arrancar con un mensaje claro, no a las
tres horas de entrenamiento. Todas las validaciones cruzadas que se pueden hacer
sin tocar la GPU se hacen acá.
"""
from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator
DType = Literal["float32", "bfloat16", "float16"]
AttentionBackend = Literal["auto", "flash", "mem_efficient", "math"]
class _Base(BaseModel):
"""Base común: prohíbe campos desconocidos.
Un typo en un YAML (`n_layers` en vez de `n_layer`) tiene que ser un error
ruidoso, no un valor por defecto aplicado en silencio.
"""
model_config = ConfigDict(extra="forbid", frozen=True)
class HardwareConfig(_Base):
"""Perfil de la placa. Es la única capa que cambia al migrar de GPU.
La 2060 (Turing, sm_75) no soporta bfloat16 ni FlashAttention-2; la 5090
(Blackwell, sm_120) soporta ambos. Todo eso vive acá, no en el código.
"""
name: str
device: Literal["cuda", "cpu"] = "cuda"
dtype: DType = "bfloat16"
use_grad_scaler: bool = False
attention_backend: AttentionBackend = "auto"
compile: bool = True
matmul_precision: Literal["highest", "high", "medium"] = "high"
# El tamaño de lote es una propiedad de la placa, no del modelo: 12 GB y
# 32 GB no admiten lo mismo. El lote efectivo es micro_batch * grad_accum.
micro_batch_size: int = Field(gt=0)
grad_accum_steps: int = Field(gt=0)
# Capability mínima requerida, expresada como (major, minor). Se verifica
# contra la GPU real en train/backends.py antes de empezar.
min_compute_capability: tuple[int, int] | None = None
# Pico teórico de la placa en el dtype de entrenamiento, solo para calcular
# la métrica de MFU. Es un valor aproximado y de referencia: sirve para
# comparar corridas entre sí y detectar regresiones de rendimiento, no para
# afirmar nada sobre el hardware.
peak_tflops: float | None = Field(default=None, gt=0.0)
@property
def effective_batch_size(self) -> int:
return self.micro_batch_size * self.grad_accum_steps
@model_validator(mode="after")
def _check_dtype_scaler(self) -> HardwareConfig:
# float16 sin GradScaler diverge; bfloat16 con GradScaler es un
# sinsentido (bf16 tiene el mismo rango que fp32).
if self.dtype == "float16" and not self.use_grad_scaler:
raise ValueError(
f"perfil '{self.name}': dtype=float16 exige use_grad_scaler=true. "
"Sin escalado de gradiente el entrenamiento en fp16 diverge."
)
if self.dtype != "float16" and self.use_grad_scaler:
raise ValueError(
f"perfil '{self.name}': use_grad_scaler solo aplica a dtype=float16 "
f"(este perfil usa {self.dtype})."
)
return self
class ModelConfig(_Base):
"""Arquitectura. La misma clase Transformer sirve para 50M y para 350M."""
name: str
vocab_size: int = Field(gt=0)
n_layer: int = Field(gt=0)
n_head: int = Field(gt=0)
n_kv_head: int = Field(gt=0)
d_model: int = Field(gt=0)
seq_len: int = Field(gt=0)
# SwiGLU tiene tres matrices en vez de dos, así que el multiplicador
# habitual es 8/3 en vez de 4 para conservar el mismo número de parámetros.
ffn_mult: float = 8 / 3
ffn_multiple_of: int = 64
rope_theta: float = 10_000.0
norm_eps: float = 1e-5
tie_embeddings: bool = True
qk_norm: bool = True
z_loss_weight: float = Field(default=1e-4, ge=0.0)
dropout: float = Field(default=0.0, ge=0.0, lt=1.0)
init_std: float = 0.02
@property
def head_dim(self) -> int:
return self.d_model // self.n_head
@property
def ffn_dim(self) -> int:
"""Dimensión oculta del FFN, redondeada a un múltiplo eficiente."""
raw = int(self.d_model * self.ffn_mult)
m = self.ffn_multiple_of
return ((raw + m - 1) // m) * m
@model_validator(mode="after")
def _check_shapes(self) -> ModelConfig:
if self.d_model % self.n_head != 0:
raise ValueError(
f"modelo '{self.name}': d_model={self.d_model} no es divisible "
f"por n_head={self.n_head}."
)
if self.n_head % self.n_kv_head != 0:
raise ValueError(
f"modelo '{self.name}': n_head={self.n_head} no es divisible por "
f"n_kv_head={self.n_kv_head} (GQA exige que cada grupo de queries "
"comparta exactamente una cabeza de keys/values)."
)
if self.head_dim % 2 != 0:
raise ValueError(
f"modelo '{self.name}': head_dim={self.head_dim} debe ser par "
"(RoPE rota los canales de a pares)."
)
return self
class OptimizerConfig(_Base):
lr: float = Field(gt=0.0)
beta1: float = Field(default=0.9, gt=0.0, lt=1.0)
beta2: float = Field(default=0.95, gt=0.0, lt=1.0)
eps: float = 1e-8
weight_decay: float = Field(default=0.1, ge=0.0)
grad_clip: float = Field(default=1.0, ge=0.0)
class ScheduleConfig(_Base):
"""WarmupStableDecay.
Se elige sobre cosine porque cosine obliga a fijar el total de pasos por
adelantado: extender una corrida o ramificar un checkpoint deja de ser
posible. Con WSD el LR se mantiene plano y solo decae al final, que es lo
que hace baratas las consolidaciones periódicas (Etapa 7b del plan).
"""
kind: Literal["wsd"] = "wsd"
warmup_steps: int = Field(ge=0)
decay_steps: int = Field(ge=0)
min_lr_ratio: float = Field(default=0.0, ge=0.0, le=1.0)
class DataConfig(_Base):
"""De dónde salen los tokens. `source` decide qué loader se usa."""
source: Literal["shards", "chars"] = "shards"
# source="shards": directorio con los .bin uint16 + index.json
shards_dir: str | None = None
# source="chars": un archivo de texto plano, para el smoke test
text_path: str | None = None
val_fraction: float = Field(default=0.005, gt=0.0, lt=0.5)
@model_validator(mode="after")
def _check_source(self) -> DataConfig:
required = {"shards": "shards_dir", "chars": "text_path"}[self.source]
if getattr(self, required) is None:
raise ValueError(f"data.source='{self.source}' exige data.{required}.")
return self
class TrainConfig(_Base):
run_name: str
out_dir: str = "runs"
seed: int = 1337
max_steps: int = Field(gt=0)
optimizer: OptimizerConfig
schedule: ScheduleConfig
log_every: int = Field(default=10, gt=0)
eval_every: int = Field(default=500, gt=0)
eval_batches: int = Field(default=20, gt=0)
checkpoint_every: int = Field(default=1000, gt=0)
sample_every: int = Field(default=0, ge=0) # 0 = no generar muestras
@model_validator(mode="after")
def _check_schedule_fits(self) -> TrainConfig:
s = self.schedule
if s.warmup_steps + s.decay_steps > self.max_steps:
raise ValueError(
f"schedule: warmup({s.warmup_steps}) + decay({s.decay_steps}) = "
f"{s.warmup_steps + s.decay_steps} supera max_steps={self.max_steps}; "
"no quedaría fase estable."
)
return self
class Config(_Base):
"""Config raíz: la composición de todas las capas."""
hardware: HardwareConfig
model: ModelConfig
train: TrainConfig
data: DataConfig
@model_validator(mode="after")
def _check_cross_layer(self) -> Config:
if self.data.source == "chars" and self.model.vocab_size > 1024:
raise ValueError(
f"data.source='chars' con model.vocab_size={self.model.vocab_size}: "
"el smoke test a nivel de caracteres usa un vocabulario chico "
"derivado del texto. Usá un modelo pensado para eso."
)
return self
+5
View File
@@ -0,0 +1,5 @@
"""Arquitectura del modelo. Una sola clase para todas las escalas."""
from enlace.model.transformer import RMSNorm, SwiGLU, Transformer, build_rope_cache
__all__ = ["RMSNorm", "SwiGLU", "Transformer", "build_rope_cache"]
+48
View File
@@ -0,0 +1,48 @@
"""Selección del backend de atención.
Es el único punto del modelo que conoce diferencias entre placas. FlashAttention-2
exige Ampere (sm_80) o superior: en la RTX 2060 (Turing, sm_75) hay que usar el
backend mem_efficient, que también es O(n) en memoria pero más lento. El perfil de
hardware decide; el modelo no sabe en qué placa corre.
"""
from __future__ import annotations
from contextlib import contextmanager
from typing import Iterator
import torch
from torch.nn.attention import SDPBackend, sdpa_kernel
_BACKENDS: dict[str, list[SDPBackend]] = {
# "auto" deja elegir a PyTorch: prueba flash, después mem_efficient, después
# la implementación matemática. Es lo correcto salvo que se quiera forzar
# una ruta concreta para medir o para evitar un kernel con bugs.
"auto": [SDPBackend.FLASH_ATTENTION, SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH],
"flash": [SDPBackend.FLASH_ATTENTION],
"mem_efficient": [SDPBackend.EFFICIENT_ATTENTION],
"math": [SDPBackend.MATH],
}
@contextmanager
def attention_backend(name: str) -> Iterator[None]:
"""Fija el backend de SDPA dentro del bloque.
En CPU no hay backends alternativos que elegir, así que es un no-op: forzar
uno ahí solo produce advertencias inútiles durante los tests.
"""
if not torch.cuda.is_available():
yield
return
try:
backends = _BACKENDS[name]
except KeyError:
raise ValueError(
f"backend de atención desconocido: {name!r}. "
f"Válidos: {sorted(_BACKENDS)}."
) from None
with sdpa_kernel(backends):
yield
+249
View File
@@ -0,0 +1,249 @@
"""Transformer decoder-only estilo Llama.
Una sola clase para todas las escalas del proyecto: la config decide si es el
modelo de 50M de la 2060 o el de 350M de la 5090. Las piezas son las que hoy
son estándar y por buenas razones:
- RMSNorm pre-norm: más barato que LayerNorm y más estable en profundidad.
- SwiGLU: mejor calidad por parámetro que un MLP con GELU.
- RoPE: posiciones relativas sin parámetros, y extensible después sin
reentrenar (importante para ampliar el contexto cuando la memoria crezca).
- GQA: reduce la caché KV en inferencia, que es el cuello de botella real
de un asistente que responde en tiempo real.
- QK-norm y z-loss: estabilizan el entrenamiento en float16, que en la RTX
2060 no es opcional porque Turing no soporta bfloat16.
"""
from __future__ import annotations
import math
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from enlace.config.schema import ModelConfig
from enlace.model.attention import attention_backend
class RMSNorm(nn.Module):
"""Normalización por raíz cuadrática media, sin término de sesgo."""
def __init__(self, dim: int, eps: float) -> None:
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
# El cálculo se hace en float32 aunque el tensor venga en fp16: la suma
# de cuadrados sobre d_model desborda el rango de fp16 con facilidad.
dtype = x.dtype
x32 = x.float()
normed = x32 * torch.rsqrt(x32.pow(2).mean(-1, keepdim=True) + self.eps)
return (normed * self.weight.float()).to(dtype)
def build_rope_cache(seq_len: int, head_dim: int, theta: float) -> tuple[Tensor, Tensor]:
"""Precalcula cos y sin de RoPE. Devuelve (seq_len, head_dim)."""
inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
pos = torch.arange(seq_len, dtype=torch.float32)
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim/2)
# Se duplica para la convención de mitades: [x1..xd/2, xd/2+1..xd].
emb = torch.cat((freqs, freqs), dim=-1)
return emb.cos(), emb.sin()
def _rotate_half(x: Tensor) -> Tensor:
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat((-x2, x1), dim=-1)
def apply_rope(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
"""Aplica RoPE a (batch, n_head, seq, head_dim)."""
cos = cos[None, None, :, :].to(x.dtype)
sin = sin[None, None, :, :].to(x.dtype)
return x * cos + _rotate_half(x) * sin
class Attention(nn.Module):
"""Atención causal con GQA y QK-norm."""
def __init__(self, cfg: ModelConfig, backend: str) -> None:
super().__init__()
self.n_head = cfg.n_head
self.n_kv_head = cfg.n_kv_head
self.head_dim = cfg.head_dim
self.n_rep = cfg.n_head // cfg.n_kv_head
self.backend = backend
self.dropout_p = cfg.dropout
self.wq = nn.Linear(cfg.d_model, cfg.n_head * cfg.head_dim, bias=False)
self.wk = nn.Linear(cfg.d_model, cfg.n_kv_head * cfg.head_dim, bias=False)
self.wv = nn.Linear(cfg.d_model, cfg.n_kv_head * cfg.head_dim, bias=False)
self.wo = nn.Linear(cfg.n_head * cfg.head_dim, cfg.d_model, bias=False)
# QK-norm: normalizar queries y keys antes de RoPE acota el crecimiento
# de los logits de atención, que es la causa habitual de los picos de
# loss en fp16. Cuesta casi nada y evita perder días de entrenamiento.
if cfg.qk_norm:
self.q_norm: nn.Module = RMSNorm(cfg.head_dim, cfg.norm_eps)
self.k_norm: nn.Module = RMSNorm(cfg.head_dim, cfg.norm_eps)
else:
self.q_norm = nn.Identity()
self.k_norm = nn.Identity()
def forward(self, x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
b, t, _ = x.shape
q = self.wq(x).view(b, t, self.n_head, self.head_dim).transpose(1, 2)
k = self.wk(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
v = self.wv(x).view(b, t, self.n_kv_head, self.head_dim).transpose(1, 2)
q = apply_rope(self.q_norm(q), cos, sin)
k = apply_rope(self.k_norm(k), cos, sin)
# GQA: cada cabeza de key/value se comparte entre n_rep cabezas de query.
if self.n_rep > 1:
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
with attention_backend(self.backend):
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout_p if self.training else 0.0
)
out = out.transpose(1, 2).contiguous().view(b, t, -1)
return self.wo(out)
class SwiGLU(nn.Module):
"""FFN con compuerta SiLU: w2(silu(w1(x)) * w3(x))."""
def __init__(self, cfg: ModelConfig) -> None:
super().__init__()
hidden = cfg.ffn_dim
self.w1 = nn.Linear(cfg.d_model, hidden, bias=False)
self.w3 = nn.Linear(cfg.d_model, hidden, bias=False)
self.w2 = nn.Linear(hidden, cfg.d_model, bias=False)
def forward(self, x: Tensor) -> Tensor:
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class Block(nn.Module):
def __init__(self, cfg: ModelConfig, backend: str) -> None:
super().__init__()
self.attn_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.attn = Attention(cfg, backend)
self.ffn_norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.ffn = SwiGLU(cfg)
def forward(self, x: Tensor, cos: Tensor, sin: Tensor) -> Tensor:
x = x + self.attn(self.attn_norm(x), cos, sin)
return x + self.ffn(self.ffn_norm(x))
class Transformer(nn.Module):
def __init__(self, cfg: ModelConfig, attention_backend: str = "auto") -> None:
super().__init__()
self.cfg = cfg
self.tok_emb = nn.Embedding(cfg.vocab_size, cfg.d_model)
self.drop = nn.Dropout(cfg.dropout)
self.blocks = nn.ModuleList(Block(cfg, attention_backend) for _ in range(cfg.n_layer))
self.norm = RMSNorm(cfg.d_model, cfg.norm_eps)
self.lm_head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
# Con vocab 32k y d_model 512 la tabla de embeddings son 16.7M
# parámetros: atarla con la salida ahorra un tercio de un modelo de 50M.
if cfg.tie_embeddings:
self.lm_head.weight = self.tok_emb.weight
cos, sin = build_rope_cache(cfg.seq_len, cfg.head_dim, cfg.rope_theta)
self.register_buffer("rope_cos", cos, persistent=False)
self.register_buffer("rope_sin", sin, persistent=False)
self.apply(self._init_weights)
# Las proyecciones residuales se inicializan más chicas: sin esto la
# varianza de la corriente residual crece con la profundidad.
scale = 1.0 / math.sqrt(2 * cfg.n_layer)
for name, p in self.named_parameters():
if name.endswith(("attn.wo.weight", "ffn.w2.weight")):
torch.nn.init.normal_(p, mean=0.0, std=cfg.init_std * scale)
def _init_weights(self, module: nn.Module) -> None:
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=self.cfg.init_std)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(module.weight, mean=0.0, std=self.cfg.init_std)
def num_parameters(self, non_embedding: bool = False) -> int:
total = sum(p.numel() for p in self.parameters())
if non_embedding:
total -= self.tok_emb.weight.numel()
return total
def forward(
self, idx: Tensor, targets: Tensor | None = None
) -> tuple[Tensor, Tensor | None]:
"""Devuelve (logits, loss). `loss` es None si no hay targets."""
_, t = idx.shape
if t > self.cfg.seq_len:
raise ValueError(
f"secuencia de largo {t} supera seq_len={self.cfg.seq_len}. "
"Recortá la entrada o ampliá el contexto por RoPE scaling."
)
cos = self.rope_cos[:t]
sin = self.rope_sin[:t]
x = self.drop(self.tok_emb(idx))
for block in self.blocks:
x = block(x, cos, sin)
x = self.norm(x)
logits = self.lm_head(x)
if targets is None:
return logits, None
# La entropía cruzada se calcula en float32: en fp16 el logsumexp sobre
# 32k clases pierde precisión y el gradiente se degrada.
flat_logits = logits.float().view(-1, logits.size(-1))
flat_targets = targets.reshape(-1)
loss = F.cross_entropy(flat_logits, flat_targets, ignore_index=-100)
# z-loss: penaliza que los logits crezcan en magnitud absoluta. Es lo
# que impide que el softmax sature y desborde entrenando en fp16.
if self.cfg.z_loss_weight > 0.0:
valid = flat_targets != -100
if valid.any():
z = torch.logsumexp(flat_logits[valid], dim=-1)
loss = loss + self.cfg.z_loss_weight * z.pow(2).mean()
return logits, loss
@torch.no_grad()
def generate(
self,
idx: Tensor,
max_new_tokens: int,
temperature: float = 1.0,
top_k: int | None = None,
) -> Tensor:
"""Generación sin caché KV: sirve para las muestras de control durante
el entrenamiento. La inferencia real usa la caché (enlace/serve/)."""
self.eval()
for _ in range(max_new_tokens):
window = idx[:, -self.cfg.seq_len :]
logits, _ = self(window)
logits = logits[:, -1, :] / max(temperature, 1e-5)
if top_k is not None:
k = min(top_k, logits.size(-1))
threshold = torch.topk(logits, k, dim=-1).values[:, -1:]
logits = logits.masked_fill(logits < threshold, float("-inf"))
probs = F.softmax(logits, dim=-1)
idx = torch.cat((idx, torch.multinomial(probs, num_samples=1)), dim=1)
return idx
+1
View File
@@ -0,0 +1 @@
"""Entrenamiento: bucle, schedules, checkpointing y detección de hardware."""
+100
View File
@@ -0,0 +1,100 @@
"""Todo lo que depende de la placa concreta vive acá.
Es el único archivo del entrenamiento que conoce diferencias entre GPUs. Cuando
llegue la RTX 5090, este archivo y `configs/hardware/` son lo único que debería
necesitar revisión: si hace falta tocar algo más, es una fuga de abstracción y
conviene arreglarla acá en vez de propagarla.
"""
from __future__ import annotations
import torch
from enlace.config.schema import HardwareConfig
_DTYPES: dict[str, torch.dtype] = {
"float32": torch.float32,
"bfloat16": torch.bfloat16,
"float16": torch.float16,
}
class HardwareError(RuntimeError):
"""El hardware real no coincide con lo que declara el perfil."""
def resolve_dtype(name: str) -> torch.dtype:
return _DTYPES[name]
def describe_device(cfg: HardwareConfig) -> str:
if cfg.device == "cpu":
return "cpu"
props = torch.cuda.get_device_properties(0)
cap = torch.cuda.get_device_capability(0)
return (
f"{props.name} (sm_{cap[0]}{cap[1]}, "
f"{props.total_memory / 1024**3:.1f} GB, torch {torch.__version__})"
)
def setup_device(cfg: HardwareConfig) -> torch.device:
"""Valida que el perfil corresponda a la placa real y prepara el device.
Falla temprano y explícito: descubrir a mitad de una corrida de dos días que
el perfil era el de otra placa es caro.
"""
if cfg.device == "cpu":
torch.set_float32_matmul_precision(cfg.matmul_precision)
return torch.device("cpu")
if not torch.cuda.is_available():
raise HardwareError(
f"el perfil '{cfg.name}' pide device=cuda pero PyTorch no ve ninguna GPU. "
"Para trabajar en una máquina sin GPU usá configs/hardware/cpu.yaml."
)
cap = torch.cuda.get_device_capability(0)
if cfg.min_compute_capability is not None:
required = tuple(cfg.min_compute_capability)
if cap < required:
raise HardwareError(
f"el perfil '{cfg.name}' exige compute capability "
f"sm_{required[0]}{required[1]} o superior, pero la GPU es "
f"sm_{cap[0]}{cap[1]}."
)
# Turing (sm_75) no tiene unidades bfloat16: PyTorch lo emula y el
# entrenamiento se vuelve inservible en vez de fallar. Mejor fallar.
if cfg.dtype == "bfloat16" and not torch.cuda.is_bf16_supported():
raise HardwareError(
f"el perfil '{cfg.name}' pide bfloat16, que esta GPU (sm_{cap[0]}{cap[1]}) "
"no soporta de forma nativa. Usá float16 con use_grad_scaler=true."
)
# FlashAttention-2 requiere Ampere o superior.
if cfg.attention_backend == "flash" and cap < (8, 0):
raise HardwareError(
f"el perfil '{cfg.name}' pide attention_backend=flash, que exige "
f"sm_80 o superior; esta GPU es sm_{cap[0]}{cap[1]}. "
"Usá mem_efficient."
)
torch.set_float32_matmul_precision(cfg.matmul_precision)
return torch.device("cuda")
def autocast_context(cfg: HardwareConfig, device: torch.device):
"""Contexto de precisión mixta acorde al perfil."""
if cfg.dtype == "float32":
return torch.autocast(device_type=device.type, enabled=False)
return torch.autocast(device_type=device.type, dtype=resolve_dtype(cfg.dtype))
def build_grad_scaler(cfg: HardwareConfig) -> torch.amp.GradScaler:
"""GradScaler solo tiene sentido en float16.
En bfloat16 el rango del exponente es el de float32, así que no hay nada que
escalar; el schema ya prohíbe la combinación, esto es la contraparte.
"""
return torch.amp.GradScaler(enabled=cfg.use_grad_scaler)
+131
View File
@@ -0,0 +1,131 @@
"""Checkpointing con reanudación exacta.
"Exacta" significa que reanudar en el paso N y seguir hasta N+K produce los
mismos pesos que una corrida ininterrumpida hasta N+K. Para eso no alcanza con
guardar el modelo: hacen falta el optimizer, el escalador de gradiente, la
posición en el stream de datos y el estado de los generadores aleatorios.
Se verifica a propósito y temprano (tests/test_resume.py). Un checkpoint que no
reanuda exacto es peor que no tener checkpoint, porque el daño es silencioso.
"""
from __future__ import annotations
import random
from pathlib import Path
from typing import Any
import numpy as np
import torch
from enlace.config.load import to_yaml
from enlace.config.schema import Config
CHECKPOINT_FORMAT = 1
def _rng_state() -> dict[str, Any]:
state: dict[str, Any] = {
"python": random.getstate(),
"numpy": np.random.get_state(),
"torch": torch.get_rng_state(),
}
if torch.cuda.is_available():
state["cuda"] = torch.cuda.get_rng_state_all()
return state
def _restore_rng(state: dict[str, Any]) -> None:
random.setstate(state["python"])
np.random.set_state(state["numpy"])
torch.set_rng_state(state["torch"])
if "cuda" in state and torch.cuda.is_available():
torch.cuda.set_rng_state_all(state["cuda"])
def save(
path: str | Path,
*,
config: Config,
step: int,
model: torch.nn.Module,
optimizer: torch.optim.Optimizer,
scaler: torch.amp.GradScaler,
stream_state: dict[str, Any],
metrics: dict[str, float],
) -> Path:
"""Escribe el checkpoint de forma atómica.
Se escribe a un temporal y se renombra: si el proceso muere a mitad de la
escritura, el checkpoint anterior sigue intacto en vez de quedar truncado.
"""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
# torch.compile envuelve el modelo y prefija las claves con "_orig_mod.";
# se guarda el módulo original para que el checkpoint no dependa de si la
# corrida usó compilación o no.
module = getattr(model, "_orig_mod", model)
payload = {
"format": CHECKPOINT_FORMAT,
"step": step,
"model": module.state_dict(),
"optimizer": optimizer.state_dict(),
"scaler": scaler.state_dict(),
"stream": stream_state,
"rng": _rng_state(),
"metrics": metrics,
# Un checkpoint sin su config no es reproducible.
"config": config.model_dump(mode="json"),
}
tmp = path.with_suffix(path.suffix + ".tmp")
torch.save(payload, tmp)
tmp.replace(path)
# Copia legible al lado, para poder inspeccionar la corrida sin abrir el .pt
path.with_suffix(".yaml").write_text(to_yaml(config))
return path
def load(
path: str | Path,
*,
model: torch.nn.Module,
optimizer: torch.optim.Optimizer | None = None,
scaler: torch.amp.GradScaler | None = None,
map_location: str | torch.device = "cpu",
) -> dict[str, Any]:
"""Restaura el estado y devuelve los metadatos (step, stream, metrics)."""
path = Path(path)
payload = torch.load(path, map_location=map_location, weights_only=False)
fmt = payload.get("format")
if fmt != CHECKPOINT_FORMAT:
raise ValueError(
f"{path}: formato de checkpoint {fmt}, se esperaba {CHECKPOINT_FORMAT}."
)
module = getattr(model, "_orig_mod", model)
module.load_state_dict(payload["model"])
if optimizer is not None:
optimizer.load_state_dict(payload["optimizer"])
if scaler is not None:
scaler.load_state_dict(payload["scaler"])
if "rng" in payload:
_restore_rng(payload["rng"])
return {
"step": payload["step"],
"stream": payload.get("stream", {}),
"metrics": payload.get("metrics", {}),
"config": payload.get("config", {}),
}
def latest(run_dir: str | Path) -> Path | None:
"""Devuelve el checkpoint más reciente de una corrida, si existe."""
candidates = sorted(Path(run_dir).glob("ckpt-*.pt"))
return candidates[-1] if candidates else None
+59
View File
@@ -0,0 +1,59 @@
"""Schedule de learning rate.
WSD (WarmupStableDecay) en vez de cosine, y la razón es arquitectónica: con
cosine hay que fijar el total de pasos por adelantado, así que extender una
corrida o ramificar un checkpoint deja de ser posible sin romper el schedule.
WSD sube el LR, lo mantiene plano la mayor parte de la corrida, y solo lo decae
al final. Eso habilita tres cosas que este proyecto necesita:
- extender un pretraining que quedó corto, sin recalcular nada;
- el annealing con el corpus de estilo, que ocurre durante la fase de decay;
- las consolidaciones periódicas, que son una fase de decay corta desde el
snapshot activo en vez de un reentrenamiento completo.
Advertencia práctica: el inicio del decay se ubica en `max_steps - decay_steps`,
así que cambiar `max_steps` mueve esa frontera. Para extender una corrida hay
que ramificar desde un checkpoint anterior al inicio del decay; reanudar con
otro `max_steps` sobre un checkpoint ya decaído da un LR incoherente con lo que
el modelo venía viendo. Reanudar una corrida interrumpida exige la config
idéntica (lo verifica tests/test_resume.py).
"""
from __future__ import annotations
import math
from enlace.config.schema import ScheduleConfig
def wsd_lr(step: int, base_lr: float, cfg: ScheduleConfig, max_steps: int) -> float:
"""LR para `step` (0-indexado).
El decaimiento usa 1 - sqrt(progreso), que empíricamente rinde mejor que el
lineal: mantiene el LR alto más tiempo y después cae rápido.
"""
if step < cfg.warmup_steps:
# Warmup lineal. El +1 evita un primer paso con LR exactamente cero,
# que desperdicia una actualización.
return base_lr * (step + 1) / cfg.warmup_steps
decay_start = max_steps - cfg.decay_steps
if step < decay_start or cfg.decay_steps == 0:
return base_lr
progress = (step - decay_start) / cfg.decay_steps
progress = min(max(progress, 0.0), 1.0)
factor = 1.0 - math.sqrt(progress)
return base_lr * (cfg.min_lr_ratio + (1.0 - cfg.min_lr_ratio) * factor)
def build_lr_fn(cfg: ScheduleConfig, base_lr: float, max_steps: int):
"""Devuelve una función step -> lr, para usar en el bucle de entrenamiento."""
if cfg.kind != "wsd":
raise ValueError(f"schedule desconocido: {cfg.kind}")
def lr_at(step: int) -> float:
return wsd_lr(step, base_lr, cfg, max_steps)
return lr_at
+276
View File
@@ -0,0 +1,276 @@
"""Bucle de entrenamiento.
No hay un solo hiperparámetro en este archivo: todo viene de la config. Lo que
sí vive acá es la mecánica que tiene que ser correcta pase lo que pase —
acumulación de gradiente, precisión mixta, recorte de norma, checkpointing
atómico y reanudación exacta.
python -m enlace.train.train configs/runs/smoke-2060.yaml
python -m enlace.train.train configs/runs/pretrain-2060.yaml --resume
"""
from __future__ import annotations
import json
import math
import random
import sys
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from enlace.config.load import ConfigError, load_config
from enlace.config.schema import Config
from enlace.data.loaders import build_stream
from enlace.model.transformer import Transformer
from enlace.train import checkpoint
from enlace.train.backends import (
autocast_context,
build_grad_scaler,
describe_device,
setup_device,
)
from enlace.train.schedules import build_lr_fn
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def build_optimizer(model: torch.nn.Module, cfg: Config) -> torch.optim.AdamW:
"""AdamW con weight decay solo sobre las matrices.
Aplicar decay a normas y sesgos (tensores de una dimensión) degrada la
calidad sin ahorrar nada: son pocos parámetros y penalizarlos solo distorsiona
la normalización.
"""
decay, no_decay = [], []
for param in model.parameters():
if not param.requires_grad:
continue
(decay if param.dim() >= 2 else no_decay).append(param)
opt = cfg.train.optimizer
return torch.optim.AdamW(
[
{"params": decay, "weight_decay": opt.weight_decay},
{"params": no_decay, "weight_decay": 0.0},
],
lr=opt.lr,
betas=(opt.beta1, opt.beta2),
eps=opt.eps,
)
def flops_per_token(model: Transformer) -> float:
"""Estimación estándar de FLOPs por token en forward+backward.
6*N por las multiplicaciones de matrices y 12*L*H*Q*T por la atención, que
no depende del número de parámetros sino del largo de contexto.
"""
cfg = model.cfg
n = model.num_parameters(non_embedding=True)
attn = 12 * cfg.n_layer * cfg.n_head * cfg.head_dim * cfg.seq_len
return 6 * n + attn
@torch.no_grad()
def evaluate(model: torch.nn.Module, stream, cfg: Config, device: torch.device) -> float:
"""Loss promedio en el split de validación."""
model.eval()
total = 0.0
for _ in range(cfg.train.eval_batches):
x, y = stream.next_batch("val")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
total += loss.item()
model.train()
return total / cfg.train.eval_batches
@torch.no_grad()
def sample(model: torch.nn.Module, stream, device: torch.device, n_tokens: int = 160) -> str:
"""Genera una muestra corta para inspección humana.
Leer estas muestras en cada checkpoint es la forma más rápida de detectar
que algo se rompió: las curvas de loss pueden verse bien mientras el modelo
produce basura.
"""
module = getattr(model, "_orig_mod", model)
start = torch.zeros((1, 1), dtype=torch.long, device=device)
out = module.generate(start, max_new_tokens=n_tokens, temperature=0.8, top_k=50)
model.train()
# Se descarta el token semilla: no es salida del modelo y ensucia la lectura.
generado = out[0, 1:].tolist()
decode = getattr(stream, "decode", None)
if decode is None:
return f"<sin decodificador> {generado[:16]}"
return decode(generado)
def train(cfg: Config, resume: bool = False) -> Path:
run_dir = Path(cfg.train.out_dir) / cfg.train.run_name
run_dir.mkdir(parents=True, exist_ok=True)
log_path = run_dir / "metrics.jsonl"
seed_everything(cfg.train.seed)
device = setup_device(cfg.hardware)
print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}")
print(f"[enlace] perfil: {cfg.hardware.name} | dtype {cfg.hardware.dtype}")
stream = build_stream(cfg, device)
if stream.vocab_size != cfg.model.vocab_size:
raise ConfigError(
f"el corpus tiene vocab_size={stream.vocab_size} pero el modelo "
f"declara {cfg.model.vocab_size}. Son incompatibles."
)
model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device)
n_params = model.num_parameters()
print(f"[enlace] modelo {cfg.model.name}: {n_params:,} parámetros")
optimizer = build_optimizer(model, cfg)
scaler = build_grad_scaler(cfg.hardware)
lr_at = build_lr_fn(cfg.train.schedule, cfg.train.optimizer.lr, cfg.train.max_steps)
start_step = 0
if resume:
ckpt_path = checkpoint.latest(run_dir)
if ckpt_path is None:
print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero")
else:
meta = checkpoint.load(
ckpt_path, model=model, optimizer=optimizer, scaler=scaler, map_location=device
)
stream.load_state_dict(meta["stream"])
start_step = meta["step"]
print(f"[enlace] reanudado desde {ckpt_path.name} en el paso {start_step}")
if cfg.hardware.compile:
print("[enlace] compilando el modelo (la primera iteración tarda)...")
model = torch.compile(model) # type: ignore[assignment]
tokens_per_step = (
cfg.hardware.effective_batch_size * cfg.model.seq_len
)
fpt = flops_per_token(getattr(model, "_orig_mod", model))
model.train()
print(
f"[enlace] {cfg.train.max_steps} pasos x {tokens_per_step:,} tokens "
f"= {cfg.train.max_steps * tokens_per_step / 1e9:.2f}B tokens"
)
t_last = time.perf_counter()
for step in range(start_step, cfg.train.max_steps):
lr = lr_at(step)
for group in optimizer.param_groups:
group["lr"] = lr
optimizer.zero_grad(set_to_none=True)
loss_sum = 0.0
for _ in range(cfg.hardware.grad_accum_steps):
x, y = stream.next_batch("train")
with autocast_context(cfg.hardware, device):
_, loss = model(x, y)
# Se divide por los micro-pasos para que el gradiente acumulado
# sea el promedio y no la suma: si no, el LR efectivo dependería
# de grad_accum_steps y las curvas no serían comparables entre
# perfiles de hardware.
loss = loss / cfg.hardware.grad_accum_steps
scaler.scale(loss).backward()
loss_sum += loss.item()
grad_norm = float("nan")
if cfg.train.optimizer.grad_clip > 0:
# Hay que deshacer el escalado antes de medir la norma, o el recorte
# se aplicaría sobre gradientes inflados por el GradScaler.
scaler.unscale_(optimizer)
grad_norm = float(
torch.nn.utils.clip_grad_norm_(
model.parameters(), cfg.train.optimizer.grad_clip
)
)
scaler.step(optimizer)
scaler.update()
if (step + 1) % cfg.train.log_every == 0:
if device.type == "cuda":
torch.cuda.synchronize()
now = time.perf_counter()
dt = (now - t_last) / cfg.train.log_every
t_last = now
record: dict[str, Any] = {
"step": step + 1,
"loss": round(loss_sum, 4),
"lr": lr,
"grad_norm": round(grad_norm, 4),
"tokens_per_s": round(tokens_per_step / dt),
"seconds_per_step": round(dt, 4),
}
if cfg.hardware.peak_tflops:
mfu = (fpt * tokens_per_step / dt) / (cfg.hardware.peak_tflops * 1e12)
record["mfu"] = round(mfu, 4)
with log_path.open("a") as fh:
fh.write(json.dumps(record) + "\n")
mfu_txt = f" mfu {record['mfu']:.1%}" if "mfu" in record else ""
print(
f"paso {step + 1:>7} | loss {loss_sum:.4f} | lr {lr:.2e} | "
f"|g| {grad_norm:.2f} | {record['tokens_per_s']:,} tok/s{mfu_txt}"
)
if (step + 1) % cfg.train.eval_every == 0:
val = evaluate(model, stream, cfg, device)
print(f"paso {step + 1:>7} | val_loss {val:.4f} | ppl {math.exp(min(val, 20)):.2f}")
with log_path.open("a") as fh:
fh.write(json.dumps({"step": step + 1, "val_loss": round(val, 4)}) + "\n")
t_last = time.perf_counter()
if cfg.train.sample_every and (step + 1) % cfg.train.sample_every == 0:
text = sample(model, stream, device)
(run_dir / "samples.txt").open("a").write(f"--- paso {step + 1} ---\n{text}\n\n")
print(f"muestra: {text[:120]!r}")
t_last = time.perf_counter()
if (step + 1) % cfg.train.checkpoint_every == 0 or (step + 1) == cfg.train.max_steps:
path = checkpoint.save(
run_dir / f"ckpt-{step + 1:08d}.pt",
config=cfg,
step=step + 1,
model=model,
optimizer=optimizer,
scaler=scaler,
stream_state=stream.state_dict(),
metrics={"loss": loss_sum},
)
print(f"[enlace] checkpoint -> {path.name}")
t_last = time.perf_counter()
return run_dir
def main() -> int:
args = [a for a in sys.argv[1:] if a != "--resume"]
resume = "--resume" in sys.argv[1:]
if not args:
print("uso: python -m enlace.train.train <config.yaml> [--resume] [clave=valor ...]")
return 2
try:
cfg = load_config(args[0], args[1:])
except ConfigError as exc:
print(f"[enlace] {exc}", file=sys.stderr)
return 2
train(cfg, resume=resume)
return 0
if __name__ == "__main__":
raise SystemExit(main())