From 38546c03d1f76bd1e2a92316dcd851ded00f10f1 Mon Sep 17 00:00:00 2001 From: wassname <1103714+wassname@users.noreply.github.com> Date: Thu, 15 Jan 2026 12:28:24 +0800 Subject: [PATCH] add val loss to info --- antipasto/train/train_adapter.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/antipasto/train/train_adapter.py b/antipasto/train/train_adapter.py index 0e4033d..b7c54ed 100644 --- a/antipasto/train/train_adapter.py +++ b/antipasto/train/train_adapter.py @@ -873,6 +873,10 @@ def compute_validation_loss( avg_total = total_loss / n_batches if n_batches > 0 else float("inf") avg_components = {k: np.mean(v) for k, v in loss_components.items() if not isinstance(v[0], str)} + # Mark all val infos with phase='val' for later filtering + for info in all_infos: + info['phase'] = 'val' + # Extract per-coefficient breakdown (log validation table inline) df_coef, coef_metrics = extract_coef_metrics( all_infos, log_table=log_tables, @@ -880,7 +884,7 @@ def compute_validation_loss( ) if all_infos else (None, {}) val_summary = summarize_phase_for_compare(all_infos, phase="val") - return avg_total, avg_components, df_coef, coef_metrics, val_summary + return avg_total, avg_components, df_coef, coef_metrics, val_summary, all_infos def train_epoch( @@ -934,6 +938,9 @@ def train_epoch( loss_subspace=loss_subspace, scale_adapter_fn=scale_adapter_fn, ) + # Mark train infos with phase='train' for filtering in df_hist + for info in batch_infos: + info['phase'] = 'train' infos.extend(batch_infos) # Epoch-start snapshot: print per-coef table on the very first batch. @@ -997,12 +1004,14 @@ def train_epoch( if val_dataloader is not None and opt_step % val_n_steps == 0 and opt_step > 0 and step % config.grad_accum_steps == 0: # Keep validation compute cadence (early stopping + wandb), but don't spam tables. log_val = False - val_loss, val_components, val_df_coef, val_coef_metrics, val_summary = compute_validation_loss( + val_loss, val_components, val_df_coef, val_coef_metrics, val_summary, val_infos = compute_validation_loss( model=model, val_dataloader=val_dataloader, loss_layers=loss_layers, loss_layer_indices=loss_layer_indices, config=config, step=opt_step, loss_subspace=loss_subspace, flip_stats=flip_stats, scale_adapter_fn=scale_adapter_fn, log_tables=log_val, total_steps=total_steps, ) + # Add val infos to main infos list for saving to df_hist + infos.extend(val_infos) last_val_summary = val_summary last_val_loss = val_loss @@ -1073,7 +1082,7 @@ def train_epoch( # No validation ran during the epoch (e.g., tiny epoch). Compute once here. # We do not print per-coef tables here to keep epoch-end output compact. step_for_log = (epoch + 1) * len(train_dataloader) - 1 - val_loss, _, _, _, last_val_summary = compute_validation_loss( + val_loss, _, _, _, last_val_summary, val_infos = compute_validation_loss( model=model, val_dataloader=val_dataloader, loss_layers=loss_layers, @@ -1086,6 +1095,8 @@ def train_epoch( scale_adapter_fn=scale_adapter_fn, log_tables=False, ) + # Add val infos to main infos list for saving to df_hist + infos.extend(val_infos) last_val_loss = val_loss last_val_step = step_for_log