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.
|
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}")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user