Phase 4: merge LoRA adapter and run full evaluation gates on the merged checkpoint #3
@@ -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()
|
||||
Reference in New Issue
Block a user