diff --git a/scripts/21_quantize_nvfp4.py b/scripts/21_quantize_nvfp4.py index 27be137..4a70ecb 100644 --- a/scripts/21_quantize_nvfp4.py +++ b/scripts/21_quantize_nvfp4.py @@ -24,17 +24,24 @@ Algoritmo: truncados a MAX_SEQUENCE_LENGTH. 4. oneshot(..., moe_calibrate_all_experts=True) -- obligatorio, si no la mayoria de los 256 expertos ruteados quedan sin calibrar. - 5. model.save_pretrained(OUTPUT_PATH) (separa visual a model_visual.safetensors - nativamente), processor.save_pretrained(OUTPUT_PATH) (copia el - chat_template.jinja de produccion, no el de masking), + 5. model.save_pretrained(OUTPUT_PATH), processor.save_pretrained(OUTPUT_PATH) + (copia el chat_template.jinja de produccion, no el de masking), save_mtp_tensors_to_checkpoint(source_model=MODEL_PATH, dest_dir=OUTPUT_PATH) (copia mtp.* directo del checkpoint origen a model_mtp.safetensors, ya que - esta clase no los instancia). + esta clase no los instancia). NOTA: a diferencia del checkpoint de referencia + de RedHatAI (que trae vision en su propio model_visual.safetensors), en esta + version de transformers save_pretrained() escribe lenguaje+vision juntos en + el/los shard(s) de model.safetensors -- comportamiento igualmente valido (el + index.json mapea cada tensor a su shard real), verificado por conteo de + tensores en vez de por nombre de archivo. 6. Verificacion automatica (aborta si algo no cuadra): quantization_config.format - == nvfp4-pack-quantized; model_mtp.safetensors y model_visual.safetensors - presentes con conteos de tensores razonables; muestra de tensores - cuantizados decodifica sin NaN/Inf; chat_template.jinja NO identico al de - masking de training y SI identico al de MODEL_PATH. + == nvfp4-pack-quantized; model_mtp.safetensors presente y conteo de tensores + de vision razonable (via el index, sin asumir un archivo separado); muestra de + tensores cuantizados decodifica sin NaN/Inf; chat_template.jinja NO identico + al de masking de training y SI identico al de MODEL_PATH. + +Soporta --verify-only (o env var VERIFY_ONLY=1) para re-correr solo las +verificaciones sobre un OUTPUT_PATH ya generado, sin repetir la calibracion. """ import json import os @@ -199,24 +206,27 @@ def verify_quantization_config(): def verify_mtp_and_visual_shards(): mtp_path = OUTPUT_PATH / "model_mtp.safetensors" - visual_path = OUTPUT_PATH / "model_visual.safetensors" if not mtp_path.exists(): raise AssertionError(f"falta {mtp_path} -- los tensores MTP no se reinyectaron, --speculative-config no arrancara") - if not visual_path.exists(): - raise AssertionError(f"falta {visual_path} -- los pesos de vision no se preservaron") with safe_open(str(mtp_path), framework="pt") as f: mtp_count = len(f.keys()) - with safe_open(str(visual_path), framework="pt") as f: - visual_count = len(f.keys()) - print(f"[INFO] model_mtp.safetensors: {mtp_count} tensores") - print(f"[INFO] model_visual.safetensors: {visual_count} tensores") - - # Referencia (Fase 3/4): ~19 tensores MTP, ~333 tensores de vision. Rango amplio - # a proposito -- lo que importa es que no esten vacios ni truncados a un puñado. + # Referencia (Fase 3/4): ~19 tensores MTP. Rango amplio a proposito -- lo que + # importa es que no este vacio ni truncado a un puñado. if not (10 <= mtp_count <= 40): raise AssertionError(f"conteo de tensores MTP fuera de rango razonable: {mtp_count} (esperado ~19)") + + # A diferencia del checkpoint de referencia de RedHatAI (que trae vision en su + # propio model_visual.safetensors), esta version de transformers + # (Qwen3_5MoeForConditionalGeneration.save_pretrained) escribe lenguaje+vision + # juntos en el/los shard(s) de model.safetensors -- comportamiento igualmente + # valido (el index.json mapea cada tensor a su shard real sin importar el nombre + # de archivo), asi que se verifica por conteo de tensores via el index en vez de + # exigir un archivo separado. + all_names = iter_output_tensor_names() + visual_count = sum(1 for name, _shard in all_names if ".visual." in name or name.startswith("visual.")) + print(f"[INFO] tensores de vision encontrados (en los shards de model.safetensors): {visual_count}") if not (250 <= visual_count <= 450): raise AssertionError(f"conteo de tensores de vision fuera de rango razonable: {visual_count} (esperado ~333)") @@ -294,54 +304,72 @@ def verify_chat_template(): ) +def parse_args(): + import argparse + + parser = argparse.ArgumentParser() + parser.add_argument( + "--verify-only", + action="store_true", + default=os.environ.get("VERIFY_ONLY", "") not in ("", "0", "false", "False"), + help="saltar calibracion/guardado, solo re-correr las verificaciones sobre OUTPUT_PATH ya existente", + ) + return parser.parse_args() + + def main(): + args = parse_args() print(f"[INFO] MODEL_PATH={MODEL_PATH}") print(f"[INFO] OUTPUT_PATH={OUTPUT_PATH}") - print(f"[INFO] TRAIN_DATA_PATH={TRAIN_DATA_PATH}") - print(f"[INFO] NUM_CALIBRATION_SAMPLES={NUM_CALIBRATION_SAMPLES} MAX_SEQUENCE_LENGTH={MAX_SEQUENCE_LENGTH}") - print(f"[INFO] cargando processor desde {MODEL_PATH}") - processor = AutoProcessor.from_pretrained(str(MODEL_PATH), trust_remote_code=True) - tokenizer = processor.tokenizer + if args.verify_only: + print("[INFO] --verify-only: saltando calibracion/guardado, solo verificando OUTPUT_PATH existente") + else: + print(f"[INFO] TRAIN_DATA_PATH={TRAIN_DATA_PATH}") + print(f"[INFO] NUM_CALIBRATION_SAMPLES={NUM_CALIBRATION_SAMPLES} MAX_SEQUENCE_LENGTH={MAX_SEQUENCE_LENGTH}") - calibration_dataset = load_calibration_dataset(tokenizer) + print(f"[INFO] cargando processor desde {MODEL_PATH}") + processor = AutoProcessor.from_pretrained(str(MODEL_PATH), trust_remote_code=True) + tokenizer = processor.tokenizer - print(f"[INFO] cargando modelo desde {MODEL_PATH} (dtype=auto)") - t_load = time.time() - model = Qwen3_5MoeForConditionalGeneration.from_pretrained( - str(MODEL_PATH), dtype="auto", trust_remote_code=True - ) - print(f"[INFO] modelo cargado en {time.time() - t_load:.1f}s") - report_memory("post-load") + calibration_dataset = load_calibration_dataset(tokenizer) - recipe = QuantizationModifier(targets="Linear", scheme="NVFP4", ignore=QUANTIZATION_IGNORE) + print(f"[INFO] cargando modelo desde {MODEL_PATH} (dtype=auto)") + t_load = time.time() + model = Qwen3_5MoeForConditionalGeneration.from_pretrained( + str(MODEL_PATH), dtype="auto", trust_remote_code=True + ) + print(f"[INFO] modelo cargado en {time.time() - t_load:.1f}s") + report_memory("post-load") - data_collator = CalibrationDataCollator(tokenizer.pad_token_id or tokenizer.eos_token_id) + recipe = QuantizationModifier(targets="Linear", scheme="NVFP4", ignore=QUANTIZATION_IGNORE) - print("[INFO] arrancando oneshot() -- calibracion NVFP4 con moe_calibrate_all_experts=True") - t_quant = time.time() - oneshot( - model=model, - recipe=recipe, - dataset=calibration_dataset, - max_seq_length=MAX_SEQUENCE_LENGTH, - num_calibration_samples=NUM_CALIBRATION_SAMPLES, - moe_calibrate_all_experts=True, - data_collator=data_collator, - ) - print(f"[INFO] oneshot() completo en {time.time() - t_quant:.1f}s") - report_memory("post-oneshot") + data_collator = CalibrationDataCollator(tokenizer.pad_token_id or tokenizer.eos_token_id) - OUTPUT_PATH.mkdir(parents=True, exist_ok=True) - print(f"[INFO] guardando modelo cuantizado en {OUTPUT_PATH}") - t_save = time.time() - model.save_pretrained(str(OUTPUT_PATH)) - processor.save_pretrained(str(OUTPUT_PATH)) - print(f"[INFO] save_pretrained completo en {time.time() - t_save:.1f}s") + print("[INFO] arrancando oneshot() -- calibracion NVFP4 con moe_calibrate_all_experts=True") + t_quant = time.time() + oneshot( + model=model, + recipe=recipe, + dataset=calibration_dataset, + max_seq_length=MAX_SEQUENCE_LENGTH, + num_calibration_samples=NUM_CALIBRATION_SAMPLES, + moe_calibrate_all_experts=True, + data_collator=data_collator, + ) + print(f"[INFO] oneshot() completo en {time.time() - t_quant:.1f}s") + report_memory("post-oneshot") - print(f"[INFO] reinyectando tensores MTP desde {MODEL_PATH}") - save_mtp_tensors_to_checkpoint(source_model=str(MODEL_PATH), dest_dir=str(OUTPUT_PATH)) - print("[INFO] tensores MTP reinyectados") + OUTPUT_PATH.mkdir(parents=True, exist_ok=True) + print(f"[INFO] guardando modelo cuantizado en {OUTPUT_PATH}") + t_save = time.time() + model.save_pretrained(str(OUTPUT_PATH)) + processor.save_pretrained(str(OUTPUT_PATH)) + print(f"[INFO] save_pretrained completo en {time.time() - t_save:.1f}s") + + print(f"[INFO] reinyectando tensores MTP desde {MODEL_PATH}") + save_mtp_tensors_to_checkpoint(source_model=str(MODEL_PATH), dest_dir=str(OUTPUT_PATH)) + print("[INFO] tensores MTP reinyectados") print("[INFO] verificando checkpoint de salida") quant_config = verify_quantization_config()