"""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) n_assistant_tokens = sum(assistant_masks) return out.loss.item(), n_assistant_tokens 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() # Cada entrada es (loss_del_ejemplo, n_tokens_assistant_del_ejemplo) -- se necesitan # ambos para poder reportar tanto el promedio simple por ejemplo (util para comparar # buckets entre si) como el promedio ponderado por token (comparable directamente # contra el eval_loss que reporta transformers.Trainer, que pondera por cantidad de # tokens validos y no por cantidad de ejemplos -- un bucket con pocos ejemplos pero # secuencias largas/dificiles no debe pesar igual que uno con muchos ejemplos cortos). losses_by_bucket = defaultdict(list) for i, example in enumerate(examples): bucket = example.get("meta", {}).get("bucket", "sin_bucket") loss, n_tokens = compute_loss_per_example(model, tokenizer, example) losses_by_bucket[bucket].append((loss, n_tokens)) 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) def weighted_avg(pairs): total_tokens = sum(n for _, n in pairs) return sum(loss * n for loss, n in pairs) / total_tokens def simple_avg(pairs): return sum(loss for loss, _ in pairs) / len(pairs) all_pairs = [pair for pairs in losses_by_bucket.values() for pair in pairs] global_avg_simple = simple_avg(all_pairs) global_avg_weighted = weighted_avg(all_pairs) 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): pairs = losses_by_bucket[bucket] n_tokens_total = sum(n for _, n in pairs) print( f" bucket={bucket:20s} n={len(pairs):4d} tokens={n_tokens_total:6d} " f"loss_avg_simple={simple_avg(pairs):.4f} loss_avg_weighted={weighted_avg(pairs):.4f}" ) replay_pairs = losses_by_bucket.get("replay") if replay_pairs: print( f" bucket=replay (aislado) n={len(replay_pairs):4d} " f"loss_avg_simple={simple_avg(replay_pairs):.4f} loss_avg_weighted={weighted_avg(replay_pairs):.4f}" ) print(f"\n loss_avg GLOBAL simple (por ejemplo) = {global_avg_simple:.4f}") print(f" loss_avg GLOBAL ponderado (por token) = {global_avg_weighted:.4f}") print(f" eval_loss Fase 3 (adapter puro, Trainer, ponderado por token) = {FASE3_EVAL_LOSS:.4f}") diff = abs(global_avg_weighted - FASE3_EVAL_LOSS) print(f" diferencia absoluta (ponderado vs Fase 3) = {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()