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, "