Phase 5: re-quantize merged checkpoint to NVFP4 with MTP/vision tensor reinjection and production-config verification #4
@@ -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,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():
|
||||
args = parse_args()
|
||||
print(f"[INFO] MODEL_PATH={MODEL_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] NUM_CALIBRATION_SAMPLES={NUM_CALIBRATION_SAMPLES} MAX_SEQUENCE_LENGTH={MAX_SEQUENCE_LENGTH}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user