Files
qwen3-6-lora/scripts/20_merge_lora.py
aleleba c65d309719 Phase 6.4: make the gates fail when they cannot verify something
A code review found seven ways these gates could pass green with something
actually wrong. All are the same family: a missing value was treated as OK.
The rule now written into all three files is that absent is not OK, absent
is "could not verify", and that either fails or is reported as an explicit
SKIP - it never slips through as green.

30_eval_suite.py:
- A bucket with no baseline of its own fell back to the global 0.2750 and
  printed it in a column headed "baseline", as if it were that bucket's
  number. Measured against the real eval.jsonl buckets: negativos going
  from 0.12 to 0.33 is a real +0.21 regression, but the computed delta was
  +0.055 and it PASSED; manejo_errores sitting unchanged at 0.42 produced
  a fabricated +0.145 FAIL that would have discarded a healthy candidate
  mid-downtime. Now such buckets print SKIP and the verdict reports how
  many went unverified.
- "VEREDICTO: FAIL" exited 0, so a runbook chaining the gate into
  quantization would have carried on to write 24 GB. Now exits 1.
- A typo in BASELINE_BUCKET_LOSSES silently matched nothing; now aborts.
- The penpot exemption is labelled honestly: those 11 rows are pre-existing
  LoRA #1 tool-calling, not new capability, so gate 1 has no regression
  coverage there and the log says so.

20_merge_lora.py dry-run (merge path untouched, verified by AST diff):
- adapter_config.get("use_rslora", False) meant a missing key passed AND
  the log printed use_rslora=False, asserting it had checked something that
  was never there. A different PEFT version omitting a key was enough.
- lora_bias was not checked at all, only bias. They are different fields:
  lora_bias puts a bias inside lora_B, which W + scaling * (B @ A) ignores.
- The 620 keys were printed but never asserted, so an adapter with extra
  tensors printed "310 + 310 = 930" and passed.
- rank_pattern/alpha_pattern were not checked. They set r per module, so
  scaling is not uniformly alpha/r while both the dry-run and the merge
  apply a single 2.0 to all 310 tensors.
- A missing family was invisible: swap linear_attn for 150 mlp.gate targets
  and the total is still 310, no norm is zero because the family is simply
  gone, and it passed. Now presence and per-family counts are asserted,
  derived from the real adapter: linear_attn 150, shared_expert 120,
  attention_qkvo 40, otros 0.
Verified against seven synthetic adapters plus the real phase 3 one; only
the correct adapter passes.

21_quantize_nvfp4.py (recipe and oneshot untouched): the calibration cache
now carries a provenance.json recording the training file's sha256, the
recipe numbers and the bucket distribution, and loading aborts on mismatch.
This is the phase's number one risk and it had no mechanical defence: the
phase 5 cache on disk has exactly 512 rows, the same as the v2 recipe, so
the only existing check could not tell them apart and reusing it would have
calibrated with zero design data and washed out the new capability
silently. Verified: that cache now aborts.

