From 2762761a8e9df48bb26d21b792a2a827998d95a6 Mon Sep 17 00:00:00 2001 From: Alejandro Lembke Barrientos Date: Wed, 29 Jul 2026 17:49:08 +0000 Subject: [PATCH] Fase 4: puerta 1 - reportar tambien loss ponderado por token (comparable a Trainer) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- scripts/30_eval_suite.py | 53 ++++++++++++++++++++++++++++------------ 1 file changed, 37 insertions(+), 16 deletions(-) diff --git a/scripts/30_eval_suite.py b/scripts/30_eval_suite.py index aeb3841..dfded3f 100644 --- a/scripts/30_eval_suite.py +++ b/scripts/30_eval_suite.py @@ -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, "