Compare commits

...

3 Commits

Author SHA1 Message Date
msaldain 03dadc93da Aplicar ruff format a todo el código
YAML / yaml (push) Failing after 1m58s
Python / calidad (push) Successful in 6s
Python / tests (push) Failing after 1m22s
Cambio mecánico, sin efecto en el comportamiento: la suite pasa igual antes y
después. Va en un commit propio para no tapar los cambios con sentido.

Se agregan además dos flujos de verificación que corren en cada carga al
repositorio:

- YAML: yamllint para sintaxis y estilo, más la carga de cada config contra su
  esquema de pydantic. Son cosas distintas — un YAML puede ser sintácticamente
  perfecto y estar roto igual, con 'run_nombre' en vez de 'run_name'. Ese paso
  no instala torch: se verificó que la capa de configuración no lo importa, así
  que corre en segundos en vez de descargar dos gigas y medio de CUDA.
- Python: ruff check, ruff format --check y la suite completa con torch de CPU.

Las rutas ignoradas de .yamllint.yml van ancladas con barra inicial. Sin
anclar, 'data/' y 'runs/' excluían configs/data/ y configs/runs/ — siete
archivos, justo los que más importa revisar — y el linter pasaba en verde sin
haber mirado nada. Es el mismo defecto que ya había aparecido en .gitignore.
2026-07-28 07:25:51 -03:00
msaldain 4048936067 Corregir lo que reportó ruff, que nunca se había ejecutado
ruff estaba configurado en pyproject.toml desde el primer commit y jamás se
había corrido. Tenía 15 hallazgos. Dos importan más allá del estilo:

- zip() sin strict= trunca en silencio al más corto. En las comparaciones de
  lotes eso significa que un test podía pasar sin haber comparado todo. Donde
  los largos deben coincidir ahora es strict=True; donde difieren a propósito
  (pares consecutivos) queda strict=False, que documenta la intención.
- Un import sin usar delataba algo peor: BraveBackend se había escrito sin una
  sola prueba. Se agregan ocho, contra una respuesta con la forma que devuelve
  la API, incluidas la limpieza de etiquetas, el caso de límite de tasa —que
  tiene que distinguirse de 'no respondió'— y que la credencial viaje en la
  cabecera y nunca en la URL. El respaldo tiene que funcionar justo cuando el
  primario ya falló; merecía la misma cobertura.

