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,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():
|
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}")
|
||||||
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}")
|
if args.verify_only:
|
||||||
processor = AutoProcessor.from_pretrained(str(MODEL_PATH), trust_remote_code=True)
|
print("[INFO] --verify-only: saltando calibracion/guardado, solo verificando OUTPUT_PATH existente")
|
||||||
tokenizer = processor.tokenizer
|
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)")
|
calibration_dataset = load_calibration_dataset(tokenizer)
|
||||||
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")
|
|
||||||
|
|
||||||
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")
|
data_collator = CalibrationDataCollator(tokenizer.pad_token_id or tokenizer.eos_token_id)
|
||||||
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")
|
|
||||||
|
|
||||||
OUTPUT_PATH.mkdir(parents=True, exist_ok=True)
|
print("[INFO] arrancando oneshot() -- calibracion NVFP4 con moe_calibrate_all_experts=True")
|
||||||
print(f"[INFO] guardando modelo cuantizado en {OUTPUT_PATH}")
|
t_quant = time.time()
|
||||||
t_save = time.time()
|
oneshot(
|
||||||
model.save_pretrained(str(OUTPUT_PATH))
|
model=model,
|
||||||
processor.save_pretrained(str(OUTPUT_PATH))
|
recipe=recipe,
|
||||||
print(f"[INFO] save_pretrained completo en {time.time() - t_save:.1f}s")
|
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}")
|
OUTPUT_PATH.mkdir(parents=True, exist_ok=True)
|
||||||
save_mtp_tensors_to_checkpoint(source_model=str(MODEL_PATH), dest_dir=str(OUTPUT_PATH))
|
print(f"[INFO] guardando modelo cuantizado en {OUTPUT_PATH}")
|
||||||
print("[INFO] tensores MTP reinyectados")
|
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")
|
print("[INFO] verificando checkpoint de salida")
|
||||||
quant_config = verify_quantization_config()
|
quant_config = verify_quantization_config()
|
||||||
|
|||||||
Reference in New Issue
Block a user