"""Fase 4 -- suite de evaluacion en 4 puertas. Puerta 1 (--gate 1): eval-loss offline por bucket sobre el checkpoint MERGEADO (no el adapter puro) -- no necesita servir el modelo. 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/30_eval_suite.py --gate 1 Carga el checkpoint mergeado con AutoModelForCausalLM (para detectar bugs de merge que un eval sobre el adapter puro no veria), le pisa en memoria el chat_template con data/chat_template_train.jinja (igual que en training, para poder generar assistant_masks), recorre data/eval.jsonl agrupado por meta.bucket, y reporta loss promedio global y por bucket (aislando bucket=="replay"), comparado contra eval_loss=0.275 de Fase 3. Las puertas 2-4 (tool-calls, adherencia, E2E) viven en scripts separados (scripts/31_gate2_toolcalls.py, scripts/32_gate3_adherencia.py, scripts/33_gate4_e2e.py) porque necesitan el contenedor de eval sirviendo el checkpoint mergeado via HTTP, no solo lectura offline. """ import argparse import json import os import time from collections import defaultdict from pathlib import Path import torch from transformers import AutoModelForCausalLM, AutoTokenizer REPO_ROOT = Path(__file__).resolve().parent.parent OUTPUT_PATH = os.environ.get("OUTPUT_PATH", "/workspace/ft-models/Qwen3.6-35B-A3B-mcp-bf16") TRAIN_CHAT_TEMPLATE_PATH = REPO_ROOT / "data" / "chat_template_train.jinja" EVAL_FILE = REPO_ROOT / "data" / "eval.jsonl" FASE3_EVAL_LOSS = 0.275 def parse_args(): parser = argparse.ArgumentParser() parser.add_argument("--gate", type=int, required=True, choices=[1]) return parser.parse_args() def load_eval_examples(): examples = [] with open(EVAL_FILE, encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue examples.append(json.loads(line)) return examples def compute_loss_per_example(model, tokenizer, example): rendered = tokenizer.apply_chat_template( example["messages"], tools=example.get("tools"), tokenize=True, return_assistant_tokens_mask=True, return_dict=True, add_generation_prompt=False, ) input_ids = rendered["input_ids"] assistant_masks = rendered["assistant_masks"] if sum(assistant_masks) == 0: raise AssertionError("assistant_masks vacia para un ejemplo de eval.jsonl") labels = [tok if mask == 1 else -100 for tok, mask in zip(input_ids, assistant_masks)] input_ids_t = torch.tensor([input_ids], dtype=torch.long, device=model.device) labels_t = torch.tensor([labels], dtype=torch.long, device=model.device) with torch.no_grad(): out = model(input_ids=input_ids_t, labels=labels_t) return out.loss.item() def run_gate1(): print(f"[INFO] cargando checkpoint mergeado desde {OUTPUT_PATH}") tokenizer = AutoTokenizer.from_pretrained(OUTPUT_PATH) tokenizer.chat_template = TRAIN_CHAT_TEMPLATE_PATH.read_text(encoding="utf-8") if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token t0 = time.time() model = AutoModelForCausalLM.from_pretrained( OUTPUT_PATH, dtype=torch.bfloat16, attn_implementation="flash_attention_2", ) model = model.to("cuda") model.eval() load_time = time.time() - t0 print(f"[INFO] modelo cargado en {load_time:.1f}s") examples = load_eval_examples() print(f"[INFO] {len(examples)} ejemplos en {EVAL_FILE}") torch.cuda.reset_peak_memory_stats() t0 = time.time() losses_by_bucket = defaultdict(list) for i, example in enumerate(examples): bucket = example.get("meta", {}).get("bucket", "sin_bucket") loss = compute_loss_per_example(model, tokenizer, example) losses_by_bucket[bucket].append(loss) if (i + 1) % 25 == 0: print(f"[INFO] {i + 1}/{len(examples)} ejemplos evaluados") eval_time = time.time() - t0 peak_mem_gb = torch.cuda.max_memory_allocated() / (1024 ** 3) all_losses = [loss for losses in losses_by_bucket.values() for loss in losses] global_avg = sum(all_losses) / len(all_losses) print("\n=== Puerta 1 -- eval-loss offline por bucket (checkpoint mergeado) ===") print(f"[INFO] tiempo de eval: {eval_time:.1f}s, memoria pico: {peak_mem_gb:.2f} GB") for bucket in sorted(losses_by_bucket): losses = losses_by_bucket[bucket] avg = sum(losses) / len(losses) print(f" bucket={bucket:20s} n={len(losses):4d} loss_avg={avg:.4f}") replay_losses = losses_by_bucket.get("replay") if replay_losses: replay_avg = sum(replay_losses) / len(replay_losses) print(f" bucket=replay (aislado) n={len(replay_losses):4d} loss_avg={replay_avg:.4f}") print(f"\n loss_avg GLOBAL (checkpoint mergeado) = {global_avg:.4f}") print(f" eval_loss Fase 3 (adapter puro, sanity) = {FASE3_EVAL_LOSS:.4f}") diff = abs(global_avg - FASE3_EVAL_LOSS) print(f" diferencia absoluta = {diff:.4f}") if diff > 0.05: print( " [WARN] diferencia > 0.05 -- senal posible de bug real en el merge, " "revisar antes de continuar a la puerta 2" ) else: print(" [OK] loss del checkpoint mergeado consistente con Fase 3 -- merge probablemente correcto") def main(): args = parse_args() if args.gate == 1: run_gate1() if __name__ == "__main__": main()