Compare commits
5 Commits
80691131e3
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 7fac220338 | |||
| 59b7c3acb6 | |||
| 03dadc93da | |||
| 4048936067 | |||
| ba61125a3b |
@@ -0,0 +1,7 @@
|
|||||||
|
{
|
||||||
|
"permissions": {
|
||||||
|
"allow": [
|
||||||
|
"Bash(*)"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -12,6 +12,12 @@
|
|||||||
# Secretos
|
# Secretos
|
||||||
.env
|
.env
|
||||||
|
|
||||||
|
# Config personal de Claude Code. Hasta ahora solo la ignoraba un gitignore
|
||||||
|
# global de esta máquina, que no existe en arkax ni en un clon nuevo: ahí el
|
||||||
|
# archivo se habría subido sin que nadie lo notara. La regla tiene que vivir en
|
||||||
|
# el repositorio, no en la configuración de una computadora.
|
||||||
|
.claude/settings.local.json
|
||||||
|
|
||||||
# Config del IDE, específica de cada máquina
|
# Config del IDE, específica de cada máquina
|
||||||
.vscode/
|
.vscode/
|
||||||
|
|
||||||
|
|||||||
@@ -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/
|
||||||
@@ -26,11 +26,15 @@ import time
|
|||||||
import urllib.error
|
import urllib.error
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
import urllib.request
|
import urllib.request
|
||||||
|
import warnings
|
||||||
from dataclasses import dataclass
|
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,13 +172,17 @@ 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
|
||||||
"q": query,
|
+ "?"
|
||||||
"count": max_results,
|
+ urllib.parse.urlencode(
|
||||||
"country": self.country,
|
{
|
||||||
"search_lang": self.language,
|
"q": query,
|
||||||
}
|
"count": max_results,
|
||||||
|
"country": self.country,
|
||||||
|
"search_lang": self.language,
|
||||||
|
}
|
||||||
|
)
|
||||||
)
|
)
|
||||||
request = urllib.request.Request(
|
request = urllib.request.Request(
|
||||||
url,
|
url,
|
||||||
@@ -283,9 +291,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 +363,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
|
||||||
)
|
)
|
||||||
@@ -371,7 +375,18 @@ def build_backend(config) -> SearchBackend:
|
|||||||
|
|
||||||
Con un solo backend devuelve ese backend pelado, para no envolver en una
|
Con un solo backend devuelve ese backend pelado, para no envolver en una
|
||||||
cadena algo que no la necesita.
|
cadena algo que no la necesita.
|
||||||
|
|
||||||
|
Los respaldos sin credencial ya vienen filtrados de la config, pero se avisa
|
||||||
|
acá: quedarse sin red de seguridad no puede pasar en silencio, y este es el
|
||||||
|
punto por el que pasa cualquier entrypoint que use búsqueda.
|
||||||
"""
|
"""
|
||||||
|
for nombre, variable in config.respaldos_omitidos:
|
||||||
|
warnings.warn(
|
||||||
|
f"búsqueda: el respaldo '{nombre}' queda deshabilitado porque falta "
|
||||||
|
f"{variable}. Si el primario '{config.backend}' falla, no hay red.",
|
||||||
|
RuntimeWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
cadena = [_construir_uno(nombre, config) for nombre in config.cadena]
|
cadena = [_construir_uno(nombre, config) for nombre in config.cadena]
|
||||||
return cadena[0] if len(cadena) == 1 else CadenaDeBackends(cadena)
|
return cadena[0] if len(cadena) == 1 else CadenaDeBackends(cadena)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+52
-24
@@ -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"
|
||||||
@@ -252,33 +250,63 @@ class SearchConfig(_Base):
|
|||||||
searxng_url: str | None = None
|
searxng_url: str | None = None
|
||||||
brave_api_key: str | None = None
|
brave_api_key: str | None = None
|
||||||
|
|
||||||
|
# Qué credencial necesita cada backend, y de qué variable de entorno sale.
|
||||||
|
_CREDENCIALES = {
|
||||||
|
"searxng": ("searxng_url", "ENLACE_SEARXNG_URL"),
|
||||||
|
"brave": ("brave_api_key", "ENLACE_BRAVE_API_KEY"),
|
||||||
|
}
|
||||||
|
|
||||||
|
def _falta_credencial(self, nombre: str) -> str | None:
|
||||||
|
requisito = self._CREDENCIALES.get(nombre)
|
||||||
|
if requisito and not getattr(self, requisito[0]):
|
||||||
|
return requisito[1]
|
||||||
|
return None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def cadena(self) -> list[str]:
|
def cadena(self) -> list[str]:
|
||||||
"""El primario seguido de los respaldos, sin repetidos."""
|
"""El primario seguido de los respaldos utilizables, sin repetidos.
|
||||||
orden = [self.backend, *self.fallbacks]
|
|
||||||
|
Un respaldo sin su credencial se omite en vez de romper la carga. La
|
||||||
|
razón es concreta: la config vive en el repositorio y las credenciales
|
||||||
|
no, así que exigirlas para poder *leerla* dejaba el proyecto sin poder
|
||||||
|
clonarse — el CI y cualquier máquina nueva fallaban antes de empezar.
|
||||||
|
|
||||||
|
Omitir no es esconder: `respaldos_omitidos` los expone y `build_backend`
|
||||||
|
avisa. El primario es otra cosa y sigue siendo error fatal (ver abajo):
|
||||||
|
sin él no hay búsqueda posible.
|
||||||
|
"""
|
||||||
vistos: list[str] = []
|
vistos: list[str] = []
|
||||||
for nombre in orden:
|
for nombre in [self.backend, *self.fallbacks]:
|
||||||
if nombre not in vistos:
|
if nombre in vistos:
|
||||||
vistos.append(nombre)
|
continue
|
||||||
|
if nombre != self.backend and self._falta_credencial(nombre):
|
||||||
|
continue
|
||||||
|
vistos.append(nombre)
|
||||||
return vistos
|
return vistos
|
||||||
|
|
||||||
|
@property
|
||||||
|
def respaldos_omitidos(self) -> list[tuple[str, str]]:
|
||||||
|
"""Respaldos descartados por falta de credencial, con qué variable falta."""
|
||||||
|
omitidos = []
|
||||||
|
for nombre in self.fallbacks:
|
||||||
|
if nombre == self.backend:
|
||||||
|
continue
|
||||||
|
variable = self._falta_credencial(nombre)
|
||||||
|
if variable:
|
||||||
|
omitidos.append((nombre, variable))
|
||||||
|
return omitidos
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def _check_backends(self) -> SearchConfig:
|
def _check_backend_primario(self) -> SearchConfig:
|
||||||
# Un backend sin credencial se descubriría recién en la primera consulta
|
# El primario sí es fatal: sin él no queda ninguna búsqueda en pie, y
|
||||||
# real, que es justo cuando el primario ya falló y el respaldo tiene que
|
# descubrirlo en la primera consulta real es tarde.
|
||||||
# funcionar. Se valida al arrancar.
|
variable = self._falta_credencial(self.backend)
|
||||||
requisitos = {
|
if variable:
|
||||||
"searxng": ("searxng_url", "ENLACE_SEARXNG_URL"),
|
campo = self._CREDENCIALES[self.backend][0]
|
||||||
"brave": ("brave_api_key", "ENLACE_BRAVE_API_KEY"),
|
raise ValueError(
|
||||||
}
|
f"search: el backend primario '{self.backend}' exige {campo} "
|
||||||
for nombre in self.cadena:
|
f"(definí {variable} en .env)."
|
||||||
requisito = requisitos.get(nombre)
|
)
|
||||||
if requisito and not getattr(self, requisito[0]):
|
|
||||||
rol = "backend" if nombre == self.backend else "fallback"
|
|
||||||
raise ValueError(
|
|
||||||
f"search: el {rol} '{nombre}' exige {requisito[0]} "
|
|
||||||
f"(definí {requisito[1]} en .env)."
|
|
||||||
)
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+6
-11
@@ -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():
|
||||||
|
|||||||
@@ -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 = {
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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"])
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
@@ -148,7 +163,7 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
|
|||||||
Resultado(
|
Resultado(
|
||||||
True,
|
True,
|
||||||
"Rendimiento",
|
"Rendimiento",
|
||||||
f"{tok_s:,.0f} tok/s | {medio*1000:.0f} ms/paso"
|
f"{tok_s:,.0f} tok/s | {medio * 1000:.0f} ms/paso"
|
||||||
+ (f" | MFU {mfu:.1%}" if mfu else ""),
|
+ (f" | MFU {mfu:.1%}" if mfu else ""),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -157,7 +172,7 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]:
|
|||||||
True,
|
True,
|
||||||
"Proyección",
|
"Proyección",
|
||||||
f"{cfg.train.max_steps:,} pasos x {tokens_por_paso:,} tok = "
|
f"{cfg.train.max_steps:,} pasos x {tokens_por_paso:,} tok = "
|
||||||
f"{tokens_totales/1e9:.2f}B tokens en ~{horas:.1f} h ({horas/24:.1f} días)",
|
f"{tokens_totales / 1e9:.2f}B tokens en ~{horas:.1f} h ({horas / 24:.1f} días)",
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
if mfu is not None:
|
if mfu is not None:
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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'"
|
||||||
;;
|
;;
|
||||||
|
|
||||||
|
|||||||
Executable
+76
@@ -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())
|
||||||
@@ -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",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
@@ -52,5 +50,5 @@ def test_extender_la_corrida_mueve_el_inicio_del_decay():
|
|||||||
"""
|
"""
|
||||||
corto = _lr_fn(max_steps=1000)
|
corto = _lr_fn(max_steps=1000)
|
||||||
largo = _lr_fn(max_steps=2000)
|
largo = _lr_fn(max_steps=2000)
|
||||||
assert corto(850) < corto(700) # ya está decayendo
|
assert corto(850) < corto(700) # ya está decayendo
|
||||||
assert largo(850) == largo(700) # todavía en la fase estable
|
assert largo(850) == largo(700) # todavía en la fase estable
|
||||||
|
|||||||
+143
-10
@@ -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
|
||||||
@@ -192,7 +193,7 @@ def test_el_primario_no_necesita_credenciales():
|
|||||||
assert isinstance(primero, DuckDuckGoBackend)
|
assert isinstance(primero, DuckDuckGoBackend)
|
||||||
|
|
||||||
|
|
||||||
def test_searxng_exige_url():
|
def test_searxng_como_primario_exige_url():
|
||||||
with pytest.raises(ValueError, match="searxng_url"):
|
with pytest.raises(ValueError, match="searxng_url"):
|
||||||
SearchConfig(backend="searxng")
|
SearchConfig(backend="searxng")
|
||||||
|
|
||||||
@@ -306,19 +307,151 @@ def test_el_fallback_repetido_no_se_duplica():
|
|||||||
assert cfg.cadena == ["duckduckgo", "brave"]
|
assert cfg.cadena == ["duckduckgo", "brave"]
|
||||||
|
|
||||||
|
|
||||||
def test_un_fallback_sin_credencial_falla_al_arrancar():
|
def test_un_respaldo_sin_credencial_se_omite_en_vez_de_romper():
|
||||||
"""No en la primera consulta real — que es justo cuando el primario ya
|
"""La config vive en el repositorio y las credenciales no.
|
||||||
falló y el respaldo tiene que funcionar."""
|
|
||||||
with pytest.raises(ValueError, match="fallback 'brave'"):
|
Exigir la clave para poder *leer* la config dejaba el proyecto sin poder
|
||||||
SearchConfig(backend="duckduckgo", fallbacks=["brave"])
|
clonarse: el CI y cualquier máquina nueva fallaban antes de empezar. El
|
||||||
|
respaldo se cae de la cadena; el primario sigue funcionando.
|
||||||
|
"""
|
||||||
|
cfg = SearchConfig(backend="duckduckgo", fallbacks=["brave"])
|
||||||
|
assert cfg.cadena == ["duckduckgo"]
|
||||||
|
|
||||||
|
|
||||||
def test_brave_como_primario_tambien_exige_credencial():
|
def test_el_respaldo_omitido_queda_a_la_vista():
|
||||||
with pytest.raises(ValueError, match="backend 'brave'"):
|
"""Omitir no es esconder: quedarse sin red de seguridad tiene que verse."""
|
||||||
|
cfg = SearchConfig(backend="duckduckgo", fallbacks=["brave"])
|
||||||
|
assert cfg.respaldos_omitidos == [("brave", "ENLACE_BRAVE_API_KEY")]
|
||||||
|
|
||||||
|
|
||||||
|
def test_construir_la_cadena_avisa_del_respaldo_faltante():
|
||||||
|
cfg = SearchConfig(backend="duckduckgo", fallbacks=["brave"])
|
||||||
|
with pytest.warns(RuntimeWarning, match="ENLACE_BRAVE_API_KEY"):
|
||||||
|
build_backend(cfg)
|
||||||
|
|
||||||
|
|
||||||
|
def test_con_credencial_el_respaldo_si_entra():
|
||||||
|
cfg = SearchConfig(backend="duckduckgo", fallbacks=["brave"], brave_api_key="x")
|
||||||
|
assert cfg.cadena == ["duckduckgo", "brave"]
|
||||||
|
assert cfg.respaldos_omitidos == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_el_primario_sin_credencial_si_es_fatal():
|
||||||
|
"""Sin primario no queda ninguna búsqueda en pie: eso no se degrada."""
|
||||||
|
with pytest.raises(ValueError, match="primario 'brave'"):
|
||||||
SearchConfig(backend="brave")
|
SearchConfig(backend="brave")
|
||||||
|
|
||||||
|
|
||||||
def test_la_config_del_repo_declara_brave_de_respaldo():
|
def test_la_config_del_repo_carga_sin_credenciales(monkeypatch):
|
||||||
|
"""Un clon limpio, sin .env, tiene que poder leer la config del repositorio.
|
||||||
|
|
||||||
|
Es exactamente lo que rompió en el CI: la validación exigía la clave de
|
||||||
|
Brave para poder cargar el archivo.
|
||||||
|
"""
|
||||||
|
monkeypatch.delenv("ENLACE_BRAVE_API_KEY", raising=False)
|
||||||
|
monkeypatch.setattr("enlace.config.load.load_dotenv", lambda *a, **k: [])
|
||||||
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 # la intención queda declarada
|
||||||
|
assert cfg.search.cadena == ["duckduckgo"] # pero no se usa sin credencial
|
||||||
|
|
||||||
|
|
||||||
|
# --- 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 — 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 "—" 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"
|
||||||
|
|||||||
Reference in New Issue
Block a user