"""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()