diff --git a/scripts/20_merge_lora.py b/scripts/20_merge_lora.py new file mode 100644 index 0000000..ea13056 --- /dev/null +++ b/scripts/20_merge_lora.py @@ -0,0 +1,225 @@ +"""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. +""" +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" + +# 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 main(): + print(f"[INFO] MODEL_PATH={MODEL_PATH}") + print(f"[INFO] ADAPTER_PATH={ADAPTER_PATH}") + print(f"[INFO] OUTPUT_PATH={OUTPUT_PATH}") + + 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()