- 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.
147 lines
5.4 KiB
Python
147 lines
5.4 KiB
Python
"""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()
|