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.
This commit is contained in:
2026-07-29 17:49:08 +00:00
parent ceba5f80cd
commit 2762761a8e
+37 -16
View File
@@ -72,7 +72,8 @@ def compute_loss_per_example(model, tokenizer, example):
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()
n_assistant_tokens = sum(assistant_masks)
return out.loss.item(), n_assistant_tokens
def run_gate1():
@@ -98,35 +99,55 @@ def run_gate1():
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 = compute_loss_per_example(model, tokenizer, example)
losses_by_bucket[bucket].append(loss)
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)
all_losses = [loss for losses in losses_by_bucket.values() for loss in losses]
global_avg = sum(all_losses) / len(all_losses)
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):
losses = losses_by_bucket[bucket]
avg = sum(losses) / len(losses)
print(f" bucket={bucket:20s} n={len(losses):4d} loss_avg={avg:.4f}")
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_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}")
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 (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}")
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, "