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:
2026-07-27 23:27:40 -03:00
parent 8805b0541b
commit d0a05cdd97
10 changed files with 882 additions and 4 deletions
+9 -4
View File
@@ -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
+5
View File
@@ -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
+36
View File
@@ -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
+5
View File
@@ -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
+7
View File
@@ -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
+13
View File
@@ -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
+21
View File
@@ -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
+5
View File
@@ -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"]
+566
View File
@@ -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())
+215
View File
@@ -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,
)