90ac2a6582
Base del proyecto ENLACE: un modelo de lenguaje propio entrenado desde cero, en español, para asistencia general y familiar. El plan completo está en docs/PLAN.md. Esta etapa establece el andamiaje y lo verifica de punta a punta: - Configuración por capas (hardware × model × train × data) validada con pydantic. Ningún hiperparámetro vive en el código y una config inválida falla al arrancar, no a las tres horas de entrenamiento. - Perfiles de hardware que aíslan el salto de GPU: la RTX 2060 (Turing) no soporta bfloat16 ni FlashAttention-2, así que entrena en float16 con GradScaler y backend mem_efficient; el perfil de la 5090 ya está escrito. backends.py valida el perfil contra la GPU real antes de empezar. - Transformer decoder-only estilo Llama: RMSNorm, SwiGLU, RoPE, GQA, embeddings atados, QK-norm y z-loss. Los dos últimos son lo que mantiene estable el entrenamiento en float16. - Entrenador con schedule WSD, acumulación de gradiente, precisión mixta, checkpointing atómico y reanudación exacta. - Cargadores de datos con estado serializable: bytes para el smoke test y shards uint16 para el corpus real. 48 tests, entre ellos el crítico: reanudar desde un checkpoint reproduce los pesos de una corrida ininterrumpida, parámetro por parámetro. Verificado en CPU: 300 pasos sobre texto en español, loss 3.07 -> 1.63. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
144 lines
4.4 KiB
Python
144 lines
4.4 KiB
Python
"""Cargadores de datos: determinismo, reanudación y recorrido de shards."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from enlace.data.loaders import ByteStream, ShardStream
|
|
|
|
|
|
def _byte_stream(texto, seed=1337):
|
|
return ByteStream(
|
|
text_path=texto,
|
|
batch_size=2,
|
|
seq_len=16,
|
|
val_fraction=0.1,
|
|
seed=seed,
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
|
|
def test_los_targets_son_la_entrada_desplazada_un_token(texto_es):
|
|
stream = _byte_stream(texto_es)
|
|
x, y = stream.next_batch("train")
|
|
assert x.shape == y.shape == (2, 16)
|
|
assert torch.equal(x[:, 1:], y[:, :-1])
|
|
|
|
|
|
def test_la_misma_semilla_da_los_mismos_lotes(texto_es):
|
|
a = _byte_stream(texto_es)
|
|
b = _byte_stream(texto_es)
|
|
for _ in range(3):
|
|
xa, _ = a.next_batch("train")
|
|
xb, _ = b.next_batch("train")
|
|
assert torch.equal(xa, xb)
|
|
|
|
|
|
def test_restaurar_el_estado_continua_la_misma_secuencia(texto_es):
|
|
a = _byte_stream(texto_es)
|
|
for _ in range(5):
|
|
a.next_batch("train")
|
|
estado = a.state_dict()
|
|
esperado = [a.next_batch("train")[0] for _ in range(3)]
|
|
|
|
b = _byte_stream(texto_es)
|
|
b.load_state_dict(estado)
|
|
obtenido = [b.next_batch("train")[0] for _ in range(3)]
|
|
|
|
for e, o in zip(esperado, obtenido):
|
|
assert torch.equal(e, o)
|
|
|
|
|
|
def test_evaluar_no_altera_la_secuencia_de_entrenamiento(texto_es):
|
|
"""Generadores separados por split: cambiar eval_every no debe mover los
|
|
lotes de entrenamiento, o dos corridas dejarían de ser comparables."""
|
|
a = _byte_stream(texto_es)
|
|
esperado = [a.next_batch("train")[0] for _ in range(3)]
|
|
|
|
b = _byte_stream(texto_es)
|
|
lotes = []
|
|
for _ in range(3):
|
|
b.next_batch("val")
|
|
lotes.append(b.next_batch("train")[0])
|
|
|
|
for e, o in zip(esperado, lotes):
|
|
assert torch.equal(e, o)
|
|
|
|
|
|
def test_el_vocabulario_de_bytes_es_siempre_256(texto_es):
|
|
assert _byte_stream(texto_es).vocab_size == 256
|
|
|
|
|
|
def test_decodifica_utf8_con_acentos():
|
|
assert ByteStream.decode(list("El niño está acá.".encode())) == "El niño está acá."
|
|
|
|
|
|
def test_texto_inexistente_da_un_error_util(tmp_path):
|
|
with pytest.raises(FileNotFoundError, match="prepare_smoke_data"):
|
|
_byte_stream(tmp_path / "no-existe.txt")
|
|
|
|
|
|
def _shard_stream(directory, batch_size=2, seq_len=8):
|
|
return ShardStream(
|
|
shards_dir=directory,
|
|
batch_size=batch_size,
|
|
seq_len=seq_len,
|
|
val_fraction=0.1,
|
|
device=torch.device("cpu"),
|
|
)
|
|
|
|
|
|
def test_los_shards_se_recorren_en_orden(shards_dir):
|
|
stream = _shard_stream(shards_dir)
|
|
x, y = stream.next_batch("train")
|
|
# El fixture escribe tokens consecutivos: 0, 1, 2, ...
|
|
assert x[0].tolist() == list(range(8))
|
|
assert y[0].tolist() == list(range(1, 9))
|
|
assert x[1].tolist() == list(range(8, 16))
|
|
|
|
x2, _ = stream.next_batch("train")
|
|
assert x2[0].tolist() == list(range(16, 24))
|
|
|
|
|
|
def test_la_lectura_cruza_el_limite_entre_shards(shards_dir):
|
|
"""Un lote que empieza cerca del final de un shard tiene que continuar en el
|
|
siguiente sin saltarse ni repetir tokens."""
|
|
stream = _shard_stream(shards_dir, batch_size=1, seq_len=8)
|
|
stream.load_state_dict({"train": 4090})
|
|
x, _ = stream.next_batch("train")
|
|
assert x[0].tolist() == list(range(4090, 4098)) # cruza de shard 0 a shard 1
|
|
|
|
|
|
def test_restaurar_la_posicion_reanuda_donde_iba(shards_dir):
|
|
a = _shard_stream(shards_dir)
|
|
for _ in range(4):
|
|
a.next_batch("train")
|
|
estado = a.state_dict()
|
|
esperado, _ = a.next_batch("train")
|
|
|
|
b = _shard_stream(shards_dir)
|
|
b.load_state_dict(estado)
|
|
obtenido, _ = b.next_batch("train")
|
|
assert torch.equal(esperado, obtenido)
|
|
|
|
|
|
def test_al_terminar_la_epoca_vuelve_al_principio(shards_dir):
|
|
stream = _shard_stream(shards_dir, batch_size=1, seq_len=8)
|
|
primero, _ = stream.next_batch("train")
|
|
# El split de train son 7372 tokens; posicionarse casi al final.
|
|
stream.load_state_dict({"train": 7370})
|
|
stream.next_batch("train")
|
|
reiniciado, _ = stream.next_batch("train")
|
|
assert reiniciado[0, 0].item() == 8 # segundo lote de la época nueva
|
|
|
|
|
|
def test_shards_faltantes_dan_un_error_util(tmp_path):
|
|
with pytest.raises(FileNotFoundError, match="prepare_data"):
|
|
_shard_stream(tmp_path)
|
|
|
|
|
|
def test_corpus_demasiado_chico_para_el_lote(shards_dir):
|
|
with pytest.raises(ValueError, match="menos que los"):
|
|
_shard_stream(shards_dir, batch_size=64, seq_len=1024)
|