El resto es orden de imports, collections.abc y líneas largas.
2026-07-28 07:25:51 -03:00
msaldain ba61125a3b Desactivar el buffering de salida en las corridas remotas
Sin terminal, Python retiene la salida estándar, así que el comando de registro
no mostraba nada durante minutos aunque el entrenamiento estuviera avanzando —
parecía colgado. Se comprobó comparando el archivo de registro, atrasado,
contra las métricas, que se escriben sin buffer y sí avanzaban.
2026-07-28 07:25:51 -03:00
23 changed files with 428 additions and 89 deletions
+64
View File
@@ -0,0 +1,64 @@
# Calidad del código Python, en cada carga.
#
# Existe por una razón concreta: ruff estaba configurado en pyproject.toml desde
# el primer commit y nunca se había ejecutado. Cuando por fin se corrió tenía 15
# hallazgos, uno de ellos real —`zip()` sin `strict=` trunca en silencio al más
# corto, así que un test que compara dos listas de distinto largo pasaría sin
# comparar todo— y otro que delataba un backend sin ninguna prueba.
#
# Una regla que no se ejecuta no es una regla.
name: Python
on:
push:
pull_request:
jobs:
calidad:
runs-on: ubuntu-latest
steps:
- name: Descargar el repositorio
uses: actions/checkout@v4
- name: Preparar Python
uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Instalar ruff
run: python -m pip install --upgrade pip "ruff>=0.6"
- name: Estilo y errores
# Las reglas viven en pyproject.toml, no acá: un solo lugar donde
# cambiarlas, y el mismo resultado corriendo `ruff check .` en local.
run: ruff check --output-format=github .
- name: Formato consistente
run: ruff format --check --diff .
tests:
runs-on: ubuntu-latest
steps:
- name: Descargar el repositorio
uses: actions/checkout@v4
- name: Preparar Python
uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Instalar dependencias
# torch de CPU explícitamente: la variante con CUDA pesa unos 2,5 GB y
# acá no hay placa. La suite entera corre en CPU a propósito.
run: |
python -m pip install --upgrade pip
python -m pip install torch --index-url https://download.pytorch.org/whl/cpu
python -m pip install -e ".[dev]"
- name: Suite de tests
# Se desactiva la carga automática de plugins: en entornos con otros
# paquetes instalados (ROS, por ejemplo) pytest intenta cargar sus
# plugins y falla al recolectar por motivos ajenos al proyecto.
env:
PYTEST_DISABLE_PLUGIN_AUTOLOAD: "1"
run: pytest -q
+45
View File
@@ -0,0 +1,45 @@
# Verificación de los YAML del repositorio, en cada carga.
#
# Dos niveles, porque atrapan cosas distintas:
#
# 1. yamllint mira la sintaxis y el estilo: indentación, duplicados, tabs.
# 2. La validación contra el esquema carga cada config con pydantic. Un YAML
# puede ser sintácticamente perfecto y estar roto igual — `n_layers` en vez
# de `n_layer`, un schedule que no entra en max_steps, float16 sin
# GradScaler. Eso solo lo ve el esquema.
#
# El segundo paso no instala torch a propósito: la capa de configuración no lo
# importa, así que el CI corre en segundos en vez de descargar 2,5 GB de CUDA.
name: YAML
on:
push:
pull_request:
jobs:
yaml:
runs-on: ubuntu-latest
steps:
- name: Descargar el repositorio
uses: actions/checkout@v4
- name: Preparar Python
uses: actions/setup-python@v5
with:
python-version: "3.12"
- name: Instalar dependencias
# Solo lo que hace falta: yamllint para la sintaxis, pydantic y
# omegaconf para el esquema. Nada de torch.
run: |
python -m pip install --upgrade pip
python -m pip install "yamllint>=1.35" "pydantic>=2.7" "omegaconf>=2.3"
- name: Sintaxis y estilo de los YAML
# --strict convierte las advertencias en errores: si una regla no vale
# la pena, se desactiva en .yamllint.yml, no se deja avisando para
# siempre.
run: yamllint --strict .
- name: Las configs cargan y pasan su esquema
run: python scripts/validar_configs.py
+45
View File
@@ -0,0 +1,45 @@
# Reglas de estilo para los YAML del repositorio.
#
# Se parte de las reglas por defecto de yamllint y se ajusta solo lo que choca
# con las convenciones ya establecidas del proyecto.
extends: default
rules:
# 100 caracteres, igual que ruff en pyproject.toml. Los comentarios de las
# configs explican decisiones y no entran cómodos en 80.
line-length:
max: 100
allow-non-breakable-words: true
# No se usa el marcador `---` al inicio de los archivos.
document-start: disable
# `on:` en los workflows es una clave legítima, no el booleano de YAML 1.1.
truthy:
check-keys: false
# Las listas dentro de un mapeo se escriben sin indentar de más:
# fallbacks:
# - brave
# y también en línea, `fallbacks: [brave]`. Ambas se aceptan.
indentation:
spaces: 2
indent-sequences: consistent
comments:
min-spaces-from-content: 1
# Los archivos generados por herramientas externas no siempre traen salto
# final; no vale la pena fallar el CI por eso.
new-line-at-end-of-file: enable
# Anclados con "/" inicial a propósito: sin anclar, "data/" y "runs/" excluyen
# cualquier directorio con ese nombre en cualquier nivel — incluidos
# configs/data/ y configs/runs/, que son justamente los que más importa
# revisar. El síntoma es un linter que pasa en verde sin haber mirado nada.
ignore: |
/.venv/
/data/
/runs/
/snapshots/
/checkpoints/
+11 -8
View File
@@ -30,7 +30,10 @@ from dataclasses import dataclass
from typing import Protocol from typing import Protocol
# Un navegador real: el endpoint lite rechaza clientes sin User-Agent. # Un navegador real: el endpoint lite rechaza clientes sin User-Agent.
_USER_AGENT = "Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/125 Safari/537.36" _USER_AGENT = (
"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 "
"(KHTML, like Gecko) Chrome/125 Safari/537.36"
)
_RE_LINK = re.compile( _RE_LINK = re.compile(
r"<a[^>]*?href=[\"'](?P<url>[^\"']+)[\"'][^>]*?class=['\"]result-link['\"][^>]*?>" r"<a[^>]*?href=[\"'](?P<url>[^\"']+)[\"'][^>]*?class=['\"]result-link['\"][^>]*?>"
@@ -168,7 +171,10 @@ class BraveBackend:
self.language = language self.language = language
def search(self, query: str, max_results: int) -> list[SearchResult]: def search(self, query: str, max_results: int) -> list[SearchResult]:
url = self.ENDPOINT + "?" + urllib.parse.urlencode( url = (
self.ENDPOINT
+ "?"
+ urllib.parse.urlencode(
{ {
"q": query, "q": query,
"count": max_results, "count": max_results,
@@ -176,6 +182,7 @@ class BraveBackend:
"search_lang": self.language, "search_lang": self.language,
} }
) )
)
request = urllib.request.Request( request = urllib.request.Request(
url, url,
headers={ headers={
@@ -283,9 +290,7 @@ class CadenaDeBackends:
self.ultimo_backend = nombre self.ultimo_backend = nombre
return resultados return resultados
raise SearchError( raise SearchError("ningún backend de búsqueda respondió — " + " | ".join(fallos))
"ningún backend de búsqueda respondió — " + " | ".join(fallos)
)
class WebSearch: class WebSearch:
@@ -357,9 +362,7 @@ def _construir_uno(nombre: str, config) -> SearchBackend:
) )
if nombre == "searxng": if nombre == "searxng":
if not config.searxng_url: if not config.searxng_url:
raise SearchError( raise SearchError("backend searxng sin searxng_url: definí ENLACE_SEARXNG_URL en .env")
"backend searxng sin searxng_url: definí ENLACE_SEARXNG_URL en .env"
)
return SearxNGBackend( return SearxNGBackend(
base_url=config.searxng_url, timeout=config.timeout, language=config.language base_url=config.searxng_url, timeout=config.timeout, language=config.language
) )
-1
View File
@@ -16,7 +16,6 @@ from typing import Any
from omegaconf import DictConfig, OmegaConf from omegaconf import DictConfig, OmegaConf
from enlace.config.env import load_dotenv from enlace.config.env import load_dotenv
from enlace.config.schema import AgentConfig, Config, DistillConfig from enlace.config.schema import AgentConfig, Config, DistillConfig
# Familias de capas que puede declarar un YAML raíz, en el orden en que se # Familias de capas que puede declarar un YAML raíz, en el orden en que se
+1 -3
View File
@@ -235,9 +235,7 @@ class SearchConfig(_Base):
""" """
backend: Literal["duckduckgo", "brave", "searxng"] = "duckduckgo" backend: Literal["duckduckgo", "brave", "searxng"] = "duckduckgo"
fallbacks: list[Literal["duckduckgo", "brave", "searxng"]] = Field( fallbacks: list[Literal["duckduckgo", "brave", "searxng"]] = Field(default_factory=list)
default_factory=list
)
region: str = "es-es" region: str = "es-es"
country: str = "uy" country: str = "uy"
language: str = "es" language: str = "es"
+6 -11
View File
@@ -33,9 +33,10 @@ from __future__ import annotations
import hashlib import hashlib
import json import json
import time import time
from collections.abc import Iterable, Iterator
from dataclasses import asdict, dataclass, field from dataclasses import asdict, dataclass, field
from pathlib import Path from pathlib import Path
from typing import Any, Iterable, Iterator, Protocol from typing import Any, Protocol
from enlace.config.env import load_dotenv from enlace.config.env import load_dotenv
from enlace.config.schema import DistillConfig from enlace.config.schema import DistillConfig
@@ -200,8 +201,7 @@ def _normalizar_resultado(item: Any) -> dict[str, Any]:
"input_tokens": uso.input_tokens, "input_tokens": uso.input_tokens,
"output_tokens": uso.output_tokens, "output_tokens": uso.output_tokens,
"cache_read_input_tokens": getattr(uso, "cache_read_input_tokens", 0) or 0, "cache_read_input_tokens": getattr(uso, "cache_read_input_tokens", 0) or 0,
"cache_creation_input_tokens": getattr(uso, "cache_creation_input_tokens", 0) "cache_creation_input_tokens": getattr(uso, "cache_creation_input_tokens", 0) or 0,
or 0,
} }
elif tipo == "errored": elif tipo == "errored":
salida["error_type"] = item.result.error.type salida["error_type"] = item.result.error.type
@@ -268,8 +268,7 @@ class Trabajo:
def leer_peticiones(self) -> list[Peticion]: def leer_peticiones(self) -> list[Peticion]:
if not self.peticiones_path.is_file(): if not self.peticiones_path.is_file():
raise DistillError( raise DistillError(
f"el trabajo '{self.nombre}' no tiene manifiesto: " f"el trabajo '{self.nombre}' no tiene manifiesto: falta {self.peticiones_path}"
f"falta {self.peticiones_path}"
) )
peticiones = [] peticiones = []
with self.peticiones_path.open(encoding="utf-8") as fh: with self.peticiones_path.open(encoding="utf-8") as fh:
@@ -340,9 +339,7 @@ def submit(
else: else:
pendientes = unicas pendientes = unicas
ya_enviados = { ya_enviados = {cid for lote in trabajo.leer_lotes() for cid in lote["custom_ids"]}
cid for lote in trabajo.leer_lotes() for cid in lote["custom_ids"]
}
pendientes = [p for p in pendientes if p.custom_id not in ya_enviados] pendientes = [p for p in pendientes if p.custom_id not in ya_enviados]
if not pendientes: if not pendientes:
@@ -376,9 +373,7 @@ def submit(
return [lote["batch_id"] for lote in lotes] return [lote["batch_id"] for lote in lotes]
def status( def status(trabajo: Trabajo, cliente: ClienteLotes) -> list[dict[str, Any]]:
trabajo: Trabajo, cliente: ClienteLotes
) -> list[dict[str, Any]]:
"""Estado de cada lote del trabajo.""" """Estado de cada lote del trabajo."""
estados = [] estados = []
for lote in trabajo.leer_lotes(): for lote in trabajo.leer_lotes():
+1 -3
View File
@@ -56,9 +56,7 @@ class ByteStream:
) )
raw = np.frombuffer(path.read_bytes(), dtype=np.uint8) raw = np.frombuffer(path.read_bytes(), dtype=np.uint8)
if len(raw) < seq_len * 4: if len(raw) < seq_len * 4:
raise ValueError( raise ValueError(f"{path} tiene {len(raw)} bytes: muy poco para seq_len={seq_len}.")
f"{path} tiene {len(raw)} bytes: muy poco para seq_len={seq_len}."
)
split_at = int(len(raw) * (1.0 - val_fraction)) split_at = int(len(raw) * (1.0 - val_fraction))
self._data = { self._data = {
+2 -3
View File
@@ -8,8 +8,8 @@ hardware decide; el modelo no sabe en qué placa corre.
from __future__ import annotations from __future__ import annotations
from collections.abc import Iterator
from contextlib import contextmanager from contextlib import contextmanager
from typing import Iterator
import torch import torch
from torch.nn.attention import SDPBackend, sdpa_kernel from torch.nn.attention import SDPBackend, sdpa_kernel
@@ -40,8 +40,7 @@ def attention_backend(name: str) -> Iterator[None]:
backends = _BACKENDS[name] backends = _BACKENDS[name]
except KeyError: except KeyError:
raise ValueError( raise ValueError(
f"backend de atención desconocido: {name!r}. " f"backend de atención desconocido: {name!r}. Válidos: {sorted(_BACKENDS)}."
f"Válidos: {sorted(_BACKENDS)}."
) from None ) from None
with sdpa_kernel(backends): with sdpa_kernel(backends):
+1 -3
View File
@@ -186,9 +186,7 @@ class Transformer(nn.Module):
total -= self.tok_emb.weight.numel() total -= self.tok_emb.weight.numel()
return total return total
def forward( def forward(self, idx: Tensor, targets: Tensor | None = None) -> tuple[Tensor, Tensor | None]:
self, idx: Tensor, targets: Tensor | None = None
) -> tuple[Tensor, Tensor | None]:
"""Devuelve (logits, loss). `loss` es None si no hay targets.""" """Devuelve (logits, loss). `loss` es None si no hay targets."""
_, t = idx.shape _, t = idx.shape
if t > self.cfg.seq_len: if t > self.cfg.seq_len:
+1 -3
View File
@@ -122,9 +122,7 @@ def load(
fmt = payload.get("format") fmt = payload.get("format")
if fmt != CHECKPOINT_FORMAT: if fmt != CHECKPOINT_FORMAT:
raise ValueError( raise ValueError(f"{path}: formato de checkpoint {fmt}, se esperaba {CHECKPOINT_FORMAT}.")
f"{path}: formato de checkpoint {fmt}, se esperaba {CHECKPOINT_FORMAT}."
)
module = getattr(model, "_orig_mod", model) module = getattr(model, "_orig_mod", model)
module.load_state_dict(payload["model"]) module.load_state_dict(payload["model"])
+3 -9
View File
@@ -146,9 +146,7 @@ def train(cfg: Config, resume: bool = False) -> Path:
if ckpt_path is None: if ckpt_path is None:
print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero") print(f"[enlace] --resume sin checkpoints en {run_dir}: se empieza de cero")
else: else:
meta = checkpoint.load( meta = checkpoint.load(ckpt_path, model=model, optimizer=optimizer, scaler=scaler)
ckpt_path, model=model, optimizer=optimizer, scaler=scaler
)
stream.load_state_dict(meta["stream"]) stream.load_state_dict(meta["stream"])
start_step = meta["step"] start_step = meta["step"]
print(f"[enlace] reanudado desde {ckpt_path.name} en el paso {start_step}") print(f"[enlace] reanudado desde {ckpt_path.name} en el paso {start_step}")
@@ -157,9 +155,7 @@ def train(cfg: Config, resume: bool = False) -> Path:
print("[enlace] compilando el modelo (la primera iteración tarda)...") print("[enlace] compilando el modelo (la primera iteración tarda)...")
model = torch.compile(model) # type: ignore[assignment] model = torch.compile(model) # type: ignore[assignment]
tokens_per_step = ( tokens_per_step = cfg.hardware.effective_batch_size * cfg.model.seq_len
cfg.hardware.effective_batch_size * cfg.model.seq_len
)
fpt = flops_per_token(getattr(model, "_orig_mod", model)) fpt = flops_per_token(getattr(model, "_orig_mod", model))
model.train() model.train()
@@ -194,9 +190,7 @@ def train(cfg: Config, resume: bool = False) -> Path:
# se aplicaría sobre gradientes inflados por el GradScaler. # se aplicaría sobre gradientes inflados por el GradScaler.
scaler.unscale_(optimizer) scaler.unscale_(optimizer)
grad_norm = float( grad_norm = float(
torch.nn.utils.clip_grad_norm_( torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.train.optimizer.grad_clip)
model.parameters(), cfg.train.optimizer.grad_clip
)
) )
scaler.step(optimizer) scaler.step(optimizer)
scaler.update() scaler.update()
+39 -14
View File
@@ -32,7 +32,12 @@ import torch
from enlace.config.load import ConfigError, load_config from enlace.config.load import ConfigError, load_config
from enlace.config.schema import Config from enlace.config.schema import Config
from enlace.model.transformer import Transformer from enlace.model.transformer import Transformer
from enlace.train.backends import autocast_context, build_grad_scaler, describe_device, setup_device from enlace.train.backends import (
autocast_context,
build_grad_scaler,
describe_device,
setup_device,
)
from enlace.train.train import build_optimizer, flops_per_token from enlace.train.train import build_optimizer, flops_per_token
@@ -52,7 +57,10 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
resultados: list[Resultado] = [] resultados: list[Resultado] = []
print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}") print(f"[enlace] dispositivo: {describe_device(cfg.hardware)}")
print(f"[enlace] perfil {cfg.hardware.name} | modelo {cfg.model.name} | dtype {cfg.hardware.dtype}") print(
f"[enlace] perfil {cfg.hardware.name} | modelo {cfg.model.name} "
f"| dtype {cfg.hardware.dtype}"
)
model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device) model = Transformer(cfg.model, cfg.hardware.attention_backend).to(device)
optimizer = build_optimizer(model, cfg) optimizer = build_optimizer(model, cfg)
@@ -76,7 +84,8 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
def micro_lote(): def micro_lote():
datos = torch.randint( datos = torch.randint(
0, cfg.model.vocab_size, 0,
cfg.model.vocab_size,
(cfg.hardware.micro_batch_size, cfg.model.seq_len + 1), (cfg.hardware.micro_batch_size, cfg.model.seq_len + 1),
generator=generador, generator=generador,
) )
@@ -131,16 +140,22 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
uso < 0.92, uso < 0.92,
"VRAM", "VRAM",
f"pico reservado {pico_reservado:.2f} GB de {total:.1f} GB ({uso:.0%}). " f"pico reservado {pico_reservado:.2f} GB de {total:.1f} GB ({uso:.0%}). "
+ ("Entra con margen." if uso < 0.85 else + (
"Ajustado: un pico de fragmentación puede tirar la corrida." if uso < 0.92 else "Entra con margen."
"NO ENTRA con seguridad — bajá micro_batch_size o seq_len."), if uso < 0.85
else "Ajustado: un pico de fragmentación puede tirar la corrida."
if uso < 0.92
else "NO ENTRA con seguridad — bajá micro_batch_size o seq_len."
),
) )
) )
# --- rendimiento --- # --- rendimiento ---
medio = sum(tiempos) / len(tiempos) medio = sum(tiempos) / len(tiempos)
tok_s = tokens_por_paso / medio tok_s = tokens_por_paso / medio
mfu = (fpt * tokens_por_paso / medio) / (cfg.hardware.peak_tflops * 1e12) if cfg.hardware.peak_tflops else None mfu = None
if cfg.hardware.peak_tflops:
mfu = (fpt * tokens_por_paso / medio) / (cfg.hardware.peak_tflops * 1e12)
horas = cfg.train.max_steps * medio / 3600 horas = cfg.train.max_steps * medio / 3600
tokens_totales = cfg.train.max_steps * tokens_por_paso tokens_totales = cfg.train.max_steps * tokens_por_paso
@@ -165,14 +180,23 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
Resultado( Resultado(
mfu > 0.15, mfu > 0.15,
"MFU", "MFU",
f"{mfu:.1%}{'razonable para esta placa' if mfu > 0.15 else 'bajo: revisá backend de atención, compile o tamaño de lote'}", f"{mfu:.1%}"
+ (
"razonable para esta placa"
if mfu > 0.15
else "bajo: revisá backend de atención, compile o tamaño de lote"
),
) )
) )
# --- estabilidad --- # --- estabilidad ---
finitas = all(p == p and abs(p) != float("inf") for p in perdidas) finitas = all(p == p and abs(p) != float("inf") for p in perdidas)
resultados.append( resultados.append(
Resultado(finitas, "Loss finita", f"{len(perdidas)} pasos medidos, todas finitas" if finitas else "apareció NaN o inf") Resultado(
finitas,
"Loss finita",
f"{len(perdidas)} pasos medidos, todas finitas" if finitas else "apareció NaN o inf",
)
) )
if cfg.hardware.use_grad_scaler: if cfg.hardware.use_grad_scaler:
@@ -182,15 +206,16 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
tasa <= 0.5, tasa <= 0.5,
"Estabilidad fp16", "Estabilidad fp16",
f"{salteados}/{pasos} pasos salteados por el GradScaler" f"{salteados}/{pasos} pasos salteados por el GradScaler"
+ (" (normal al calibrar la escala)" if tasa <= 0.5 + (
else " — tasa alta: la corrida perdería actualizaciones"), " (normal al calibrar la escala)"
if tasa <= 0.5
else " — tasa alta: la corrida perdería actualizaciones"
),
) )
) )
# --- parámetros --- # --- parámetros ---
resultados.append( resultados.append(Resultado(True, "Modelo", f"{n_params:,} parámetros ({n_params / 1e6:.1f}M)"))
Resultado(True, "Modelo", f"{n_params:,} parámetros ({n_params/1e6:.1f}M)")
)
return resultados return resultados
+1
View File
@@ -30,6 +30,7 @@ distill = [
dev = [ dev = [
"pytest>=8.0", "pytest>=8.0",
"ruff>=0.6", "ruff>=0.6",
"yamllint>=1.35",
] ]
[build-system] [build-system]
+4 -1
View File
@@ -105,9 +105,12 @@ run)
# distintas no se pisan y `logs` sabe a cuál conectarse. # distintas no se pisan y `logs` sabe a cuál conectarse.
sesion="enlace-$(basename "$config" .yaml)" sesion="enlace-$(basename "$config" .yaml)"
# tmux es lo que hace que cortar el SSH no mate el entrenamiento. # tmux es lo que hace que cortar el SSH no mate el entrenamiento.
# `-u` desactiva el buffering de stdout: sin eso Python retiene la salida
# porque no hay terminal, y `remote.sh logs` no muestra nada durante
# minutos aunque la corrida esté avanzando.
ssh "${SSH_OPTS[@]}" -t "$HOST" "cd '$DIR' && mkdir -p runs && \ ssh "${SSH_OPTS[@]}" -t "$HOST" "cd '$DIR' && mkdir -p runs && \
tmux new-session -d -s '$sesion' \ tmux new-session -d -s '$sesion' \
\"$PY -m enlace.train.train '$config' $* 2>&1 | tee -a 'runs/$sesion.log'\" \ \"$PY -u -m enlace.train.train '$config' $* 2>&1 | tee -a 'runs/$sesion.log'\" \
&& echo 'corriendo en tmux: $sesion'" && echo 'corriendo en tmux: $sesion'"
;; ;;
+76
View File
@@ -0,0 +1,76 @@
#!/usr/bin/env python3
"""Valida que todas las configs del repositorio carguen y pasen su esquema.
Un YAML sintácticamente correcto puede seguir estando roto: `n_layers` en vez de
`n_layer`, un schedule que no entra en `max_steps`, float16 sin GradScaler. Eso
no lo ve un linter de sintaxis — lo ven los validadores de pydantic, que ya
existen. Esto los ejecuta sobre cada config del repositorio.
No importa torch a propósito: así el CI valida configs en segundos en vez de
descargar los dos gigas y medio de CUDA.
python scripts/validar_configs.py
"""
from __future__ import annotations
import sys
from pathlib import Path
RAIZ = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(RAIZ))
from enlace.config.load import ( # noqa: E402
ConfigError,
load_agent_config,
load_config,
load_distill_config,
)
def validar() -> int:
fallas: list[tuple[Path, str]] = []
verificadas = 0
# Las configs de corrida componen las cuatro capas, así que validarlas
# cubre transitivamente hardware, modelo, entrenamiento y datos.
for path in sorted((RAIZ / "configs" / "runs").glob("*.yaml")):
try:
cfg = load_config(path)
verificadas += 1
print(
f" OK {path.relative_to(RAIZ)} "
f"({cfg.hardware.name} / {cfg.model.name} / {cfg.train.run_name})"
)
except ConfigError as exc:
fallas.append((path, str(exc)))
print(f" FALLA {path.relative_to(RAIZ)}")
# Los árboles que no son de entrenamiento se cargan aparte.
for etiqueta, cargador, path in (
("agente", load_agent_config, RAIZ / "configs" / "agent" / "tools.yaml"),
("destilación", load_distill_config, RAIZ / "configs" / "data" / "distill.yaml"),
):
try:
cargador(path)
verificadas += 1
print(f" OK {path.relative_to(RAIZ)} ({etiqueta})")
except ConfigError as exc:
fallas.append((path, str(exc)))
print(f" FALLA {path.relative_to(RAIZ)}")
print()
if fallas:
for path, motivo in fallas:
print(f"--- {path.relative_to(RAIZ)} ---")
print(motivo)
print()
print(f"{len(fallas)} config(s) inválida(s) de {verificadas + len(fallas)}.")
return 1
print(f"{verificadas} configs válidas.")
return 0
if __name__ == "__main__":
raise SystemExit(validar())
+4 -1
View File
@@ -107,7 +107,10 @@ def test_rechaza_gqa_incoherente(tmp_path):
def test_rechaza_schedule_que_no_entra(tmp_path): def test_rechaza_schedule_que_no_entra(tmp_path):
_espera_error( _espera_error(
_raiz(tmp_path, train={"max_steps": 100, "schedule": {"warmup_steps": 80, "decay_steps": 80}}), _raiz(
tmp_path,
train={"max_steps": 100, "schedule": {"warmup_steps": 80, "decay_steps": 80}},
),
"no quedaría fase estable", "no quedaría fase estable",
) )
+5 -9
View File
@@ -8,8 +8,9 @@ no el SDK de Anthropic.
from __future__ import annotations from __future__ import annotations
import json import json
from collections.abc import Iterator
from pathlib import Path from pathlib import Path
from typing import Any, Iterator from typing import Any
import pytest import pytest
@@ -105,8 +106,7 @@ def trabajo(tmp_path: Path) -> Trabajo:
def _peticiones(n: int) -> list[Peticion]: def _peticiones(n: int) -> list[Peticion]:
return [ return [
Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i}) Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i}) for i in range(n)
for i in range(n)
] ]
@@ -229,9 +229,7 @@ def test_collect_espera_a_los_lotes_en_curso(trabajo, config):
cliente = ClienteFalso(rondas_hasta_terminar=2) cliente = ClienteFalso(rondas_hasta_terminar=2)
submit(trabajo, _peticiones(3), cliente, config) submit(trabajo, _peticiones(3), cliente, config)
dormidas = [] dormidas = []
resumen = collect( resumen = collect(trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append)
trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append
)
assert resumen["ok"] == 3 assert resumen["ok"] == 3
assert dormidas # efectivamente esperó assert dormidas # efectivamente esperó
@@ -241,9 +239,7 @@ def test_collect_respeta_el_tope_de_espera(trabajo, config):
submit(trabajo, _peticiones(3), cliente, config) submit(trabajo, _peticiones(3), cliente, config)
reloj = iter([0.0, 1e9, 1e9]) reloj = iter([0.0, 1e9, 1e9])
with pytest.raises(DistillError, match="max_wait_hours"): with pytest.raises(DistillError, match="max_wait_hours"):
collect( collect(trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None)
trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None
)
def test_collect_sin_lotes_da_un_error_util(trabajo, config): def test_collect_sin_lotes_da_un_error_util(trabajo, config):
+2 -2
View File
@@ -46,7 +46,7 @@ def test_restaurar_el_estado_continua_la_misma_secuencia(texto_es):
b.load_state_dict(estado) b.load_state_dict(estado)
obtenido = [b.next_batch("train")[0] for _ in range(3)] obtenido = [b.next_batch("train")[0] for _ in range(3)]
for e, o in zip(esperado, obtenido): for e, o in zip(esperado, obtenido, strict=True):
assert torch.equal(e, o) assert torch.equal(e, o)
@@ -62,7 +62,7 @@ def test_evaluar_no_altera_la_secuencia_de_entrenamiento(texto_es):
b.next_batch("val") b.next_batch("val")
lotes.append(b.next_batch("train")[0]) lotes.append(b.next_batch("train")[0])
for e, o in zip(esperado, lotes): for e, o in zip(esperado, lotes, strict=True):
assert torch.equal(e, o) assert torch.equal(e, o)
+2 -1
View File
@@ -144,9 +144,10 @@ def test_z_loss_penaliza_logits_grandes():
def test_el_modelo_del_plan_pesa_lo_esperado(): def test_el_modelo_del_plan_pesa_lo_esperado():
"""tiny-50m tiene que estar cerca de 50M: es lo que entra en la 2060.""" """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 from pathlib import Path
from enlace.config.load import load_config
cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml") cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml")
model = Transformer(cfg.model, "math") model = Transformer(cfg.model, "math")
assert 45e6 < model.num_parameters() < 55e6 assert 45e6 < model.num_parameters() < 55e6
+1 -3
View File
@@ -92,9 +92,7 @@ def test_el_estado_rng_se_restaura_desde_cpu(monkeypatch):
recibidos = [] recibidos = []
original = torch.set_rng_state original = torch.set_rng_state
monkeypatch.setattr( monkeypatch.setattr(torch, "set_rng_state", lambda s: (recibidos.append(s), original(s))[1])
torch, "set_rng_state", lambda s: (recibidos.append(s), original(s))[1]
)
_restore_rng(_rng_state()) _restore_rng(_rng_state())
assert recibidos, "no se llamó a set_rng_state" assert recibidos, "no se llamó a set_rng_state"
+2 -4
View File
@@ -7,9 +7,7 @@ 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): def _lr_fn(max_steps=1000, warmup=100, decay=200, min_ratio=0.0, base=1e-3):
cfg = ScheduleConfig( cfg = ScheduleConfig(kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio)
kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio
)
return build_lr_fn(cfg, base, max_steps) return build_lr_fn(cfg, base, max_steps)
@@ -28,7 +26,7 @@ def test_la_fase_estable_es_plana():
def test_el_decay_baja_de_forma_monotona_hasta_cero(): def test_el_decay_baja_de_forma_monotona_hasta_cero():
lr = _lr_fn() lr = _lr_fn()
valores = [lr(s) for s in range(800, 1000)] valores = [lr(s) for s in range(800, 1000)]
assert all(a >= b for a, b in zip(valores, valores[1:])) assert all(a >= b for a, b in zip(valores, valores[1:], strict=False))
assert valores[0] == 1e-3 assert valores[0] == 1e-3
assert valores[-1] < 1e-4 assert valores[-1] < 1e-4
+102
View File
@@ -7,6 +7,7 @@ no en producción. Nada en esta suite sale a internet.
from __future__ import annotations from __future__ import annotations
import json
from pathlib import Path from pathlib import Path
import pytest import pytest
@@ -322,3 +323,104 @@ def test_la_config_del_repo_declara_brave_de_respaldo():
cfg = load_agent_config() cfg = load_agent_config()
assert cfg.search.backend == "duckduckgo" assert cfg.search.backend == "duckduckgo"
assert "brave" in cfg.search.fallbacks assert "brave" in cfg.search.fallbacks
# --- parser de Brave -------------------------------------------------------
#
# El respaldo tiene que funcionar justo cuando el primario ya falló, así que su
# parser merece la misma cobertura. La respuesta de abajo reproduce la forma que
# devuelve la API de Brave.
_BRAVE = {
"web": {
"results": [
{
"title": "Home Assistant",
"url": "https://www.home-assistant.io/",
"description": "Domótica <strong>libre</strong> y con privacidad.",
},
{
"title": "Home Assistant &mdash; Wikipedia",
"url": "https://es.wikipedia.org/wiki/Home_Assistant",
"description": "Software de automatización del hogar.",
},
{"title": "Sin url", "url": "", "description": "no debería aparecer"},
]
}
}
class _RespuestaFalsa:
def __init__(self, payload: dict) -> None:
self._cuerpo = json.dumps(payload).encode()
def read(self) -> bytes:
return self._cuerpo
def __enter__(self):
return self
def __exit__(self, *_):
return False
def _brave_con(payload, monkeypatch):
backend = BraveBackend(api_key="clave-de-prueba")
monkeypatch.setattr("urllib.request.urlopen", lambda *a, **k: _RespuestaFalsa(payload))
return backend
def test_brave_extrae_titulo_url_y_descripcion(monkeypatch):
resultados = _brave_con(_BRAVE, monkeypatch).search("home assistant", 5)
assert len(resultados) == 2 # el que no tiene url se descarta
assert resultados[0].title == "Home Assistant"
assert resultados[0].url == "https://www.home-assistant.io/"
def test_brave_limpia_etiquetas_y_entidades(monkeypatch):
resultados = _brave_con(_BRAVE, monkeypatch).search("q", 5)
assert "<strong>" not in resultados[0].snippet
assert "&mdash;" not in resultados[1].title
# Los espacios múltiples se colapsan, igual que en el parser de DuckDuckGo.
assert " " not in resultados[1].snippet
def test_brave_respeta_el_maximo(monkeypatch):
assert len(_brave_con(_BRAVE, monkeypatch).search("q", 1)) == 1
def test_brave_sin_resultados_devuelve_lista_vacia(monkeypatch):
assert _brave_con({"web": {"results": []}}, monkeypatch).search("q", 5) == []
def test_brave_respuesta_sin_la_clave_web_no_rompe(monkeypatch):
assert _brave_con({"type": "search"}, monkeypatch).search("q", 5) == []
def test_brave_distingue_el_limite_de_tasa(monkeypatch):
"""Si el respaldo también está limitado, el mensaje tiene que decirlo: es
un diagnóstico distinto de 'no respondió'."""
import urllib.error
def falla_429(*a, **k):
raise urllib.error.HTTPError("u", 429, "Too Many Requests", {}, None)
monkeypatch.setattr("urllib.request.urlopen", falla_429)
with pytest.raises(SearchError, match="límite de tasa"):
BraveBackend(api_key="x").search("q", 5)
def test_brave_manda_la_credencial_en_la_cabecera(monkeypatch):
"""La clave va en X-Subscription-Token, nunca en la URL: una credencial en
la query string termina en los registros de todos los intermediarios."""
capturado = {}
def espia(request, *a, **k):
capturado["url"] = request.full_url
capturado["headers"] = request.headers
return _RespuestaFalsa({"web": {"results": []}})
monkeypatch.setattr("urllib.request.urlopen", espia)
BraveBackend(api_key="secreta-123").search("q", 5)
assert "secreta-123" not in capturado["url"]
assert capturado["headers"]["X-subscription-token"] == "secreta-123"