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,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())
|
||||
Reference in New Issue
Block a user