Anclar las rutas de .gitignore al raíz del repo
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 <noreply@anthropic.com>
This commit is contained in:
@@ -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"]
|
||||
@@ -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 <trabajo.jsonl> --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())
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user