diff --git a/scripts/21_quantize_nvfp4.py b/scripts/21_quantize_nvfp4.py index 09fe2de..d30968b 100644 --- a/scripts/21_quantize_nvfp4.py +++ b/scripts/21_quantize_nvfp4.py @@ -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