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.
387 lines
18 KiB
Python
387 lines
18 KiB
Python
"""Fase 5: cuantiza a NVFP4 el checkpoint mergeado de Fase 4 via llm-compressor,
|
|
clonando la receta exacta de RedHatAI (Qwen3.6-35B-A3B-NVFP4), y reinyecta los
|
|
tensores MTP (bf16, aparte) que la clase de carga no instancia.
|
|
|
|
Corre DENTRO del contenedor `qwen-lora-train` en spark (requiere `llmcompressor` y
|
|
`compressed-tensors` ya instalados ahi via pip --no-deps, ver Docmost de esta fase):
|
|
|
|
docker exec qwen-lora-train python3 \
|
|
/workspace/ai-projects/qwen3-6-lora/scripts/21_quantize_nvfp4.py
|
|
|
|
Diferencia obligatoria frente a la receta original de RedHatAI: la calibracion usa
|
|
una muestra de data/train.jsonl (el propio dataset de fine-tuning) en vez del
|
|
corpus generico ultrachat_200k, aplicando el chat template de PRODUCCION (el que
|
|
ya trae MODEL_PATH, no el de masking de training).
|
|
|
|
Algoritmo:
|
|
1. Cargar Qwen3_5MoeForConditionalGeneration.from_pretrained(MODEL_PATH,
|
|
dtype="auto") + AutoProcessor.from_pretrained(MODEL_PATH).
|
|
2. Recipe QuantizationModifier(targets="Linear", scheme="NVFP4", ignore=[...])
|
|
-- ignore list identica a la de RedHatAI (recipe.yaml del checkpoint real).
|
|
3. Muestra de calibracion: NUM_CALIBRATION_SAMPLES ejemplos aleatorios
|
|
(random.seed(42)) de TRAIN_DATA_PATH, renderizados con
|
|
processor.tokenizer.apply_chat_template (template de produccion) y
|
|
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), 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). 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 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
|
|
import random
|
|
import time
|
|
from pathlib import Path
|
|
|
|
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
|
|
|
import torch
|
|
from datasets import Dataset
|
|
from safetensors import safe_open
|
|
from transformers import AutoProcessor, Qwen3_5MoeForConditionalGeneration
|
|
|
|
# Compat shim: llmcompressor 0.12.0 (ultima version en PyPI) importa incondicionalmente
|
|
# GraniteMoeParallelExperts al armar su registro de arquitecturas MoE "linearizables"
|
|
# (llmcompressor/modeling/moe/granitemoe.py), incluso para modelos que no son GraniteMoe
|
|
# (como este Qwen3.5 MoE). transformers 5.14.1 renombro esa clase a GraniteMoeExperts,
|
|
# rompiendo ese import y abortando oneshot() para CUALQUIER modelo. Como Qwen3.5 MoE no
|
|
# esta en ese registro (solo granitemoe/llama4), el alias nunca se usa realmente -- solo
|
|
# hace falta que el nombre exista para que el import no explote.
|
|
import transformers.models.granitemoe.modeling_granitemoe as _granitemoe_mod # noqa: E402
|
|
|
|
if not hasattr(_granitemoe_mod, "GraniteMoeParallelExperts"):
|
|
_granitemoe_mod.GraniteMoeParallelExperts = _granitemoe_mod.GraniteMoeExperts
|
|
|
|
from compressed_tensors.utils import save_mtp_tensors_to_checkpoint # noqa: E402
|
|
from llmcompressor import oneshot # noqa: E402
|
|
from llmcompressor.modifiers.quantization import QuantizationModifier # noqa: E402
|
|
|
|
REPO_ROOT = Path(__file__).resolve().parent.parent
|
|
MODEL_PATH = Path(os.environ.get("MODEL_PATH", "/workspace/ft-models/Qwen3.6-35B-A3B-mcp-bf16"))
|
|
OUTPUT_PATH = Path(os.environ.get("OUTPUT_PATH", "/workspace/ft-models/Qwen3.6-35B-A3B-mcp-NVFP4"))
|
|
TRAIN_DATA_PATH = Path(os.environ.get("TRAIN_DATA_PATH", str(REPO_ROOT / "data" / "train.jsonl")))
|
|
NUM_CALIBRATION_SAMPLES = int(os.environ.get("NUM_CALIBRATION_SAMPLES", "256"))
|
|
MAX_SEQUENCE_LENGTH = int(os.environ.get("MAX_SEQUENCE_LENGTH", "8192"))
|
|
CALIBRATION_SEED = 42
|
|
|
|
# Receta identica a la de RedHatAI/Qwen3.6-35B-A3B-NVFP4 (recipe.yaml leido por SSH,
|
|
# citado integro en PLAN.md). Las capas Gated DeltaNet (linear_attn) quedan en BF16
|
|
# a proposito -- el LoRA se entreno ahi, pero al no cuantizarse no se agrega perdida
|
|
# de precision adicional sobre lo ya mergeado en Fase 4.
|
|
QUANTIZATION_IGNORE = [
|
|
"re:.*lm_head",
|
|
"re:visual.*",
|
|
"re:model.visual.*",
|
|
"re:.*mlp.gate$",
|
|
"re:.*embed_tokens$",
|
|
"re:.*shared_expert_gate$",
|
|
"re:.*linear_attn.*",
|
|
]
|
|
|
|
|
|
def load_calibration_dataset(tokenizer):
|
|
print(f"[INFO] cargando ejemplos de calibracion desde {TRAIN_DATA_PATH}")
|
|
examples = []
|
|
with open(TRAIN_DATA_PATH, encoding="utf-8") as f:
|
|
for line in f:
|
|
line = line.strip()
|
|
if line:
|
|
examples.append(json.loads(line))
|
|
print(f"[INFO] {len(examples)} ejemplos disponibles en {TRAIN_DATA_PATH}")
|
|
|
|
rng = random.Random(CALIBRATION_SEED)
|
|
rng.shuffle(examples)
|
|
sampled = examples[:NUM_CALIBRATION_SAMPLES]
|
|
if len(sampled) < NUM_CALIBRATION_SAMPLES:
|
|
raise AssertionError(
|
|
f"se pidieron {NUM_CALIBRATION_SAMPLES} muestras de calibracion pero "
|
|
f"{TRAIN_DATA_PATH} solo tiene {len(examples)} ejemplos"
|
|
)
|
|
|
|
from collections import Counter
|
|
|
|
bucket_counts = Counter(ex.get("meta", {}).get("bucket", "?") for ex in sampled)
|
|
print(f"[INFO] distribucion de buckets en la muestra de calibracion: {dict(bucket_counts)}")
|
|
|
|
input_ids_list = []
|
|
attention_mask_list = []
|
|
for ex in sampled:
|
|
text = tokenizer.apply_chat_template(
|
|
ex["messages"],
|
|
tools=ex.get("tools"),
|
|
tokenize=False,
|
|
add_generation_prompt=False,
|
|
)
|
|
encoded = tokenizer(
|
|
text,
|
|
truncation=True,
|
|
max_length=MAX_SEQUENCE_LENGTH,
|
|
add_special_tokens=False,
|
|
)
|
|
input_ids_list.append(encoded["input_ids"])
|
|
attention_mask_list.append(encoded["attention_mask"])
|
|
|
|
lengths = [len(ids) for ids in input_ids_list]
|
|
print(
|
|
f"[INFO] longitudes de calibracion: min={min(lengths)} max={max(lengths)} "
|
|
f"avg={sum(lengths) / len(lengths):.1f}"
|
|
)
|
|
return Dataset.from_dict({"input_ids": input_ids_list, "attention_mask": attention_mask_list})
|
|
|
|
|
|
class CalibrationDataCollator:
|
|
"""Padding simple a la derecha -- sin labels, oneshot solo necesita forward pass."""
|
|
|
|
def __init__(self, pad_token_id):
|
|
self.pad_token_id = pad_token_id
|
|
|
|
def __call__(self, features):
|
|
max_len = max(len(f["input_ids"]) for f in features)
|
|
input_ids = []
|
|
attention_mask = []
|
|
for f in features:
|
|
ids = f["input_ids"]
|
|
mask = f["attention_mask"]
|
|
pad_len = max_len - len(ids)
|
|
input_ids.append(ids + [self.pad_token_id] * pad_len)
|
|
attention_mask.append(mask + [0] * pad_len)
|
|
return {
|
|
"input_ids": torch.tensor(input_ids, dtype=torch.long),
|
|
"attention_mask": torch.tensor(attention_mask, dtype=torch.long),
|
|
}
|
|
|
|
|
|
def report_memory(label):
|
|
if torch.cuda.is_available():
|
|
alloc_gb = torch.cuda.memory_allocated() / (1024 ** 3)
|
|
peak_gb = torch.cuda.max_memory_allocated() / (1024 ** 3)
|
|
print(f"[INFO] memoria CUDA en '{label}': alloc={alloc_gb:.2f}GB peak={peak_gb:.2f}GB")
|
|
|
|
|
|
def iter_output_tensor_names():
|
|
"""Nombres de todos los tensores en OUTPUT_PATH (via el/los indice(s) de shards)."""
|
|
names = set()
|
|
for index_name in ("model.safetensors.index.json", "model_visual.safetensors.index.json"):
|
|
index_path = OUTPUT_PATH / index_name
|
|
if index_path.exists():
|
|
weight_map = json.loads(index_path.read_text())["weight_map"]
|
|
for key, shard in weight_map.items():
|
|
names.add((key, shard))
|
|
for single_name in ("model.safetensors", "model_visual.safetensors"):
|
|
single_path = OUTPUT_PATH / single_name
|
|
if single_path.exists():
|
|
with safe_open(str(single_path), framework="pt") as f:
|
|
for key in f.keys():
|
|
names.add((key, single_name))
|
|
return names
|
|
|
|
|
|
def verify_quantization_config():
|
|
config = json.loads((OUTPUT_PATH / "config.json").read_text())
|
|
quant_config = config.get("quantization_config")
|
|
if quant_config is None:
|
|
raise AssertionError("config.json de salida no tiene quantization_config -- la cuantizacion no se aplico")
|
|
fmt = quant_config.get("format")
|
|
if fmt != "nvfp4-pack-quantized":
|
|
raise AssertionError(f"quantization_config.format inesperado: {fmt!r} (se esperaba 'nvfp4-pack-quantized')")
|
|
print(f"[INFO] quantization_config.format={fmt!r} confirmado")
|
|
return quant_config
|
|
|
|
|
|
def verify_mtp_and_visual_shards():
|
|
mtp_path = OUTPUT_PATH / "model_mtp.safetensors"
|
|
if not mtp_path.exists():
|
|
raise AssertionError(f"falta {mtp_path} -- los tensores MTP no se reinyectaron, --speculative-config no arrancara")
|
|
|
|
with safe_open(str(mtp_path), framework="pt") as f:
|
|
mtp_count = len(f.keys())
|
|
print(f"[INFO] model_mtp.safetensors: {mtp_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.
|
|
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)")
|
|
|
|
|
|
def verify_no_nan_inf_in_sample(quant_config, sample_size=20):
|
|
"""Decodifica una muestra de tensores cuantizados (weight_packed) y verifica
|
|
que no haya NaN/Inf tras la des-cuantizacion -- catch de errores numericos
|
|
silenciosos en la calibracion (scales cero/Inf, experts sin calibrar, etc.)."""
|
|
from compressed_tensors.compressors.nvfp4 import NVFP4PackedCompressor
|
|
from compressed_tensors.quantization.quant_scheme import PRESET_SCHEMES, QuantizationScheme
|
|
|
|
preset = PRESET_SCHEMES["NVFP4"]
|
|
scheme = QuantizationScheme(
|
|
targets=["Linear"],
|
|
weights=preset["weights"],
|
|
input_activations=preset.get("input_activations"),
|
|
)
|
|
|
|
all_names = iter_output_tensor_names()
|
|
packed_names = sorted(name for name, _shard in all_names if name.endswith(".weight_packed"))
|
|
if not packed_names:
|
|
raise AssertionError("no se encontro ningun tensor '.weight_packed' en OUTPUT_PATH -- nada se cuantizo")
|
|
print(f"[INFO] {len(packed_names)} tensores cuantizados (weight_packed) encontrados en total")
|
|
|
|
rng = random.Random(CALIBRATION_SEED)
|
|
sample = rng.sample(packed_names, min(sample_size, len(packed_names)))
|
|
|
|
shard_by_name = dict(all_names)
|
|
checked = 0
|
|
for packed_name in sample:
|
|
base_name = packed_name[: -len(".weight_packed")]
|
|
shard = shard_by_name[packed_name]
|
|
shard_path = OUTPUT_PATH / shard
|
|
with safe_open(str(shard_path), framework="pt") as f:
|
|
state_dict = {"weight_packed": f.get_tensor(packed_name)}
|
|
scale_name = f"{base_name}.weight_scale"
|
|
global_scale_name = f"{base_name}.weight_global_scale"
|
|
if scale_name in f.keys():
|
|
state_dict["weight_scale"] = f.get_tensor(scale_name)
|
|
if global_scale_name in f.keys():
|
|
state_dict["weight_global_scale"] = f.get_tensor(global_scale_name)
|
|
|
|
decompressed = NVFP4PackedCompressor.decompress(state_dict, scheme)
|
|
weight = decompressed["weight"]
|
|
if not torch.isfinite(weight).all():
|
|
raise AssertionError(f"tensor cuantizado {base_name} tiene NaN/Inf tras des-cuantizar")
|
|
checked += 1
|
|
|
|
print(f"[INFO] {checked} tensores cuantizados des-cuantizados sin NaN/Inf (muestra aleatoria)")
|
|
|
|
for ignored_suffix in ("mlp.gate.weight", "shared_expert_gate.weight", "embed_tokens.weight"):
|
|
if any(name.endswith(f"{ignored_suffix}_packed") for name, _s in all_names):
|
|
raise AssertionError(f"un tensor ignorado ({ignored_suffix}) fue cuantizado -- la ignore list no se aplico bien")
|
|
|
|
|
|
def verify_chat_template():
|
|
train_template_path = REPO_ROOT / "data" / "chat_template_train.jinja"
|
|
output_template_path = OUTPUT_PATH / "chat_template.jinja"
|
|
model_template_path = MODEL_PATH / "chat_template.jinja"
|
|
|
|
output_template = output_template_path.read_bytes()
|
|
train_template = train_template_path.read_bytes()
|
|
model_template = model_template_path.read_bytes()
|
|
|
|
if output_template == train_template:
|
|
raise AssertionError(
|
|
"chat_template.jinja de salida es BYTE-IDENTICO al template de masking de "
|
|
"training -- se copio el template equivocado"
|
|
)
|
|
if output_template != model_template:
|
|
raise AssertionError("chat_template.jinja de salida no coincide con el de MODEL_PATH (produccion)")
|
|
print(
|
|
f"[INFO] chat_template.jinja verificado: {len(output_template)} bytes, identico al "
|
|
"de produccion (MODEL_PATH), distinto del de masking de training"
|
|
)
|
|
|
|
|
|
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}")
|
|
|
|
print(f"[INFO] cargando processor desde {MODEL_PATH}")
|
|
processor = AutoProcessor.from_pretrained(str(MODEL_PATH), trust_remote_code=True)
|
|
tokenizer = processor.tokenizer
|
|
|
|
calibration_dataset = load_calibration_dataset(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")
|
|
|
|
recipe = QuantizationModifier(targets="Linear", scheme="NVFP4", ignore=QUANTIZATION_IGNORE)
|
|
|
|
data_collator = CalibrationDataCollator(tokenizer.pad_token_id or tokenizer.eos_token_id)
|
|
|
|
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")
|
|
|
|
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()
|
|
verify_mtp_and_visual_shards()
|
|
verify_no_nan_inf_in_sample(quant_config)
|
|
verify_chat_template()
|
|
|
|
total_size_gb = sum(f.stat().st_size for f in OUTPUT_PATH.rglob("*") if f.is_file()) / (1024 ** 3)
|
|
print(f"[INFO] tamano total de OUTPUT_PATH: {total_size_gb:.2f}GB")
|
|
print("[INFO] cuantizacion NVFP4 completa y verificada")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|