From d0a05cdd97fed3d01c62b980170347f0cced0282 Mon Sep 17 00:00:00 2001 From: Mateo Saldain Date: Mon, 27 Jul 2026 23:27:40 -0300 Subject: [PATCH] =?UTF-8?q?Anclar=20las=20rutas=20de=20.gitignore=20al=20r?= =?UTF-8?q?a=C3=ADz=20del=20repo?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Un patrón sin "/" inicial en .gitignore matchea cualquier directorio con ese nombre en cualquier nivel. "data/" y "runs/" estaban excluyendo, además de lo que se quería: enlace/data/ los cargadores de datos (ByteStream, ShardStream) configs/data/ corpus.yaml, smoke.yaml, distill.yaml configs/runs/ todos los configs ejecutables Es decir que el repo commiteado no contenía ni la capa de datos ni un solo config con el que arrancar un entrenamiento. Como scripts/remote.sh sincroniza por git push/pull, el servidor habría recibido un paquete que falla al importar, y el síntoma habría aparecido recién allá. Se anclan las cuatro rutas con "/" inicial y se documenta el porqué en el propio archivo. Co-Authored-By: Claude Opus 5 --- .gitignore | 13 +- configs/data/corpus.yaml | 5 + configs/data/distill.yaml | 36 ++ configs/data/smoke.yaml | 5 + configs/runs/pretrain-2060.yaml | 7 + configs/runs/smoke-2060.yaml | 13 + configs/runs/smoke-cpu.yaml | 21 ++ enlace/data/__init__.py | 5 + enlace/data/distill.py | 566 ++++++++++++++++++++++++++++++++ enlace/data/loaders.py | 215 ++++++++++++ 10 files changed, 882 insertions(+), 4 deletions(-) create mode 100644 configs/data/corpus.yaml create mode 100644 configs/data/distill.yaml create mode 100644 configs/data/smoke.yaml create mode 100644 configs/runs/pretrain-2060.yaml create mode 100644 configs/runs/smoke-2060.yaml create mode 100644 configs/runs/smoke-cpu.yaml create mode 100644 enlace/data/__init__.py create mode 100644 enlace/data/distill.py create mode 100644 enlace/data/loaders.py diff --git a/.gitignore b/.gitignore index 0093185..315b1fd 100644 --- a/.gitignore +++ b/.gitignore @@ -1,8 +1,13 @@ # Datos, pesos y bases del sistema: viven en el servidor, nunca en git. -data/ -checkpoints/ -snapshots/ -runs/ +# +# Las rutas van ancladas con "/" inicial a propósito. Sin anclar, "data/" +# excluye CUALQUIER directorio llamado data en cualquier nivel — incluidos +# enlace/data/ (los cargadores) y configs/data/. El síntoma sería un servidor +# que recibe el código por git y falla al importar. +/data/ +/checkpoints/ +/snapshots/ +/runs/ # Secretos .env diff --git a/configs/data/corpus.yaml b/configs/data/corpus.yaml new file mode 100644 index 0000000..68dc528 --- /dev/null +++ b/configs/data/corpus.yaml @@ -0,0 +1,5 @@ +# Corpus tokenizado en shards uint16 (Etapa 1 del plan). +# Los shards viven en el servidor; nunca se descargan enteros. +source: shards +shards_dir: data/shards +val_fraction: 0.005 diff --git a/configs/data/distill.yaml b/configs/data/distill.yaml new file mode 100644 index 0000000..39d26b9 --- /dev/null +++ b/configs/data/distill.yaml @@ -0,0 +1,36 @@ +# Generación de datos sintéticos con la Batch API (Etapa 3 del plan). +# +# Se usa la Batch API y no llamadas sueltas porque cuesta la mitad y generar un +# dataset no es sensible a la latencia. A cambio el trabajo es asincrónico y +# puede tardar horas: el estado vive en disco y es reanudable. +# +# Esto NO hace falta para entrenar ni para correr ENLACE. Solo se usa al armar +# el dataset SFT, y solo para la parte difícil: identidad y casos límite. El +# grueso sale de plantillas (cero modelo) y de un 7-8B local en la 2060. + +model: claude-opus-5 + +# En este modelo el pensamiento está activo por defecto y max_tokens acota +# pensamiento + respuesta juntos: sin holgura la salida se corta a la mitad. +# 8000 alcanza de sobra para ejemplos de SFT, que son cortos. +max_tokens: 8000 + +# Generar parafraseos y ejemplos etiquetados no es trabajo intensivo en +# razonamiento. `low` baja el costo bastante sin perder calidad acá; subilo a +# `medium` para el set de identidad, que sí exige coherencia fina. +effort: low + +# Tope de la API: 100.000 pedidos o 256 MB por lote. Se parte bien por debajo +# para no depender de que una request HTTP enorme llegue entera. +requests_per_batch: 10000 + +poll_interval_seconds: 60 +max_wait_hours: 26 # la API promete < 24 h; el resto es red de seguridad + +output_dir: data/distill + +# Solo para el reporte de costo estimado. Los precios reales los factura +# Anthropic; esto sirve para saber en qué orden de magnitud estamos. +usd_per_mtok_input: 5.0 +usd_per_mtok_output: 25.0 +batch_discount: 0.5 diff --git a/configs/data/smoke.yaml b/configs/data/smoke.yaml new file mode 100644 index 0000000..d8f4018 --- /dev/null +++ b/configs/data/smoke.yaml @@ -0,0 +1,5 @@ +# Texto plano en español para el smoke test a nivel de caracteres. +# El vocabulario se deriva del propio texto al cargarlo. +source: chars +text_path: data/smoke/texto.txt +val_fraction: 0.05 diff --git a/configs/runs/pretrain-2060.yaml b/configs/runs/pretrain-2060.yaml new file mode 100644 index 0000000..2071fda --- /dev/null +++ b/configs/runs/pretrain-2060.yaml @@ -0,0 +1,7 @@ +# Pretraining real del modelo base en la RTX 2060 (Etapa 2 del plan). +# ~1-2 días de cómputo. Correr siempre bajo tmux: ver scripts/remote.sh. +include: + hardware: hardware/turing-2060.yaml + model: model/tiny-50m.yaml + train: train/pretrain.yaml + data: data/corpus.yaml diff --git a/configs/runs/smoke-2060.yaml b/configs/runs/smoke-2060.yaml new file mode 100644 index 0000000..c6595e3 --- /dev/null +++ b/configs/runs/smoke-2060.yaml @@ -0,0 +1,13 @@ +# Smoke test de la Etapa 0 en el servidor con la RTX 2060. +# Criterio: loss < 1.5 y texto legible en menos de 15 minutos. +include: + hardware: hardware/turing-2060.yaml + model: model/char-smoke.yaml + train: train/smoke.yaml + data: data/smoke.yaml + +# El modelo char es diminuto: en la 2060 entra un micro-lote grande sin +# acumulación. Sobrescribe el perfil de la placa solo en lo que corresponde. +hardware: + micro_batch_size: 64 + grad_accum_steps: 1 diff --git a/configs/runs/smoke-cpu.yaml b/configs/runs/smoke-cpu.yaml new file mode 100644 index 0000000..f6f4755 --- /dev/null +++ b/configs/runs/smoke-cpu.yaml @@ -0,0 +1,21 @@ +# Smoke test en CPU — para verificar el código en la máquina de trabajo, +# que no tiene GPU. No entrena nada útil: comprueba que el stack funciona. +include: + hardware: hardware/cpu.yaml + model: model/char-smoke.yaml + train: train/smoke.yaml + data: data/smoke.yaml + +train: + run_name: smoke-cpu + max_steps: 300 + schedule: + kind: wsd + warmup_steps: 30 + decay_steps: 60 + min_lr_ratio: 0.0 + log_every: 20 + eval_every: 100 + eval_batches: 5 + checkpoint_every: 150 + sample_every: 150 diff --git a/enlace/data/__init__.py b/enlace/data/__init__.py new file mode 100644 index 0000000..7fe7924 --- /dev/null +++ b/enlace/data/__init__.py @@ -0,0 +1,5 @@ +"""Datos: descarga, limpieza, tokenización a shards y carga con reanudación.""" + +from enlace.data.loaders import ByteStream, ShardStream, build_stream + +__all__ = ["ByteStream", "ShardStream", "build_stream"] diff --git a/enlace/data/distill.py b/enlace/data/distill.py new file mode 100644 index 0000000..a612b35 --- /dev/null +++ b/enlace/data/distill.py @@ -0,0 +1,566 @@ +"""Generación de datos sintéticos con la Batch API. + +Se usa la Batch API y no llamadas sueltas por dos razones: cuesta la mitad, y +generar un dataset no es sensible a la latencia — es exactamente su caso de uso. +A cambio el trabajo es asincrónico y puede tardar horas, así que **todo el +estado vive en disco y el trabajo es reanudable**. Matar el proceso, perder la +conexión o reiniciar el servidor no pierde nada ni gasta de nuevo. + +Cuatro decisiones que valen la aclaración: + + 1. **El `custom_id` se deriva del contenido del pedido** (hash del prompt). + Eso hace el trabajo idempotente: re-generar la misma especificación produce + los mismos identificadores, así que un pedido ya resuelto se reconoce y no + se vuelve a pagar. + 2. **Los resultados se indexan por `custom_id`, nunca por posición.** La API + los devuelve en cualquier orden; emparejarlos por índice mezcla las + respuestas en silencio, y un dataset mal alineado es peor que uno vacío + porque parece correcto. + 3. **Los fallos se registran, no se descartan.** Van a `errors.jsonl` con su + motivo, y `submit --retry-failed` vuelve a enviar solo esos. + 4. **El prompt de sistema compartido va marcado para caché.** Es el mismo + para miles de pedidos; a ~0,1x del precio de entrada, deja de contar. + +Uso: + + python -m enlace.data.distill submit --name parafraseo + python -m enlace.data.distill status --name parafraseo + python -m enlace.data.distill collect --name parafraseo +""" + +from __future__ import annotations + +import hashlib +import json +import time +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import Any, Iterable, Iterator, Protocol + +from enlace.config.schema import DistillConfig + + +class DistillError(RuntimeError): + """Fallo de la generación por lotes, con mensaje legible.""" + + +# --- especificación de un pedido ------------------------------------------ + + +@dataclass(frozen=True, slots=True) +class Peticion: + """Un pedido de generación, independiente del SDK. + + El `custom_id` no se pasa: se deriva del contenido, para que el trabajo sea + idempotente y re-ejecutable sin duplicar costo. + """ + + system: str + user: str + # Esquema JSON opcional. Con él la respuesta viene estructurada y validada + # por la API, que para generar datasets es la diferencia entre parsear texto + # a mano y recibir filas listas. + schema: dict[str, Any] | None = None + # Metadatos propios: no van a la API, se conservan junto al resultado para + # saber de qué semilla salió cada ejemplo. + meta: dict[str, Any] = field(default_factory=dict) + + @property + def custom_id(self) -> str: + material = json.dumps( + {"system": self.system, "user": self.user, "schema": self.schema}, + sort_keys=True, + ensure_ascii=False, + ) + return "req-" + hashlib.sha256(material.encode()).hexdigest()[:24] + + +class ClienteLotes(Protocol): + """Lo mínimo que el runner necesita de un cliente de lotes. + + Está declarado como protocolo para que los tests corran contra un doble sin + red ni credenciales: el harness es lo que hay que verificar, no el SDK. + """ + + def crear(self, pedidos: list[dict[str, Any]]) -> str: ... + def estado(self, batch_id: str) -> dict[str, Any]: ... + def resultados(self, batch_id: str) -> Iterator[dict[str, Any]]: ... + def cancelar(self, batch_id: str) -> None: ... + + +# --- cliente real ---------------------------------------------------------- + + +class ClienteAnthropic: + """Adaptador sobre el SDK oficial. + + El SDK se importa acá adentro y no arriba: el servidor de entrenamiento no + necesita `anthropic` instalado, y el paquete tiene que importar igual. + """ + + def __init__(self, config: DistillConfig) -> None: + try: + import anthropic + except ImportError: + raise DistillError( + "falta el SDK: instalá el extra con `pip install -e '.[distill]'`" + ) from None + self._anthropic = anthropic + self._client = anthropic.Anthropic() + self.config = config + + def construir_params(self, peticion: Peticion) -> dict[str, Any]: + """Traduce una Petición a los parámetros de la Messages API.""" + params: dict[str, Any] = { + "model": self.config.model, + "max_tokens": self.config.max_tokens, + # El prompt de sistema es idéntico entre miles de pedidos: marcarlo + # para caché lo deja a ~0,1x del precio de entrada. + "system": [ + { + "type": "text", + "text": peticion.system, + "cache_control": {"type": "ephemeral"}, + } + ], + "messages": [{"role": "user", "content": peticion.user}], + "thinking": {"type": "adaptive"}, + "output_config": {"effort": self.config.effort}, + } + if peticion.schema is not None: + params["output_config"]["format"] = { + "type": "json_schema", + "schema": peticion.schema, + } + return params + + def crear(self, pedidos: list[dict[str, Any]]) -> str: + from anthropic.types.message_create_params import MessageCreateParamsNonStreaming + from anthropic.types.messages.batch_create_params import Request + + lote = self._client.messages.batches.create( + requests=[ + Request( + custom_id=p["custom_id"], + params=MessageCreateParamsNonStreaming(**p["params"]), + ) + for p in pedidos + ] + ) + return lote.id + + def estado(self, batch_id: str) -> dict[str, Any]: + lote = self._client.messages.batches.retrieve(batch_id) + conteos = lote.request_counts + return { + "processing_status": lote.processing_status, + "succeeded": conteos.succeeded, + "errored": conteos.errored, + "processing": conteos.processing, + "canceled": conteos.canceled, + "expired": conteos.expired, + } + + def resultados(self, batch_id: str) -> Iterator[dict[str, Any]]: + for item in self._client.messages.batches.results(batch_id): + yield _normalizar_resultado(item) + + def cancelar(self, batch_id: str) -> None: + self._client.messages.batches.cancel(batch_id) + + +def _normalizar_resultado(item: Any) -> dict[str, Any]: + """Aplana un resultado del SDK a un dict propio. + + Se hace acá para que el resto del código no dependa de los tipos del SDK, y + para que lo que se escribe a disco sea estable aunque el SDK cambie. + """ + tipo = item.result.type + salida: dict[str, Any] = {"custom_id": item.custom_id, "type": tipo} + + if tipo == "succeeded": + mensaje = item.result.message + texto = "".join(b.text for b in mensaje.content if b.type == "text") + uso = mensaje.usage + salida["text"] = texto + salida["stop_reason"] = mensaje.stop_reason + salida["usage"] = { + "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, + } + elif tipo == "errored": + salida["error_type"] = item.result.error.type + salida["error"] = str(getattr(item.result.error, "message", "")) + return salida + + +# --- estado en disco ------------------------------------------------------- + + +@dataclass +class Trabajo: + """Un trabajo de generación, con su estado persistido. + + El directorio es la fuente de verdad; el proceso no guarda nada en memoria + que no pueda reconstruir leyéndolo. + """ + + nombre: str + raiz: Path + + @property + def dir(self) -> Path: + return self.raiz / self.nombre + + @property + def peticiones_path(self) -> Path: + return self.dir / "requests.jsonl" + + @property + def lotes_path(self) -> Path: + return self.dir / "batches.json" + + @property + def resultados_path(self) -> Path: + return self.dir / "results.jsonl" + + @property + def errores_path(self) -> Path: + return self.dir / "errors.jsonl" + + def crear_dir(self) -> None: + self.dir.mkdir(parents=True, exist_ok=True) + + # -- peticiones -- + + def guardar_peticiones(self, peticiones: Iterable[Peticion]) -> list[Peticion]: + """Escribe el manifiesto, descartando duplicados por `custom_id`.""" + self.crear_dir() + vistos: set[str] = set() + unicas: list[Peticion] = [] + for p in peticiones: + if p.custom_id in vistos: + continue + vistos.add(p.custom_id) + unicas.append(p) + + with self.peticiones_path.open("w", encoding="utf-8") as fh: + for p in unicas: + fila = asdict(p) | {"custom_id": p.custom_id} + fh.write(json.dumps(fila, ensure_ascii=False) + "\n") + return unicas + + 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}" + ) + peticiones = [] + with self.peticiones_path.open(encoding="utf-8") as fh: + for linea in fh: + fila = json.loads(linea) + fila.pop("custom_id", None) + peticiones.append(Peticion(**fila)) + return peticiones + + # -- lotes -- + + def leer_lotes(self) -> list[dict[str, Any]]: + if not self.lotes_path.is_file(): + return [] + return json.loads(self.lotes_path.read_text()) + + def guardar_lotes(self, lotes: list[dict[str, Any]]) -> None: + self.crear_dir() + tmp = self.lotes_path.with_suffix(".json.tmp") + tmp.write_text(json.dumps(lotes, indent=2, ensure_ascii=False)) + tmp.replace(self.lotes_path) + + # -- resultados -- + + def ids_resueltos(self) -> set[str]: + """Los `custom_id` que ya tienen resultado o error registrado.""" + resueltos: set[str] = set() + for path in (self.resultados_path, self.errores_path): + if not path.is_file(): + continue + with path.open(encoding="utf-8") as fh: + for linea in fh: + if linea.strip(): + resueltos.add(json.loads(linea)["custom_id"]) + return resueltos + + def anexar(self, path: Path, filas: Iterable[dict[str, Any]]) -> int: + self.crear_dir() + n = 0 + with path.open("a", encoding="utf-8") as fh: + for fila in filas: + fh.write(json.dumps(fila, ensure_ascii=False) + "\n") + n += 1 + return n + + +# --- el runner ------------------------------------------------------------- + + +def submit( + trabajo: Trabajo, + peticiones: list[Peticion], + cliente: ClienteAnthropic | ClienteLotes, + config: DistillConfig, + *, + solo_pendientes: bool = True, +) -> list[str]: + """Escribe el manifiesto y envía los lotes. Devuelve los batch_id. + + Reanudable: los pedidos ya resueltos se saltan, y los lotes ya enviados no + se reenvían. Volver a correr esto tras una caída no cuesta nada de más. + """ + unicas = trabajo.guardar_peticiones(peticiones) + + if solo_pendientes: + resueltos = trabajo.ids_resueltos() + pendientes = [p for p in unicas if p.custom_id not in resueltos] + else: + pendientes = unicas + + 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: + return [lote["batch_id"] for lote in trabajo.leer_lotes()] + + construir = getattr(cliente, "construir_params", None) + lotes = trabajo.leer_lotes() + tamano = config.requests_per_batch + + for inicio in range(0, len(pendientes), tamano): + trozo = pendientes[inicio : inicio + tamano] + pedidos = [ + { + "custom_id": p.custom_id, + "params": construir(p) if construir else {}, + } + for p in trozo + ] + batch_id = cliente.crear(pedidos) + lotes.append( + { + "batch_id": batch_id, + "custom_ids": [p.custom_id for p in trozo], + "submitted_at": time.time(), + } + ) + # Se persiste después de cada lote, no al final: si el proceso muere en + # el segundo envío, el primero no queda huérfano y sin registrar. + trabajo.guardar_lotes(lotes) + + return [lote["batch_id"] for lote in lotes] + + +def status( + trabajo: Trabajo, cliente: ClienteLotes +) -> list[dict[str, Any]]: + """Estado de cada lote del trabajo.""" + estados = [] + for lote in trabajo.leer_lotes(): + info = cliente.estado(lote["batch_id"]) + estados.append({"batch_id": lote["batch_id"], **info}) + return estados + + +def collect( + trabajo: Trabajo, + cliente: ClienteLotes, + config: DistillConfig, + *, + esperar: bool = True, + ahora=time.monotonic, + dormir=time.sleep, +) -> dict[str, int]: + """Espera a que terminen los lotes y recoge los resultados. + + Los éxitos van a `results.jsonl` y los fallos a `errors.jsonl`, ambos + indexados por `custom_id`. Recoger dos veces no duplica: lo ya registrado + se saltea. + """ + lotes = trabajo.leer_lotes() + if not lotes: + raise DistillError( + f"el trabajo '{trabajo.nombre}' no tiene lotes enviados; corré `submit` primero" + ) + + limite = ahora() + config.max_wait_hours * 3600 + pendientes = {lote["batch_id"] for lote in lotes} + + resumen = {"ok": 0, "error": 0, "ya_estaban": 0} + + while pendientes: + terminados = set() + for batch_id in sorted(pendientes): + info = cliente.estado(batch_id) + if info["processing_status"] == "ended": + terminados.add(batch_id) + + for batch_id in terminados: + resueltos = trabajo.ids_resueltos() + ok, err, repetidos = [], [], 0 + for resultado in cliente.resultados(batch_id): + if resultado["custom_id"] in resueltos: + repetidos += 1 + continue + (ok if resultado["type"] == "succeeded" else err).append(resultado) + + resumen["ok"] += trabajo.anexar(trabajo.resultados_path, ok) + resumen["error"] += trabajo.anexar(trabajo.errores_path, err) + resumen["ya_estaban"] += repetidos + + pendientes -= terminados + + if not pendientes or not esperar: + break + if ahora() > limite: + raise DistillError( + f"se superó max_wait_hours={config.max_wait_hours} con lotes sin " + f"terminar: {sorted(pendientes)}. El trabajo sigue en el servidor; " + "volvé a correr `collect` para retomarlo." + ) + dormir(config.poll_interval_seconds) + + return resumen + + +# --- costo ----------------------------------------------------------------- + + +def resumen_de_uso(trabajo: Trabajo, config: DistillConfig) -> dict[str, Any]: + """Tokens consumidos y costo estimado, leídos de los resultados en disco.""" + entrada = salida = cache_read = cache_write = 0 + n = 0 + if trabajo.resultados_path.is_file(): + with trabajo.resultados_path.open(encoding="utf-8") as fh: + for linea in fh: + uso = json.loads(linea).get("usage") + if not uso: + continue + n += 1 + entrada += uso["input_tokens"] + salida += uso["output_tokens"] + cache_read += uso.get("cache_read_input_tokens", 0) + cache_write += uso.get("cache_creation_input_tokens", 0) + + # Los tokens leídos de caché cuestan ~0,1x y los escritos ~1,25x; el + # descuento del lote se aplica sobre todo. Es una estimación para saber en + # qué orden de magnitud estamos, no una factura. + entrada_efectiva = entrada + cache_read * 0.1 + cache_write * 1.25 + usd = ( + entrada_efectiva / 1e6 * config.usd_per_mtok_input + + salida / 1e6 * config.usd_per_mtok_output + ) * config.batch_discount + + return { + "resultados": n, + "input_tokens": entrada, + "output_tokens": salida, + "cache_read_input_tokens": cache_read, + "cache_creation_input_tokens": cache_write, + "usd_estimado": round(usd, 2), + } + + +# --- CLI ------------------------------------------------------------------- + + +def leer_especificaciones(path: str | Path) -> list[Peticion]: + """Lee un JSONL de especificaciones: {system, user, schema?, meta?}.""" + path = Path(path) + if not path.is_file(): + raise DistillError(f"no existe el archivo de peticiones: {path}") + peticiones = [] + for numero, linea in enumerate(path.read_text(encoding="utf-8").splitlines(), 1): + if not linea.strip(): + continue + try: + fila = json.loads(linea) + peticiones.append( + Peticion( + system=fila["system"], + user=fila["user"], + schema=fila.get("schema"), + meta=fila.get("meta", {}), + ) + ) + except (json.JSONDecodeError, KeyError) as exc: + raise DistillError(f"{path}:{numero}: petición inválida ({exc})") from None + return peticiones + + +def main() -> int: + import argparse + + from enlace.config.load import ConfigError, load_distill_config + + parser = argparse.ArgumentParser( + prog="python -m enlace.data.distill", + description="Generación de datos sintéticos con la Batch API.", + ) + parser.add_argument("accion", choices=["submit", "status", "collect", "usage"]) + parser.add_argument("peticiones", nargs="?", help="JSONL de peticiones (submit)") + parser.add_argument("--name", required=True, help="nombre del trabajo") + parser.add_argument("--config", default=None, help="ruta a distill.yaml") + parser.add_argument( + "--retry-failed", + action="store_true", + help="reenvía también las peticiones que fallaron", + ) + parser.add_argument( + "--no-wait", action="store_true", help="no esperar a que terminen los lotes" + ) + args = parser.parse_args() + + try: + config = load_distill_config(args.config) + trabajo = Trabajo(nombre=args.name, raiz=Path(config.output_dir)) + + if args.accion == "usage": + print(json.dumps(resumen_de_uso(trabajo, config), indent=2)) + return 0 + + cliente = ClienteAnthropic(config) + + if args.accion == "submit": + if not args.peticiones: + parser.error("submit necesita el JSONL de peticiones") + if args.retry_failed and trabajo.errores_path.is_file(): + # Los fallos se reenvían borrando su registro: `submit` saltea lo + # ya resuelto, así que basta con que dejen de estarlo. + trabajo.errores_path.unlink() + peticiones = leer_especificaciones(args.peticiones) + ids = submit(trabajo, peticiones, cliente, config) + print(f"[enlace] {len(peticiones)} peticiones en {len(ids)} lote(s)") + for batch_id in ids: + print(f" {batch_id}") + + elif args.accion == "status": + for info in status(trabajo, cliente): + print(json.dumps(info, ensure_ascii=False)) + + elif args.accion == "collect": + resumen = collect(trabajo, cliente, config, esperar=not args.no_wait) + print(f"[enlace] {resumen}") + print(json.dumps(resumen_de_uso(trabajo, config), indent=2)) + + except (DistillError, ConfigError) as exc: + print(f"[enlace] {exc}", file=__import__("sys").stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/enlace/data/loaders.py b/enlace/data/loaders.py new file mode 100644 index 0000000..d9045bf --- /dev/null +++ b/enlace/data/loaders.py @@ -0,0 +1,215 @@ +"""Cargadores de datos con reanudación exacta. + +El requisito no negociable: al reanudar desde un checkpoint, el entrenamiento +tiene que ver exactamente los mismos lotes que habría visto sin la interrupción. +Un loader que reanuda desde el principio del corpus reentrena sobre datos ya +vistos y arruina la corrida en silencio — es el bug más caro de descubrir tarde, +así que el estado del loader se guarda en el checkpoint como cualquier otro. +""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any, Protocol + +import numpy as np +import torch +from torch import Tensor + + +class TokenStream(Protocol): + """Interfaz común: el entrenador no sabe si lee bytes o shards.""" + + vocab_size: int + + def next_batch(self, split: str) -> tuple[Tensor, Tensor]: ... + def state_dict(self) -> dict[str, Any]: ... + def load_state_dict(self, state: dict[str, Any]) -> None: ... + + +class ByteStream: + """Corpus a nivel de bytes, para el smoke test de la Etapa 0. + + Se trabaja sobre bytes UTF-8 y no sobre caracteres: el vocabulario es + exactamente 256 sin depender del texto, y los acentos y la ñ se representan + como secuencias multi-byte igual que en el tokenizer BPE byte-level que usa + el modelo real. El smoke test ejercita así la misma clase de entrada. + """ + + vocab_size = 256 + + def __init__( + self, + text_path: str | Path, + batch_size: int, + seq_len: int, + val_fraction: float, + seed: int, + device: torch.device, + ) -> None: + path = Path(text_path) + if not path.is_file(): + raise FileNotFoundError( + f"no existe {path}. El smoke test necesita un texto en español; " + "ver scripts/prepare_smoke_data.sh" + ) + raw = np.frombuffer(path.read_bytes(), dtype=np.uint8) + if len(raw) < seq_len * 4: + raise ValueError( + f"{path} tiene {len(raw)} bytes: muy poco para seq_len={seq_len}." + ) + + split_at = int(len(raw) * (1.0 - val_fraction)) + self._data = { + "train": torch.from_numpy(raw[:split_at].astype(np.int64)), + "val": torch.from_numpy(raw[split_at:].astype(np.int64)), + } + self.batch_size = batch_size + self.seq_len = seq_len + self.device = device + # Un generador por split: así evaluar no altera la secuencia de lotes de + # entrenamiento, y la reanudación es exacta aunque cambie eval_every. + self._gens = { + "train": torch.Generator().manual_seed(seed), + "val": torch.Generator().manual_seed(seed + 1), + } + + def next_batch(self, split: str) -> tuple[Tensor, Tensor]: + data = self._data[split] + high = len(data) - self.seq_len - 1 + idx = torch.randint(high, (self.batch_size,), generator=self._gens[split]) + x = torch.stack([data[i : i + self.seq_len] for i in idx]) + y = torch.stack([data[i + 1 : i + 1 + self.seq_len] for i in idx]) + return x.to(self.device), y.to(self.device) + + def state_dict(self) -> dict[str, Any]: + return {k: g.get_state() for k, g in self._gens.items()} + + def load_state_dict(self, state: dict[str, Any]) -> None: + for k, g in self._gens.items(): + if k in state: + g.set_state(state[k]) + + @staticmethod + def decode(ids: list[int]) -> str: + # errors="replace" porque el modelo puede generar secuencias de bytes + # que no son UTF-8 válido, sobre todo temprano en el entrenamiento. + return bytes(i % 256 for i in ids).decode("utf-8", errors="replace") + + +class ShardStream: + """Corpus tokenizado en shards uint16 (Etapa 1 del plan). + + Los shards se leen en orden y de forma circular, con un puntero global de + posición. Guardar ese entero es todo lo que hace falta para reanudar exacto, + y a diferencia de un muestreo aleatorio garantiza que cada token se ve una + vez por época antes de repetir ninguno. + """ + + def __init__( + self, + shards_dir: str | Path, + batch_size: int, + seq_len: int, + val_fraction: float, + device: torch.device, + ) -> None: + directory = Path(shards_dir) + index_path = directory / "index.json" + if not index_path.is_file(): + raise FileNotFoundError( + f"no existe {index_path}. Los shards se generan con " + "scripts/prepare_data.sh (Etapa 1)." + ) + index = json.loads(index_path.read_text()) + self.vocab_size: int = index["vocab_size"] + + files = [directory / entry["file"] for entry in index["shards"]] + missing = [f for f in files if not f.is_file()] + if missing: + raise FileNotFoundError(f"faltan shards declarados en index.json: {missing}") + + # np.memmap: los shards pueden sumar decenas de GB y el servidor no + # necesariamente tiene RAM para tenerlos cargados. + arrays = [np.memmap(f, dtype=np.uint16, mode="r") for f in files] + total = sum(len(a) for a in arrays) + split_at = int(total * (1.0 - val_fraction)) + + self._arrays = arrays + self._lengths = [len(a) for a in arrays] + self._bounds = {"train": (0, split_at), "val": (split_at, total)} + self._pos = {"train": 0, "val": 0} + + self.batch_size = batch_size + self.seq_len = seq_len + self.device = device + + needed = batch_size * seq_len + 1 + for split, (lo, hi) in self._bounds.items(): + if hi - lo < needed: + raise ValueError( + f"el split '{split}' tiene {hi - lo} tokens, menos que los " + f"{needed} que exige un lote de {batch_size}x{seq_len}." + ) + + def _read(self, start: int, count: int) -> np.ndarray: + """Lee `count` tokens desde la posición global `start`, cruzando shards.""" + out = np.empty(count, dtype=np.int64) + written = 0 + offset = start + # Ubica el shard que contiene `offset`. + shard_idx = 0 + while offset >= self._lengths[shard_idx]: + offset -= self._lengths[shard_idx] + shard_idx += 1 + while written < count: + array = self._arrays[shard_idx] + take = min(count - written, len(array) - offset) + out[written : written + take] = array[offset : offset + take] + written += take + offset = 0 + shard_idx = (shard_idx + 1) % len(self._arrays) + return out + + def next_batch(self, split: str) -> tuple[Tensor, Tensor]: + lo, hi = self._bounds[split] + span = hi - lo + needed = self.batch_size * self.seq_len + 1 + pos = self._pos[split] + if pos + needed > span: + pos = 0 # fin de época: se vuelve al principio del split + chunk = self._read(lo + pos, needed) + self._pos[split] = pos + self.batch_size * self.seq_len + + tokens = torch.from_numpy(chunk) + x = tokens[:-1].view(self.batch_size, self.seq_len) + y = tokens[1:].view(self.batch_size, self.seq_len) + return x.to(self.device), y.to(self.device) + + def state_dict(self) -> dict[str, Any]: + return dict(self._pos) + + def load_state_dict(self, state: dict[str, Any]) -> None: + self._pos.update({k: int(v) for k, v in state.items()}) + + +def build_stream(cfg, device: torch.device) -> TokenStream: + """Construye el stream que corresponda a `data.source`.""" + data, hw, train = cfg.data, cfg.hardware, cfg.train + if data.source == "chars": + return ByteStream( + text_path=data.text_path, + batch_size=hw.micro_batch_size, + seq_len=cfg.model.seq_len, + val_fraction=data.val_fraction, + seed=train.seed, + device=device, + ) + return ShardStream( + shards_dir=data.shards_dir, + batch_size=hw.micro_batch_size, + seq_len=cfg.model.seq_len, + val_fraction=data.val_fraction, + device=device, + )