Fase 5: 21_quantize_nvfp4.py - corregir verificacion de tensores de vision

La verificacion automatica asumia (siguiendo el layout del checkpoint de
referencia de RedHatAI) que save_pretrained() separaria vision a su propio
model_visual.safetensors. En la practica, en 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). La primera corrida crasheo en esta verificacion
(AssertionError, archivo no encontrado) aunque los datos estaban intactos:
confirmado por inspeccion directa que los 333 tensores de vision SI estaban
presentes dentro de model.safetensors. Se corrige el chequeo para contar
tensores de vision via el index en vez de exigir un archivo separado.

Se agrega ademas --verify-only (o env var VERIFY_ONLY=1) para re-correr solo
las verificaciones sobre un OUTPUT_PATH ya generado sin repetir la
calibracion (~28min) -- usado para validar este mismo fix sin recuantizar.

Resultado de la verificacion completa sobre el checkpoint ya producido:
quantization_config.format=nvfp4-pack-quantized, model_mtp.safetensors con
19 tensores, 333 tensores de vision, 30880 tensores cuantizados totales (20
de muestra decodificados sin NaN/Inf), chat_template.jinja identico al de
produccion, tamano total 23.35GB.
This commit is contained in:
2026-07-29 23:58:57 +00:00
parent c9878ef98f
commit 9f073034e4
+46 -18
View File
@@ -24,17 +24,24 @@ Algoritmo:
truncados a MAX_SEQUENCE_LENGTH. truncados a MAX_SEQUENCE_LENGTH.
4. oneshot(..., moe_calibrate_all_experts=True) -- obligatorio, si no la 4. oneshot(..., moe_calibrate_all_experts=True) -- obligatorio, si no la
mayoria de los 256 expertos ruteados quedan sin calibrar. mayoria de los 256 expertos ruteados quedan sin calibrar.
5. model.save_pretrained(OUTPUT_PATH) (separa visual a model_visual.safetensors 5. model.save_pretrained(OUTPUT_PATH), processor.save_pretrained(OUTPUT_PATH)
nativamente), processor.save_pretrained(OUTPUT_PATH) (copia el (copia el chat_template.jinja de produccion, no el de masking),
chat_template.jinja de produccion, no el de masking),
save_mtp_tensors_to_checkpoint(source_model=MODEL_PATH, dest_dir=OUTPUT_PATH) 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 (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 6. Verificacion automatica (aborta si algo no cuadra): quantization_config.format
== nvfp4-pack-quantized; model_mtp.safetensors y model_visual.safetensors == nvfp4-pack-quantized; model_mtp.safetensors presente y conteo de tensores
presentes con conteos de tensores razonables; muestra de tensores de vision razonable (via el index, sin asumir un archivo separado); muestra de
cuantizados decodifica sin NaN/Inf; chat_template.jinja NO identico al de tensores cuantizados decodifica sin NaN/Inf; chat_template.jinja NO identico
masking de training y SI identico al de MODEL_PATH. 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 json
import os import os
@@ -199,24 +206,27 @@ def verify_quantization_config():
def verify_mtp_and_visual_shards(): def verify_mtp_and_visual_shards():
mtp_path = OUTPUT_PATH / "model_mtp.safetensors" mtp_path = OUTPUT_PATH / "model_mtp.safetensors"
visual_path = OUTPUT_PATH / "model_visual.safetensors"
if not mtp_path.exists(): if not mtp_path.exists():
raise AssertionError(f"falta {mtp_path} -- los tensores MTP no se reinyectaron, --speculative-config no arrancara") 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: with safe_open(str(mtp_path), framework="pt") as f:
mtp_count = len(f.keys()) 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_mtp.safetensors: {mtp_count} tensores")
print(f"[INFO] model_visual.safetensors: {visual_count} tensores") # 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.
# 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.
if not (10 <= mtp_count <= 40): if not (10 <= mtp_count <= 40):
raise AssertionError(f"conteo de tensores MTP fuera de rango razonable: {mtp_count} (esperado ~19)") 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): if not (250 <= visual_count <= 450):
raise AssertionError(f"conteo de tensores de vision fuera de rango razonable: {visual_count} (esperado ~333)") raise AssertionError(f"conteo de tensores de vision fuera de rango razonable: {visual_count} (esperado ~333)")
@@ -294,9 +304,27 @@ 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(): def main():
args = parse_args()
print(f"[INFO] MODEL_PATH={MODEL_PATH}") print(f"[INFO] MODEL_PATH={MODEL_PATH}")
print(f"[INFO] OUTPUT_PATH={OUTPUT_PATH}") print(f"[INFO] OUTPUT_PATH={OUTPUT_PATH}")
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] TRAIN_DATA_PATH={TRAIN_DATA_PATH}")
print(f"[INFO] NUM_CALIBRATION_SAMPLES={NUM_CALIBRATION_SAMPLES} MAX_SEQUENCE_LENGTH={MAX_SEQUENCE_LENGTH}") print(f"[INFO] NUM_CALIBRATION_SAMPLES={NUM_CALIBRATION_SAMPLES} MAX_SEQUENCE_LENGTH={MAX_SEQUENCE_LENGTH}")