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.
This commit is contained in:
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+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():
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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 — 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