diff --git a/enlace/agent/tools/search.py b/enlace/agent/tools/search.py index cf59796..f99560b 100644 --- a/enlace/agent/tools/search.py +++ b/enlace/agent/tools/search.py @@ -30,7 +30,10 @@ from dataclasses import dataclass from typing import Protocol # 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( r"]*?href=[\"'](?P[^\"']+)[\"'][^>]*?class=['\"]result-link['\"][^>]*?>" @@ -168,13 +171,17 @@ class BraveBackend: self.language = language def search(self, query: str, max_results: int) -> list[SearchResult]: - url = self.ENDPOINT + "?" + urllib.parse.urlencode( - { - "q": query, - "count": max_results, - "country": self.country, - "search_lang": self.language, - } + url = ( + self.ENDPOINT + + "?" + + urllib.parse.urlencode( + { + "q": query, + "count": max_results, + "country": self.country, + "search_lang": self.language, + } + ) ) request = urllib.request.Request( url, @@ -283,9 +290,7 @@ class CadenaDeBackends: self.ultimo_backend = nombre return resultados - raise SearchError( - "ningún backend de búsqueda respondió — " + " | ".join(fallos) - ) + raise SearchError("ningún backend de búsqueda respondió — " + " | ".join(fallos)) class WebSearch: @@ -357,9 +362,7 @@ def _construir_uno(nombre: str, config) -> SearchBackend: ) if nombre == "searxng": if not config.searxng_url: - raise SearchError( - "backend searxng sin searxng_url: definí ENLACE_SEARXNG_URL en .env" - ) + raise SearchError("backend searxng sin searxng_url: definí ENLACE_SEARXNG_URL en .env") return SearxNGBackend( base_url=config.searxng_url, timeout=config.timeout, language=config.language ) diff --git a/enlace/config/load.py b/enlace/config/load.py index c3f6ba9..18c1369 100644 --- a/enlace/config/load.py +++ b/enlace/config/load.py @@ -16,7 +16,6 @@ from typing import Any from omegaconf import DictConfig, OmegaConf from enlace.config.env import load_dotenv - from enlace.config.schema import AgentConfig, Config, DistillConfig # Familias de capas que puede declarar un YAML raíz, en el orden en que se diff --git a/enlace/data/distill.py b/enlace/data/distill.py index 9fee570..3c5d741 100644 --- a/enlace/data/distill.py +++ b/enlace/data/distill.py @@ -33,9 +33,10 @@ from __future__ import annotations import hashlib import json import time +from collections.abc import Iterable, Iterator from dataclasses import asdict, dataclass, field 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.schema import DistillConfig @@ -200,8 +201,7 @@ def _normalizar_resultado(item: Any) -> dict[str, Any]: "input_tokens": uso.input_tokens, "output_tokens": uso.output_tokens, "cache_read_input_tokens": getattr(uso, "cache_read_input_tokens", 0) or 0, - "cache_creation_input_tokens": getattr(uso, "cache_creation_input_tokens", 0) - or 0, + "cache_creation_input_tokens": getattr(uso, "cache_creation_input_tokens", 0) or 0, } elif tipo == "errored": salida["error_type"] = item.result.error.type @@ -268,8 +268,7 @@ class Trabajo: def leer_peticiones(self) -> list[Peticion]: if not self.peticiones_path.is_file(): raise DistillError( - f"el trabajo '{self.nombre}' no tiene manifiesto: " - f"falta {self.peticiones_path}" + f"el trabajo '{self.nombre}' no tiene manifiesto: falta {self.peticiones_path}" ) peticiones = [] with self.peticiones_path.open(encoding="utf-8") as fh: @@ -340,9 +339,7 @@ def submit( else: pendientes = unicas - ya_enviados = { - cid for lote in trabajo.leer_lotes() for cid in lote["custom_ids"] - } + ya_enviados = {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] if not pendientes: @@ -376,9 +373,7 @@ def submit( return [lote["batch_id"] for lote in lotes] -def status( - trabajo: Trabajo, cliente: ClienteLotes -) -> list[dict[str, Any]]: +def status(trabajo: Trabajo, cliente: ClienteLotes) -> list[dict[str, Any]]: """Estado de cada lote del trabajo.""" estados = [] for lote in trabajo.leer_lotes(): diff --git a/enlace/model/attention.py b/enlace/model/attention.py index 14f1116..5aa14d6 100644 --- a/enlace/model/attention.py +++ b/enlace/model/attention.py @@ -8,8 +8,8 @@ hardware decide; el modelo no sabe en qué placa corre. from __future__ import annotations +from collections.abc import Iterator from contextlib import contextmanager -from typing import Iterator import torch from torch.nn.attention import SDPBackend, sdpa_kernel @@ -40,8 +40,7 @@ def attention_backend(name: str) -> Iterator[None]: backends = _BACKENDS[name] except KeyError: raise ValueError( - f"backend de atención desconocido: {name!r}. " - f"Válidos: {sorted(_BACKENDS)}." + f"backend de atención desconocido: {name!r}. Válidos: {sorted(_BACKENDS)}." ) from None with sdpa_kernel(backends): diff --git a/enlace/train/verificacion.py b/enlace/train/verificacion.py index bbc1078..9dba548 100644 --- a/enlace/train/verificacion.py +++ b/enlace/train/verificacion.py @@ -32,7 +32,12 @@ import torch from enlace.config.load import ConfigError, load_config from enlace.config.schema import Config 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 @@ -52,7 +57,10 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: resultados: list[Resultado] = [] 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) optimizer = build_optimizer(model, cfg) @@ -76,7 +84,8 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: def micro_lote(): datos = torch.randint( - 0, cfg.model.vocab_size, + 0, + cfg.model.vocab_size, (cfg.hardware.micro_batch_size, cfg.model.seq_len + 1), generator=generador, ) @@ -131,16 +140,22 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: uso < 0.92, "VRAM", 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 - "NO ENTRA con seguridad — bajá micro_batch_size o seq_len."), + + ( + "Entra con margen." + 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 --- medio = sum(tiempos) / len(tiempos) 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 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( True, "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 ""), ) ) @@ -157,7 +172,7 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: True, "Proyección", 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: @@ -165,14 +180,23 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: Resultado( mfu > 0.15, "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 --- finitas = all(p == p and abs(p) != float("inf") for p in perdidas) 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: @@ -182,15 +206,16 @@ def verificar(cfg: Config, pasos: int = 12, warmup: int = 3) -> list[Resultado]: tasa <= 0.5, "Estabilidad fp16", 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 --- - resultados.append( - Resultado(True, "Modelo", f"{n_params:,} parámetros ({n_params/1e6:.1f}M)") - ) + resultados.append(Resultado(True, "Modelo", f"{n_params:,} parámetros ({n_params / 1e6:.1f}M)")) return resultados diff --git a/tests/test_config.py b/tests/test_config.py index 8dc905a..007bb33 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -107,7 +107,10 @@ def test_rechaza_gqa_incoherente(tmp_path): def test_rechaza_schedule_que_no_entra(tmp_path): _espera_error( - _raiz(tmp_path, train={"max_steps": 100, "schedule": {"warmup_steps": 80, "decay_steps": 80}}), + _raiz( + tmp_path, + train={"max_steps": 100, "schedule": {"warmup_steps": 80, "decay_steps": 80}}, + ), "no quedaría fase estable", ) diff --git a/tests/test_distill.py b/tests/test_distill.py index 1afba14..e8f8e30 100644 --- a/tests/test_distill.py +++ b/tests/test_distill.py @@ -8,8 +8,9 @@ no el SDK de Anthropic. from __future__ import annotations import json +from collections.abc import Iterator from pathlib import Path -from typing import Any, Iterator +from typing import Any import pytest @@ -105,8 +106,7 @@ def trabajo(tmp_path: Path) -> Trabajo: def _peticiones(n: int) -> list[Peticion]: return [ - Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i}) - for i in range(n) + Peticion(system="Sos ENLACE.", user=f"consulta {i}", meta={"semilla": i}) 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) submit(trabajo, _peticiones(3), cliente, config) dormidas = [] - resumen = collect( - trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append - ) + resumen = collect(trabajo, cliente, config, ahora=lambda: 0.0, dormir=dormidas.append) assert resumen["ok"] == 3 assert dormidas # efectivamente esperó @@ -241,9 +239,7 @@ def test_collect_respeta_el_tope_de_espera(trabajo, config): submit(trabajo, _peticiones(3), cliente, config) reloj = iter([0.0, 1e9, 1e9]) with pytest.raises(DistillError, match="max_wait_hours"): - collect( - trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None - ) + collect(trabajo, cliente, config, ahora=lambda: next(reloj), dormir=lambda _: None) def test_collect_sin_lotes_da_un_error_util(trabajo, config): diff --git a/tests/test_loaders.py b/tests/test_loaders.py index 2af87b3..1c6c972 100644 --- a/tests/test_loaders.py +++ b/tests/test_loaders.py @@ -46,7 +46,7 @@ def test_restaurar_el_estado_continua_la_misma_secuencia(texto_es): b.load_state_dict(estado) 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) @@ -62,7 +62,7 @@ def test_evaluar_no_altera_la_secuencia_de_entrenamiento(texto_es): b.next_batch("val") 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) diff --git a/tests/test_model.py b/tests/test_model.py index 03088cb..b2b11ec 100644 --- a/tests/test_model.py +++ b/tests/test_model.py @@ -144,9 +144,10 @@ def test_z_loss_penaliza_logits_grandes(): def test_el_modelo_del_plan_pesa_lo_esperado(): """tiny-50m tiene que estar cerca de 50M: es lo que entra en la 2060.""" - from enlace.config.load import load_config from pathlib import Path + from enlace.config.load import load_config + cfg = load_config(Path(__file__).resolve().parents[1] / "configs/runs/pretrain-2060.yaml") model = Transformer(cfg.model, "math") assert 45e6 < model.num_parameters() < 55e6 diff --git a/tests/test_schedules.py b/tests/test_schedules.py index a3c97d7..6a91525 100644 --- a/tests/test_schedules.py +++ b/tests/test_schedules.py @@ -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): - cfg = ScheduleConfig( - kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio - ) + cfg = ScheduleConfig(kind="wsd", warmup_steps=warmup, decay_steps=decay, min_lr_ratio=min_ratio) 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(): lr = _lr_fn() valores = [lr(s) for s in range(800, 1000)] - assert all(a >= b for a, b in zip(valores, valores[1:])) + assert all(a >= b for a, b in zip(valores, valores[1:], strict=False)) assert valores[0] == 1e-3 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) largo = _lr_fn(max_steps=2000) - assert corto(850) < corto(700) # ya está decayendo - assert largo(850) == largo(700) # todavía en la fase estable + assert corto(850) < corto(700) # ya está decayendo + assert largo(850) == largo(700) # todavía en la fase estable diff --git a/tests/test_search.py b/tests/test_search.py index 02955da..4d336b8 100644 --- a/tests/test_search.py +++ b/tests/test_search.py @@ -7,6 +7,7 @@ no en producción. Nada en esta suite sale a internet. from __future__ import annotations +import json from pathlib import Path import pytest @@ -322,3 +323,104 @@ def test_la_config_del_repo_declara_brave_de_respaldo(): cfg = load_agent_config() assert cfg.search.backend == "duckduckgo" 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 libre 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 "" 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"