Fase 4: puerta 1 (eval-loss offline por bucket) y contenedor/scripts de puertas 2-4

- scripts/30_eval_suite.py --gate 1: eval-loss sobre el checkpoint mergeado, agrupado por
  meta.bucket (aislando replay), comparado contra eval_loss=0.275 de Fase 3.
- docker-compose.eval.yml: servicio vllm-eval propio (puerto 8001), sirviendo el
  checkpoint mergeado en BF16, con tool-call-parser=qwen3_coder y reasoning-parser=qwen3.
  No se pudo leer el compose real de produccion (/data/compose/43/docker-compose.yml no
  existe en spark, probablemente vive en el host del servidor Portainer) -- flags basados
  en la arquitectura conocida del modelo.
- scripts/31_build_holdout_prompts.py: genera data/holdout_prompts.jsonl (200 prompts,
  40 por MCP, sin overlap verificado contra train.jsonl/eval.jsonl).
- scripts/32_gate2_toolcalls.py: valida tool-calls devueltas por vllm-eval (parser real
  de vLLM, nunca una regex propia) contra los 200 prompts held-out.
- scripts/33_gate3_adherencia.py: checklists de adherencia por skill + no-activacion,
  con baseline opcional contra vllm-qwen36 si esta corriendo.
- scripts/34_gate4_e2e.py: arma el plan de llamadas E2E contra los 5 MCPs y 5 skills via
  el checkpoint mergeado, para que el agente orquestador las ejecute con sus MCPs reales.
This commit is contained in:
2026-07-29 17:37:02 +00:00
parent 268451aed7
commit ceba5f80cd
7 changed files with 1090 additions and 0 deletions
+146
View File
@@ -0,0 +1,146 @@
"""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()