Files
qwen3-6-lora/scripts/30_eval_suite.py
T
aleleba 2762761a8e Fase 4: puerta 1 - reportar tambien loss ponderado por token (comparable a Trainer)
El primer resultado (promedio simple por ejemplo) daba 0.5185 vs 0.275 de Fase 3, señal de
alarma segun el propio script. La causa era metodologica, no un bug de merge: el bucket
replay concentra 112927 de los ~128849 tokens assistant del split de eval (87%), mientras
que buckets dificiles como negativos/skills_adherencia/delegacion_subagentes tienen pocos
ejemplos pero loss alto -- un promedio por ejemplo les da el mismo peso que a replay,
inflando el global. transformers.Trainer pondera por token, no por ejemplo. Con el mismo
ponderado por token: 0.2560 vs 0.275 de Fase 3 (diff=0.019, dentro del margen esperado) --
confirma que el merge es correcto.
2026-07-29 17:49:08 +00:00

168 lines
6.6 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)
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()