gate 5: retry transport failures against the Penpot MCP, which drops
connections mid-call intermittently (seen before in phase 4's gate 4).
Without it a blip on prompt 6 of 8 kills a whole run and reads like a model
failure. PluginNotConnected is deliberately not retried - that is a real
state of the world. Also unwrap the {"result":..., "log":...} envelope the
server wraps execute_code returns in; the gate was reading keys off the
outer object and rejecting a valid page setup.
2026-07-30 17:37:46 +00:00

596 lines
28 KiB
Python

"""Fase 4: mergea el adapter LoRA (out/lora-adapter/) sobre el checkpoint base BF16,
shard-a-shard, sin cargar el modelo completo via AutoModelForCausalLM.
Corre DENTRO del contenedor `qwen-lora-train` en spark:
docker exec qwen-lora-train python3 \
/workspace/ai-projects/qwen3-6-lora/.worktrees/agente-fase4-merge-eval/scripts/20_merge_lora.py
Algoritmo (opera directo sobre tensores crudos, nunca instancia el modelo):
1. Cargar adapter_model.safetensors completo (~190MB), parsear claves PEFT
(prefijo "base_model.model." + sufijo ".lora_A.weight"/".lora_B.weight") en
{nombre_tensor_base: (lora_A, lora_B)}. scaling = lora_alpha / r.
2. Leer MODEL_PATH/model.safetensors.index.json -> weight_map.
3. Por cada shard unico: cargar, mergear en fp32 los tensores LoRA-target
(W + scaling * (B @ A)) y volver a bf16; copiar el resto tal cual (esto
preserva mtp.*/visual.* automaticamente, sin logica especial). Guardar el
shard con el mismo nombre en OUTPUT_PATH.
4. Copiar sin cambios model.safetensors.index.json, config.json,
generation_config.json, archivos de tokenizer, y chat_template.jinja DESDE
MODEL_PATH (nunca desde ADAPTER_PATH -- ese es el template de masking de
training, no el de inferencia real).
5. Verificacion automatica: conteo de tensores igual; todo tensor no-target
byte-a-byte identico al base; todo tensor LoRA-target con delta no-cero;
sin NaN/Inf.
Soporta --dry-run (Fase 6): corre en ~2 segundos, SIN cargar los pesos del modelo
base (solo lee el adapter y el model.safetensors.index.json del checkpoint), y
asierte todo lo que, de estar mal, se descubriria recien despues de escribir 67 GB:
conteo de targets resueltos (310) y de claves del adapter (620), scaling, flags de
la variante de LoRA (rsLoRA/DoRA/bias/lora_bias/modules_to_save/rank_pattern/
alpha_pattern), presencia y conteo por familia de modulos, norma de lora_B por
familia, y que todas las claves remapeadas existan en el indice del checkpoint base.
Sale antes de escribir nada -- la ruta de merge real no se toca.
PRINCIPIO RECTOR del dry-run: ausente no es OK; ausente es "no se pudo verificar", y
eso tiene que fallar, nunca colarse como verde. Una clave que falta en
adapter_config.json (porque la entreno otra version de PEFT), una familia de modulos
que no aparece, o un conteo que nadie comparo son exactamente la forma en que este
chequeo produciria 67 GB con capacidad silenciosamente incompleta.
"""
import argparse
import gc
import json
import os
import re
import shutil
import time
from pathlib import Path
import torch
from safetensors import safe_open
from safetensors.torch import save_file
REPO_ROOT = Path(__file__).resolve().parent.parent
MODEL_PATH = Path(os.environ.get("MODEL_PATH", "/workspace/ft-models/Qwen--Qwen3.6-35B-A3B"))
ADAPTER_PATH = Path(os.environ.get("ADAPTER_PATH", str(REPO_ROOT / "out" / "lora-adapter")))
OUTPUT_PATH = Path(os.environ.get("OUTPUT_PATH", "/workspace/ft-models/Qwen3.6-35B-A3B-mcp-bf16"))
ADAPTER_PREFIX = "base_model.model."
LORA_A_SUFFIX = ".lora_A.weight"
LORA_B_SUFFIX = ".lora_B.weight"
# Invariantes del adapter esperados por el --dry-run. Son los de Fase 3 y los del
# LoRA #2 de Fase 6 (r/alpha sin cambios a proposito, ver PLAN.md): 310 modulos
# objetivo = 620 claves (lora_A + lora_B), con scaling = lora_alpha / r = 2.0.
#
# Los defaults literales estan aparte de los valores en uso a proposito: son
# overrideables por env y la corrida real hereda el env del contenedor, asi que si
# alguien exporta uno para "destrabar" una corrida, la asercion se vuelve tautologica.
# warn_expected_overrides() lo deja escrito en el log en vez de dejarlo pasar mudo.
DEFAULT_EXPECTED_TARGET_COUNT = 310
DEFAULT_EXPECTED_R = 32
DEFAULT_EXPECTED_LORA_ALPHA = 64
EXPECTED_TARGET_COUNT = int(os.environ.get("EXPECTED_TARGET_COUNT", str(DEFAULT_EXPECTED_TARGET_COUNT)))
EXPECTED_R = int(os.environ.get("EXPECTED_R", str(DEFAULT_EXPECTED_R)))
EXPECTED_LORA_ALPHA = int(os.environ.get("EXPECTED_LORA_ALPHA", str(DEFAULT_EXPECTED_LORA_ALPHA)))
EXPECTED_SCALING = EXPECTED_LORA_ALPHA / EXPECTED_R
# Desglose de los 310 targets por familia de modulos, derivado de los 12 sufijos de
# TARGET_MODULES (scripts/10_train.py) y de la topologia del modelo:
# linear_attn = 30 capas Gated DeltaNet x 5 sufijos
# (in_proj_qkv, in_proj_z, in_proj_a, in_proj_b, out_proj) = 150
# shared_expert = 40 capas MoE x 3 sufijos (gate_proj, up_proj, down_proj) = 120
# attention_qkvo = 10 capas de atencion completa x 4 (q/k/v/o_proj) = 40
# total = 310
# Se asierte PRESENCIA y CONTEO EXACTO de cada familia, no solo el total: un
# TARGET_MODULES mal escrito que no toque linear_attn y enganche otros 150 modulos
# deja el total en 310 y ninguna norma en cero (linear_attn simplemente no esta),
# asi que el chequeo de normas por familia no lo ve. La familia "otros" es el
# catch-all de module_family(): tiene que quedar VACIA -- cualquier cosa ahi es un
# modulo que nadie previo (por ejemplo mlp.gate, el router).
EXPECTED_FAMILY_COUNTS = {
"linear_attn": 150,
"shared_expert": 120,
"attention_qkvo": 40,
"otros": 0,
}
# Umbral relativo para "modulo efectivamente muerto": norma de lora_B por debajo de
# esta fraccion de la mediana de su familia. El chequeo de norm == 0.0 exacto atrapa
# el cero de la init de PEFT (riesgo #8), pero un modulo con norma 1e-12 pasaria
# igual de mudo y esta igual de muerto.
DEAD_MODULE_RELATIVE_THRESHOLD = 1e-6
# El adapter fue entrenado cargando el checkpoint con AutoModelForCausalLM, que expone las
# capas como "model.layers.N...."; el checkpoint base crudo (multimodal) las tiene bajo
# "model.language_model.layers.N....". Hay que remapear el nombre del tensor base antes de
# buscarlo en el mapa de shards. embed_tokens/norm top-level tienen el mismo desplazamiento;
# lm_head y mtp.*/visual.* no son target de LoRA y no necesitan remapeo.
ADAPTER_TO_CHECKPOINT_PREFIX = {
"model.layers.": "model.language_model.layers.",
"model.embed_tokens.": "model.language_model.embed_tokens.",
"model.norm.": "model.language_model.norm.",
}
def remap_adapter_name_to_checkpoint_name(name):
for adapter_prefix, checkpoint_prefix in ADAPTER_TO_CHECKPOINT_PREFIX.items():
if name.startswith(adapter_prefix):
return checkpoint_prefix + name[len(adapter_prefix):]
return name
NON_MODEL_FILES = [
"config.json",
"generation_config.json",
"configuration.json",
"tokenizer.json",
"tokenizer_config.json",
"merges.txt",
"vocab.json",
"chat_template.jinja",
"preprocessor_config.json",
"video_preprocessor_config.json",
"LICENSE",
"README.md",
]
def load_lora_deltas():
adapter_config = json.loads((ADAPTER_PATH / "adapter_config.json").read_text())
r = adapter_config["r"]
lora_alpha = adapter_config["lora_alpha"]
scaling = lora_alpha / r
print(f"[INFO] r={r} lora_alpha={lora_alpha} scaling={scaling}")
deltas = {}
with safe_open(str(ADAPTER_PATH / "adapter_model.safetensors"), framework="pt") as f:
keys = list(f.keys())
base_names = set()
for k in keys:
if k.endswith(LORA_A_SUFFIX):
base_names.add(k[len(ADAPTER_PREFIX):-len(LORA_A_SUFFIX)])
for base_name in base_names:
key_a = f"{ADAPTER_PREFIX}{base_name}{LORA_A_SUFFIX}"
key_b = f"{ADAPTER_PREFIX}{base_name}{LORA_B_SUFFIX}"
lora_a = f.get_tensor(key_a).to(torch.float32)
lora_b = f.get_tensor(key_b).to(torch.float32)
checkpoint_name = remap_adapter_name_to_checkpoint_name(f"{base_name}.weight")
deltas[checkpoint_name] = (lora_a, lora_b, scaling)
print(f"[INFO] {len(deltas)} tensores objetivo de LoRA encontrados en el adapter")
return deltas
def merge_shards(deltas):
index = json.loads((MODEL_PATH / "model.safetensors.index.json").read_text())
weight_map = index["weight_map"]
shard_files = sorted(set(weight_map.values()))
print(f"[INFO] {len(shard_files)} shards, {len(weight_map)} tensores totales")
OUTPUT_PATH.mkdir(parents=True, exist_ok=True)
merged_target_names = set()
total_tensors_in = 0
total_tensors_out = 0
checks_nontarget_sample = []
for shard_name in shard_files:
t0 = time.time()
shard_path = MODEL_PATH / shard_name
out_tensors = {}
with safe_open(str(shard_path), framework="pt") as f:
shard_keys = list(f.keys())
total_tensors_in += len(shard_keys)
for key in shard_keys:
tensor = f.get_tensor(key)
if key in deltas:
lora_a, lora_b, scaling = deltas[key]
w_fp32 = tensor.to(torch.float32)
delta = scaling * (lora_b @ lora_a)
merged = (w_fp32 + delta).to(torch.bfloat16)
if not torch.isfinite(merged).all():
raise AssertionError(f"NaN/Inf tras mergear tensor {key}")
if torch.equal(merged, tensor):
raise AssertionError(f"tensor LoRA-target {key} no cambio tras el merge (delta cero)")
out_tensors[key] = merged.contiguous()
merged_target_names.add(key)
else:
if not torch.isfinite(tensor.to(torch.float32)).all():
raise AssertionError(f"NaN/Inf en tensor no-target {key} del checkpoint base (bug pre-existente)")
out_tensors[key] = tensor.contiguous()
if len(checks_nontarget_sample) < 200:
checks_nontarget_sample.append((shard_name, key))
save_file(out_tensors, str(OUTPUT_PATH / shard_name), metadata={"format": "pt"})
total_tensors_out += len(out_tensors)
del out_tensors
gc.collect()
dt = time.time() - t0
peak_mb = torch.cuda.max_memory_allocated() / (1024 ** 2) if torch.cuda.is_available() else 0.0
print(f"[INFO] shard {shard_name}: {len(shard_keys)} tensores, {dt:.1f}s, peak_cuda={peak_mb:.0f}MB")
missing = merged_target_names.symmetric_difference(set(deltas.keys()))
if missing:
raise AssertionError(f"tensores LoRA-target no encontrados en ningun shard: {missing}")
if total_tensors_in != total_tensors_out:
raise AssertionError(f"conteo de tensores no cuadra: in={total_tensors_in} out={total_tensors_out}")
print(f"[INFO] {len(merged_target_names)} tensores mergeados, {total_tensors_out} tensores totales escritos")
return checks_nontarget_sample
def verify_nontarget_byte_identical(sample):
print(f"[INFO] verificando byte-a-byte {len(sample)} tensores no-target de muestra (incluye mtp.*/visual.*)")
mtp_or_visual_checked = 0
for shard_name, key in sample:
with safe_open(str(MODEL_PATH / shard_name), framework="pt") as f_base:
base_t = f_base.get_tensor(key)
with safe_open(str(OUTPUT_PATH / shard_name), framework="pt") as f_out:
out_t = f_out.get_tensor(key)
if not torch.equal(base_t, out_t):
raise AssertionError(f"tensor no-target {key} en {shard_name} NO es byte-identico al base")
if re.match(r"^(model\.)?mtp\.", key) or "visual" in key:
mtp_or_visual_checked += 1
print(f"[INFO] verificacion byte-a-byte ok ({mtp_or_visual_checked} tensores mtp/visual en la muestra)")
def copy_non_model_files():
for fname in NON_MODEL_FILES:
src = MODEL_PATH / fname
if src.exists():
shutil.copy2(src, OUTPUT_PATH / fname)
print(f"[INFO] copiado {fname} desde MODEL_PATH (nunca desde ADAPTER_PATH)")
shutil.copy2(
MODEL_PATH / "model.safetensors.index.json",
OUTPUT_PATH / "model.safetensors.index.json",
)
print("[INFO] copiado model.safetensors.index.json")
def verify_chat_template_is_not_training_template():
train_template = (REPO_ROOT / "data" / "chat_template_train.jinja").read_bytes()
output_template = (OUTPUT_PATH / "chat_template.jinja").read_bytes()
if output_template == train_template:
raise AssertionError(
"chat_template.jinja del checkpoint mergeado es BYTE-IDENTICO al template de "
"masking de training -- el merge tomo el template equivocado (debe venir de MODEL_PATH)"
)
base_template = (MODEL_PATH / "chat_template.jinja").read_bytes()
if output_template != base_template:
raise AssertionError("chat_template.jinja del checkpoint mergeado no coincide con el de MODEL_PATH")
print(
f"[INFO] chat_template.jinja verificado: {len(output_template)} bytes, "
"identico al de MODEL_PATH, distinto del template de training"
)
def module_family(base_name):
"""Familia de modulos a la que pertenece un target del adapter. El orden importa:
shared_expert tiene sus propios gate/up/down_proj y linear_attn sus propias
proyecciones, asi que ambos se chequean antes que la atencion q/k/v/o."""
if ".linear_attn." in base_name:
return "linear_attn"
if "shared_expert" in base_name:
return "shared_expert"
if re.search(r"\.(q|k|v|o)_proj$", base_name):
return "attention_qkvo"
return "otros"
def warn_expected_overrides():
"""Deja escrito en el log si algun EXPECTED_* viene pisado por el env. La corrida
real hereda el env del contenedor: sin este aviso, alguien que exporta un valor
para destrabar una corrida convierte la asercion en tautologia y el log sigue
diciendo [OK] igual."""
for name, literal in (
("EXPECTED_TARGET_COUNT", DEFAULT_EXPECTED_TARGET_COUNT),
("EXPECTED_R", DEFAULT_EXPECTED_R),
("EXPECTED_LORA_ALPHA", DEFAULT_EXPECTED_LORA_ALPHA),
):
in_use = globals()[name]
if in_use != literal:
print(
f"[WARN] {name} overrideado por env (valor literal {literal}, en uso {in_use}) "
"-- esta asercion NO esta verificando el invariante del proyecto"
)
def check_adapter_config_flags(adapter_config, problems):
"""Verifica las flags de adapter_config.json que cambian la semantica del merge.
Ausente NO es OK: una clave que falta (por ejemplo porque el adapter se entreno
con otra version de PEFT que la omite) es "no se pudo verificar", y se reporta
como problema. Antes, un .get(clave, False) daba verde Y ADEMAS imprimia
"clave=False", o sea que el log afirmaba haber verificado algo que nunca estuvo.
"""
reportado = {}
def leer(clave):
"""Devuelve (valor, presente). Registra el problema si la clave no esta."""
if clave not in adapter_config:
problems.append(
f"clave {clave!r} ausente del adapter_config, no se puede verificar "
"(ausente != OK: puede venir de otra version de PEFT que la omite, "
"y el merge la ignoraria en silencio)"
)
reportado[clave] = "AUSENTE"
return None, False
reportado[clave] = adapter_config[clave]
return adapter_config[clave], True
# rsLoRA escala por lora_alpha/sqrt(r) en vez de lora_alpha/r: un adapter
# entrenado con rsLoRA se mergearia con 2.0 donde corresponde 11.3 y pasaria
# TODAS las demas aserciones sin decir nada.
valor, presente = leer("use_rslora")
if presente and valor:
problems.append("use_rslora=true -- el merge aplica lora_alpha/r, rsLoRA usa lora_alpha/sqrt(r)")
# DoRA agrega un vector de magnitud que la formula W + scaling * (B @ A) ignora.
valor, presente = leer("use_dora")
if presente and valor:
problems.append("use_dora=true -- el merge ignora el vector de magnitud de DoRA")
# bias: bias del modulo BASE entrenado junto al adapter; el merge no lo aplica.
valor, presente = leer("bias")
if presente and valor != "none":
problems.append(f"bias={valor!r} -- el merge no aplica biases entrenados")
# lora_bias (PEFT >= 0.14) es OTRO campo, distinto de `bias`: agrega un termino de
# bias DENTRO de lora_B, que W + scaling * (B @ A) tampoco contempla. Riesgo #6
# del PLAN.md lo pide explicitamente; el fallo silencioso es el mismo que rsLoRA.
valor, presente = leer("lora_bias")
if presente and valor:
problems.append(
f"lora_bias={valor!r} -- PEFT agrega un bias dentro de lora_B que la formula "
"W + scaling * (B @ A) del merge ignora por completo"
)
# modules_to_save quedarian fuera del merge y se perderian en silencio.
valor, presente = leer("modules_to_save")
if presente and valor:
problems.append(f"modules_to_save={valor!r} -- esos modulos no se mergean y se perderian")
# rank_pattern / alpha_pattern permiten r y lora_alpha POR MODULO. Si estan
# poblados, scaling no es uniformemente alpha/r, pero tanto el dry-run como el
# merge real aplican un unico escalar a los 310 tensores: las capas con otro r se
# mergearian con la escala equivocada mientras el log dice scaling=2.0 [OK].
for clave in ("rank_pattern", "alpha_pattern"):
valor, presente = leer(clave)
if presente and valor:
problems.append(
f"{clave}={valor!r} no esta vacio -- define r/lora_alpha por modulo, y el merge "
f"aplica un unico scaling={EXPECTED_SCALING} a todos los targets"
)
print("[INFO] flags de la variante de LoRA en adapter_config.json:")
for clave in ("use_rslora", "use_dora", "bias", "lora_bias", "modules_to_save", "rank_pattern", "alpha_pattern"):
print(f" {clave:16s} = {reportado[clave]!r}")
def dry_run():
"""Chequeo pre-merge de ~2 segundos: no carga los pesos del modelo base, solo el
adapter (~190MB) y el indice de shards del checkpoint. Aborta ante cualquier
inconsistencia ANTES de que el merge real escriba 67 GB."""
print("[INFO] --dry-run: no se escribe nada, no se cargan los pesos del modelo base")
warn_expected_overrides()
adapter_config = json.loads((ADAPTER_PATH / "adapter_config.json").read_text())
r = adapter_config["r"]
lora_alpha = adapter_config["lora_alpha"]
scaling = lora_alpha / r
print(f"[INFO] r={r} lora_alpha={lora_alpha} scaling={scaling}")
problems = []
if r != EXPECTED_R:
problems.append(f"r={r} (se esperaba {EXPECTED_R})")
if lora_alpha != EXPECTED_LORA_ALPHA:
problems.append(f"lora_alpha={lora_alpha} (se esperaba {EXPECTED_LORA_ALPHA})")
if scaling != EXPECTED_SCALING:
problems.append(f"scaling={scaling} (se esperaba {EXPECTED_SCALING})")
check_adapter_config_flags(adapter_config, problems)
# Targets del adapter + norma de lora_B por familia. PEFT inicializa lora_B en
# CERO EXACTO, asi que una familia con norma cero significa que esos modulos
# nunca recibieron gradiente: es un bug de ENTRENAMIENTO (learning rate, masking,
# target_modules), no del merge -- aunque el sintoma aparezca aca, como el
# AssertionError de "delta cero" que tira merge_shards().
families = {}
zero_modules = []
checkpoint_names = {}
with safe_open(str(ADAPTER_PATH / "adapter_model.safetensors"), framework="pt") as f:
keys = list(f.keys())
base_names = sorted(
k[len(ADAPTER_PREFIX):-len(LORA_A_SUFFIX)] for k in keys if k.endswith(LORA_A_SUFFIX)
)
for base_name in base_names:
key_b = f"{ADAPTER_PREFIX}{base_name}{LORA_B_SUFFIX}"
if key_b not in keys:
problems.append(f"falta {key_b} en el adapter (hay lora_A sin su lora_B)")
continue
norm = f.get_tensor(key_b).to(torch.float32).norm().item()
family = families.setdefault(
module_family(base_name), {"n": 0, "norms": [], "by_module": []}
)
family["n"] += 1
family["norms"].append(norm)
family["by_module"].append((base_name, norm))
if norm == 0.0:
zero_modules.append(base_name)
checkpoint_names[base_name] = remap_adapter_name_to_checkpoint_name(f"{base_name}.weight")
# El conteo de claves se ASIERTE, no solo se imprime: un adapter con
# modules_to_save (u otros tensores extra) daria "310 + 310 = 930", una linea
# aritmeticamente falsa que hoy pasaba en verde. El invariante del plan es
# 310 targets / 620 claves.
expected_keys = 2 * len(base_names)
print(
f"[INFO] {len(base_names)} claves lora_A + {len(base_names)} lora_B = {expected_keys} claves "
f"esperadas, {len(keys)} tensores presentes en el adapter"
)
if len(keys) != expected_keys:
extras = sorted(
k for k in keys if not (k.endswith(LORA_A_SUFFIX) or k.endswith(LORA_B_SUFFIX))
)
problems.append(
f"el adapter tiene {len(keys)} tensores pero {len(base_names)} pares lora_A/lora_B "
f"implican {expected_keys} claves -- hay {len(keys) - expected_keys} tensor(es) de "
f"diferencia. No-lora_A/B encontrados (hasta 10): {extras[:10]}"
)
if expected_keys != 2 * EXPECTED_TARGET_COUNT:
problems.append(
f"claves lora_A/lora_B = {expected_keys}, se esperaban {2 * EXPECTED_TARGET_COUNT} "
f"({EXPECTED_TARGET_COUNT} targets x 2)"
)
print(f"[INFO] targets resueltos: {len(checkpoint_names)} (se esperaban {EXPECTED_TARGET_COUNT})")
if len(checkpoint_names) != EXPECTED_TARGET_COUNT:
problems.append(
f"conteo de targets resueltos = {len(checkpoint_names)}, se esperaban {EXPECTED_TARGET_COUNT}"
)
# Presencia y conteo EXACTO por familia. El total correcto no alcanza: si
# target_modules deja de enganchar linear_attn y engancha otros 150 modulos, el
# total sigue dando 310 y ninguna norma es cero (la familia simplemente no
# aparece), asi que sin este chequeo el dry-run pasa en verde. Ausente no es OK.
print("[INFO] conteo de targets por familia de modulos (esperado vs encontrado):")
for family in sorted(set(EXPECTED_FAMILY_COUNTS) | set(families)):
expected_n = EXPECTED_FAMILY_COUNTS.get(family)
found_n = families.get(family, {}).get("n", 0)
expected_txt = "no prevista" if expected_n is None else str(expected_n)
estado = "OK" if expected_n == found_n else "MAL"
print(f" familia={family:16s} esperados={expected_txt:>11s} encontrados={found_n:4d} [{estado}]")
if expected_n is None:
problems.append(
f"familia {family!r} con {found_n} targets: no esta prevista en EXPECTED_FAMILY_COUNTS "
"-- son modulos que nadie previo y que el merge tocaria igual"
)
elif found_n != expected_n:
if found_n == 0:
problems.append(
f"familia {family}: AUSENTE del adapter (se esperaban {expected_n} targets). "
"Ausente no es OK: ninguna norma da cero porque la familia ni siquiera esta, "
"asi que el chequeo de normas no lo veria. Revisar TARGET_MODULES"
)
elif expected_n == 0:
ejemplos = [n for n in base_names if module_family(n) == family][:10]
problems.append(
f"familia {family}: {found_n} targets donde se esperaban 0 -- el catch-all de "
f"module_family() no debe atrapar nada. Ejemplos: {ejemplos}"
)
else:
problems.append(
f"familia {family}: {found_n} targets, se esperaban {expected_n}"
)
print("[INFO] norma de lora_B por familia de modulos:")
for family in sorted(families):
stats = families[family]
norms = stats["norms"]
print(
f" familia={family:16s} n={stats['n']:4d} "
f"norm_total={sum(norms):10.4f} norm_min={min(norms):.6f} "
f"norm_max={max(norms):.6f} norm_avg={sum(norms) / len(norms):.6f}"
)
if max(norms) == 0.0:
problems.append(
f"familia {family}: TODAS las normas de lora_B son cero -- esos modulos nunca "
"recibieron gradiente. Es un bug de ENTRENAMIENTO (learning rate, masking o "
"target_modules), NO del merge"
)
# "Efectivamente muerto", no solo cero exacto: un modulo con norma 1e-12 frente a
# una mediana de familia de 1e-1 no aporta nada al merge, pero norm == 0.0 (igualdad
# exacta de float) no lo atrapa.
dead_modules = []
for family, stats in families.items():
norms = sorted(stats["norms"])
median = norms[len(norms) // 2]
if median <= 0.0:
continue
floor = DEAD_MODULE_RELATIVE_THRESHOLD * median
for base_name, norm in stats["by_module"]:
if 0.0 < norm < floor:
dead_modules.append((base_name, family, norm, median))
if dead_modules:
print(f"[WARN] {len(dead_modules)} modulos con norma de lora_B efectivamente muerta:")
for base_name, family, norm, median in dead_modules[:20]:
print(f" {base_name} (familia={family}, norm={norm:.3e}, mediana de familia={median:.3e})")
if len(dead_modules) > 20:
print(f" ... y {len(dead_modules) - 20} mas")
problems.append(
f"{len(dead_modules)} modulos con ||lora_B|| < {DEAD_MODULE_RELATIVE_THRESHOLD:g} x la "
"mediana de su familia -- practicamente sin gradiente. Mismo diagnostico que la norma "
"cero: es un bug de ENTRENAMIENTO, no del merge"
)
if zero_modules:
print(f"[WARN] {len(zero_modules)} modulos con norma de lora_B EXACTAMENTE cero:")
for base_name in zero_modules[:20]:
print(f" {base_name}")
if len(zero_modules) > 20:
print(f" ... y {len(zero_modules) - 20} mas")
problems.append(
f"{len(zero_modules)} modulos con ||lora_B|| == 0 -- el merge abortaria con 'delta cero'. "
"Es un bug de ENTRENAMIENTO, no del merge"
)
# Que cada clave remapeada exista en el indice del checkpoint base: es el chequeo
# que evita descubrir un mismatch de nombres recien despues de escribir 67 GB.
index = json.loads((MODEL_PATH / "model.safetensors.index.json").read_text())
weight_map = index["weight_map"]
matched = [n for n in checkpoint_names.values() if n in weight_map]
unmatched = sorted(n for n in checkpoint_names.values() if n not in weight_map)
print(
f"[INFO] claves del adapter presentes en model.safetensors.index.json: "
f"{len(matched)}/{len(checkpoint_names)} ({len(weight_map)} tensores en el indice)"
)
if unmatched:
print(f"[ERROR] {len(unmatched)} claves remapeadas NO existen en el checkpoint base:")
for name in unmatched[:20]:
print(f" {name}")
if len(unmatched) > 20:
print(f" ... y {len(unmatched) - 20} mas")
problems.append(f"{len(unmatched)} claves remapeadas ausentes del indice del checkpoint base")
if problems:
raise AssertionError("dry-run FALLIDO:\n - " + "\n - ".join(problems))
print("[OK] dry-run: todas las aserciones pasaron, el merge real puede correr")
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument(
"--dry-run",
action="store_true",
help=(
"verificar el adapter y el remapeo de claves contra el indice del checkpoint base "
"SIN cargar pesos ni escribir nada (~2s), en vez de correr el merge de 67 GB"
),
)
return parser.parse_args()
def main():
args = parse_args()
print(f"[INFO] MODEL_PATH={MODEL_PATH}")
print(f"[INFO] ADAPTER_PATH={ADAPTER_PATH}")
print(f"[INFO] OUTPUT_PATH={OUTPUT_PATH}")
if args.dry_run:
dry_run()
return
deltas = load_lora_deltas()
t0 = time.time()
nontarget_sample = merge_shards(deltas)
copy_non_model_files()
verify_chat_template_is_not_training_template()
verify_nontarget_byte_identical(nontarget_sample)
print(f"[INFO] merge completo en {time.time() - t0:.1f}s. OUTPUT_PATH={OUTPUT_PATH}")
if __name__ == "__main__":
main()