Phase 4: merge LoRA adapter and run full evaluation gates on the merged checkpoint #3
+37
-16
@@ -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, "
|
||||
|
||||
Reference in New Issue
Block a user