Fase 5: 21_quantize_nvfp4.py - cargar ultrachat_200k en modo streaming

3 intentos seguidos de calibracion mezclada (256 ultrachat + 256 propias)
crashearon con el mismo CUDA OOM reproducible, siempre en el mismo punto
exacto (setup interno de oneshot(): trace_subgraphs/disable_lm_head), tanto
con MAX_SEQUENCE_LENGTH=8192 como =2048 -- descartando el largo de secuencia
como causa. La unica variable real frente a los intentos que SI funcionaron
(256 muestras solo propias) es la inclusion de ultrachat_200k.

Causa raiz identificada: load_dataset(..., split="train_sft") sin streaming
materializa el split completo (~208k ejemplos) como Arrow local, y ademas
genera las 4 splits del repo (~2.9GB en disco). En este hardware (GB10,
memoria unificada CPU/GPU) ese cache extra parece ser suficiente para
empujar el proceso sobre el limite justo en el momento de mayor presion de
memoria del setup de oneshot(). Fix: cargar con streaming=True + shuffle de
buffer + take(n), que solo trae los N ejemplos necesarios sin materializar
el dataset completo -- probado de forma aislada (256 ejemplos en ~12s, sin
crecimiento de cache en disco).
This commit is contained in:
2026-07-30 04:00:07 +00:00
parent 78d9b0d90d
commit b246f8d97a
+13 -5
View File
@@ -141,11 +141,19 @@ def load_train_examples(n):
def load_ultrachat_examples(n):
from datasets import load_dataset
print(f"[INFO] cargando {n} muestras de {ULTRACHAT_DATASET} (split={ULTRACHAT_SPLIT})")
ds = load_dataset(ULTRACHAT_DATASET, split=ULTRACHAT_SPLIT)
ds = ds.shuffle(seed=CALIBRATION_SEED).select(range(n))
examples = [{"messages": row["messages"]} for row in ds]
print(f"[INFO] {len(examples)} muestras de {ULTRACHAT_DATASET} cargadas")
# streaming=True: el split train_sft completo tiene ~208k ejemplos (~2.9GB
# materializados como Arrow por load_dataset sin streaming, las 4 splits del
# repo se generan igual). En este hardware (GB10, memoria unificada CPU/GPU)
# ese cache extra resulto ser suficiente para tirar un CUDA OOM reproducible
# durante el setup de oneshot() (trace_subgraphs/disable_lm_head), incluso
# truncando las secuencias a 2048 tokens -- el problema no era el largo de
# secuencia sino la memoria consumida por materializar el dataset completo.
# Con streaming solo se bajan los ~n ejemplos necesarios, sin cache local.
print(f"[INFO] cargando {n} muestras de {ULTRACHAT_DATASET} (split={ULTRACHAT_SPLIT}, streaming)")
ds = load_dataset(ULTRACHAT_DATASET, split=ULTRACHAT_SPLIT, streaming=True)
ds = ds.shuffle(seed=CALIBRATION_SEED, buffer_size=10_000)
examples = [{"messages": row["messages"]} for row in ds.take(n)]
print(f"[INFO] {len(examples)} muestras de {ULTRACHAT_DATASET} cargadas (streaming, sin materializar el dataset completo)")
return examples