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:
+9
-4
@@ -1,8 +1,13 @@
|
|||||||
# Datos, pesos y bases del sistema: viven en el servidor, nunca en git.
|
# Datos, pesos y bases del sistema: viven en el servidor, nunca en git.
|
||||||
data/
|
#
|
||||||
checkpoints/
|
# Las rutas van ancladas con "/" inicial a propósito. Sin anclar, "data/"
|
||||||
snapshots/
|
# excluye CUALQUIER directorio llamado data en cualquier nivel — incluidos
|
||||||
runs/
|
# 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
|
# Secretos
|
||||||
.env
|
.env
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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