import os from pytorch_lightning import LightningModule, Trainer from pytorch_lightning.callbacks import Callback, RichProgressBar, ModelCheckpoint def build_callbacks(cfg, logger=None, phase='test', **kwargs): callbacks = [] logger = logger # Rich Progress Bar callbacks.append(progressBar()) # Checkpoint Callback if phase == 'train': callbacks.extend(getCheckpointCallback(cfg, logger=logger, **kwargs)) return callbacks def getCheckpointCallback(cfg, logger=None, **kwargs): callbacks = [] # Logging metric_monitor = { "loss_total": "total/train", "Train_jf": "recons/text2jfeats/train", "Val_jf": "recons/text2jfeats/val", "Train_rf": "recons/text2rfeats/train", "Val_rf": "recons/text2rfeats/val", "APE root": "Metrics/APE_root", "APE mean pose": "Metrics/APE_mean_pose", "AVE root": "Metrics/AVE_root", "AVE mean pose": "Metrics/AVE_mean_pose", "R_TOP_1": "Metrics/R_precision_top_1", "R_TOP_2": "Metrics/R_precision_top_2", "R_TOP_3": "Metrics/R_precision_top_3", "gt_R_TOP_3": "Metrics/gt_R_precision_top_3", "FID": "Metrics/FID", "gt_FID": "Metrics/gt_FID", "Diversity": "Metrics/Diversity", "MM dist": "Metrics/Matching_score", "Accuracy": "Metrics/accuracy", } callbacks.append( progressLogger(logger,metric_monitor=metric_monitor,log_every_n_steps=1)) chekpoint_every_steps = cfg.LOGGER.VAL_EVERY_STEPS monitor = "step" mode = "max" save_top_k = cfg.LOGGER.SAVE_TOP_K if cfg.LOGGER.get('CKPT_EVERY_EPOCHS', False) : chekpoint_every_steps = cfg.LOGGER.CKPT_EVERY_EPOCHS monitor = "total/train" mode = "min" save_top_k = -1 # Save 10 latest checkpoints checkpointParams = { 'dirpath': os.path.join(cfg.FOLDER_EXP, "checkpoints"), 'filename': "{epoch}", # 'monitor': "step", # 'mode': "max", # 'every_n_epochs': cfg.LOGGER.VAL_EVERY_STEPS, # 'save_top_k': cfg.LOGGER.SAVE_TOP_K, 'monitor': monitor, 'mode': mode, 'every_n_epochs': chekpoint_every_steps, 'save_top_k': save_top_k, 'save_last': True, 'save_on_train_epoch_end': True } callbacks.append(ModelCheckpoint(**checkpointParams)) # Save checkpoint every n*10 epochs if not cfg.LOGGER.get('CKPT_EVERY_EPOCHS', False): checkpointParams.update({ 'every_n_epochs': cfg.LOGGER.VAL_EVERY_STEPS * 10, 'save_top_k': -1, 'save_last': False }) callbacks.append(ModelCheckpoint(**checkpointParams)) # Step-based checkpointing (independent of val/epoch). # GRPO has long epochs (8765 steps each) so we need intermediate saves. # Enable via LOGGER.CKPT_EVERY_N_STEPS in the yaml. monitor=None + # save_top_k=-1 makes every trigger save unconditionally — the previous # monitor='step' variant silently skipped saves in PL 2.0 because # 'step' isn't an explicitly logged metric. ckpt_every_n_steps = int(cfg.LOGGER.get('CKPT_EVERY_N_STEPS', 0) or 0) print(f"[CKPT-DEBUG] CKPT_EVERY_N_STEPS = {ckpt_every_n_steps}", flush=True) if ckpt_every_n_steps > 0: step_ckpt_params = { 'dirpath': os.path.join(cfg.FOLDER_EXP, "checkpoints"), 'filename': "step-{step}", 'every_n_train_steps': ckpt_every_n_steps, 'save_top_k': -1, # save every trigger (no metric ranking) 'save_last': True, # also keep last.ckpt for easy resume 'save_on_train_epoch_end': False, } print(f"[CKPT-DEBUG] Adding step-based ModelCheckpoint with {step_ckpt_params}", flush=True) callbacks.append(ModelCheckpoint(**step_ckpt_params)) metrics = cfg.METRIC.TYPE metric_monitor_map = { 'TemosMetric': { 'Metrics/APE_root': { 'abbr': 'APEroot', 'mode': 'min' }, }, 'TM2TMetrics': { 'Metrics/FID': { 'abbr': 'FID', 'mode': 'min' }, 'Metrics/R_precision_top_3': { 'abbr': 'R3', 'mode': 'max' } }, 'T2MMetrics': { 'Metrics/T2M_FID': { 'abbr': 'T2M_FID', 'mode': 'min' }, 'Metrics/T2M_ADE': { 'abbr': 'T2M_ADE', 'mode': 'min' }, }, # 'M2TMetrics': { # 'Metrics/M2T_Bleu_4': { # 'abbr': 'M2T_Bleu_4', # 'mode': 'max' # }, # 'Metrics/Bleu_4': { # 'abbr': 'Bleu_4', # 'mode': 'max' # } # }, 'T2TMetrics': { 'Metrics/Bleu_1': { 'abbr': 'Bleu_1', 'mode': 'max' } }, 'ObstacleMetrics': { 'Metrics/FreeSpace_Acc_mean': { 'abbr': 'Acc', 'mode': 'max' } }, 'PredMetrics': { 'Metrics/ADE_local': { 'abbr': 'ADE', 'mode': 'min' }, 'Metrics/FID_local': { 'abbr': 'FID', 'mode': 'min' } }, 'PredMetrics_c': { 'Metrics/ADE_local': { 'abbr': 'ADE_local', 'mode': 'min' } }, 'PredMetrics_f': { # 'Metrics/ADE': { # 'abbr': 'ADE', # 'mode': 'min' # }, # 'Metrics/ADE_local': { # 'abbr': 'ADE_local', # 'mode': 'min' # }, 'Metrics/ADE_head': { 'abbr': 'ADE_head', 'mode': 'min' }, # 'Metrics/FID': { # 'abbr': 'FID', # 'mode': 'min' # } }, 'MRMetrics': { 'Metrics/MPJPE': { 'abbr': 'MPJPE', 'mode': 'min' } }, 'HUMANACTMetrics': { 'Metrics/Accuracy': { 'abbr': 'Accuracy', 'mode': 'max' } }, 'UESTCMetrics': { 'Metrics/Accuracy': { 'abbr': 'Accuracy', 'mode': 'max' } }, 'UncondMetrics': { 'Metrics/FID': { 'abbr': 'FID', 'mode': 'min' } } } checkpointParams.update({ 'every_n_epochs': cfg.LOGGER.VAL_EVERY_STEPS, 'save_top_k': 1, }) ONLY_MOTION = cfg['LOSS'].ABLATION.get("ONLY_MOTION", False) ONLY_FUTURE = cfg['LOSS'].ABLATION.get("ONLY_FUTURE", False) EGOLM_BASELINE = cfg['LOSS'].ABLATION.get("EGOLM_BASELINE", False) for metric in metrics: if metric in metric_monitor_map.keys(): if metric == "PredMetrics": task = getattr(cfg.model.params, "task", "") if task == "egovlm_4tasks": if cfg['LOSS'].ABLATION.EGOLM_BASELINE : metric = "PredMetrics_c" elif cfg['LOSS'].ABLATION.ONLY_MOTION: metric = "PredMetrics_c" elif cfg['LOSS'].ABLATION.ONLY_FUTURE: metric = "PredMetrics_f" else : metric = "PredMetrics_f" metric_monitors = dict(metric_monitor_map[metric]) # Delete R3 if training VAE if cfg.TRAIN.STAGE == 'vae' and metric == 'TM2TMetrics': del metric_monitors['Metrics/R_precision_top_3'] if cfg.TRAIN.STAGE == 'lm_pretrain' and metric == 'TM2TMetrics' and cfg.model.params.task == 'pred': del metric_monitors['Metrics/R_precision_top_3'] if cfg.TRAIN.STAGE == 'lm_instruct' and metric == 'TM2TMetrics': del metric_monitors['Metrics/R_precision_top_3'] if metric == 'M2TMetrics': task = getattr(cfg.model.params, "task", "") base_monitor = 'Metrics/M2T_Bleu_4' # base_monitor = 'Metrics/Bleu_4' IT_joint_training = cfg['LOSS'].ABLATION.get("IT_JOINT_TRAINING", False) EGOLM_BASELINE = cfg['LOSS'].ABLATION.get("EGOLM_BASELINE", False) ONLY_VIDEO = cfg['LOSS'].ABLATION.get("ONLY_VIDEO", False) TASKS_4 = cfg['LOSS'].ABLATION.get("4TASKS", False) if ONLY_VIDEO : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_c/M2T_Bleu_4' elif TASKS_4 : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_f/M2T_Bleu_4' elif ONLY_MOTION and IT_joint_training : egovlm_monitor = 'Metrics/M2TMetrics_stp2mt_c/M2T_Bleu_4' elif ONLY_MOTION and not IT_joint_training : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_c/M2T_Bleu_4' elif ONLY_MOTION and not IT_joint_training: egovlm_monitor = 'Metrics/M2TMetrics_stp2t_c/M2T_Bleu_4' elif ONLY_FUTURE and not IT_joint_training : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_f/M2T_Bleu_4' elif IT_joint_training : egovlm_monitor = 'Metrics/M2TMetrics_stp2mt_f/M2T_Bleu_4' elif cfg['LOSS'].ABLATION.get("IT_MOTION_AT_ONCE", False) : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_f/Bleu_4' elif EGOLM_BASELINE : if ONLY_MOTION : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_c/Bleu_4' else : egovlm_monitor = 'Metrics/M2TMetrics_stp2t_c/Bleu_4' else : egovlm_monitor = 'Metrics/M2TMetrics_stp2mt_f/M2T_Bleu_4' monitor_key = egovlm_monitor if task == "egovlm_4tasks" else base_monitor metric_monitors = { monitor_key: metric_monitor_map['M2TMetrics'][base_monitor] } if metric == "PredMetrics_c" : task = getattr(cfg.model.params, "task", "") remapped_monitors = {} for key, value in metric_monitors.items(): suffix = key.split('/', 1)[1] if '/' in key else key IT_joint_training = cfg['LOSS'].ABLATION.get("IT_JOINT_TRAINING", True) ONLY_VIDEO = cfg['LOSS'].ABLATION.get("ONLY_VIDEO", False) NO_TEXT = cfg['LOSS'].ABLATION.get("NO_TEXT", False) if ONLY_VIDEO or NO_TEXT : remapped_monitors[f"Metrics/PredMetrics_stp2m_c/{suffix}"] = value elif cfg['LOSS'].ABLATION.EGOLM_BASELINE : # With IT_JOINT_TRAINING the EGOLM_BASELINE val path logs its # PredMetrics under the joint m+t subtask name (stp2mt_c), not # stp2m_c — monitoring stp2m_c crashes ModelCheckpoint at the # first val end (key not found in callback metrics). _pfx = 'stp2mt_c' if IT_joint_training else 'stp2m_c' remapped_monitors[f"Metrics/PredMetrics_{_pfx}/{suffix}"] = value else : remapped_monitors[f"Metrics/PredMetrics_stp2mt_c/{suffix}"] = value metric_monitors = remapped_monitors if metric == 'PredMetrics_f': task = getattr(cfg.model.params, "task", "") if not getattr(cfg.METRIC, "DIVERSITY", False): metric_monitors.pop('Metrics/FID', None) remapped_monitors = {} for key, value in metric_monitors.items(): suffix = key.split('/', 1)[1] if '/' in key else key IT_joint_training = cfg['LOSS'].ABLATION.get("IT_JOINT_TRAINING", True) NO_TEXT = cfg['LOSS'].ABLATION.get("NO_TEXT", False) # ADE/ADE_head/FDE* are ORACLE best-of-K (custom.py: best_local_idx = dists.argmin() # against ground truth, K=num_sample_seq). Selecting on the oracle metric picks a # DIFFERENT epoch than the honest one (measured: fold0 val5 vs val1) and costs the # baseline +5.4% there. LOGGER.HONEST_SELECTION switches selection to the *_idx0 # single-shot metric, which is what we actually report. _hs = cfg.LOGGER.get('HONEST_SELECTION', False) _sfx = suffix + '_idx0' if (_hs and not suffix.endswith('_idx0')) else suffix if _sfx != suffix: # keep the checkpoint FILENAME honest too: it is built from 'abbr', so without this # an idx0-selected ckpt would still be named min-ADE_head-... and misstate its criterion value = dict(value); value['abbr'] = value['abbr'] + '_idx0' if cfg['LOSS'].ABLATION.EGOLM_BASELINE: remapped_monitors[f"Metrics/PredMetrics_stp2m_f/{_sfx}"] = value elif not IT_joint_training: remapped_monitors[f"Metrics/PredMetrics_stp2m_f/{_sfx}"] = value elif NO_TEXT : remapped_monitors[f"Metrics/PredMetrics_stp2m_f/{_sfx}"] = value else : remapped_monitors[f"Metrics/PredMetrics_stp2mt_f/{_sfx}"] = value metric_monitors = remapped_monitors if metric == 'PredMetrics': task = getattr(cfg.model.params, "task", "") if not getattr(cfg.METRIC, "DIVERSITY", False): metric_monitors.pop('Metrics/FID', None) # if task == "egovlm_4tasks": # remapped_monitors = {} # for key, value in metric_monitors.items(): # suffix = key.split('/', 1)[1] if '/' in key else key # if cfg['LOSS'].ABLATION.ONLY_MOTION : # remapped_monitors[f"Metrics/PredMetrics_stp2mt_c/{suffix}"] = value # elif cfg['LOSS'].ABLATION.EGOLM_BASELINE : # remapped_monitors[f"Metrics/PredMetrics_stp2m_c/{suffix}"] = value # else : # remapped_monitors[f"Metrics/PredMetrics_stp2mt_f/{suffix}"] = value # metric_monitors = remapped_monitors # metric_monitors.pop("Metrics/PredMetrics_stp2m_c/ADE") # metric_monitors.pop("Metrics/PredMetrics_stp2m_c/FID") for metric_monitor, monitor_cfg in metric_monitors.items(): checkpointParams.update({ 'filename': monitor_cfg['mode'] + "-" + monitor_cfg['abbr'] + "-{epoch}-{step}", 'monitor': metric_monitor, 'mode': monitor_cfg['mode'], 'save_on_train_epoch_end': False, 'every_n_epochs': 1, }) callbacks.append( ModelCheckpoint(**checkpointParams)) # --- narration-criterion checkpoints (LOGGER.NARRATION_CKPT) ------------------------- # The table scores EgoLM on UNDERSTANDING (verb F1 from generated narration), but the only # metric-selected checkpoint above is min-ADE_local (pose). 'M2TMetrics' is commented out of # metric_monitor_map, so no narration checkpoint is ever written and the deployed ckpt is # chosen by a criterion that is still improving while narration has already peaked. # These extra callbacks keep the narration-best checkpoints so selection can match the metric # we actually report. stp2mt_c = CURRENT-window narration, which is what verb F1 scores. # Bleu_1 is the closer proxy for verb F1 (unigram lemma presence); Bleu_4 is kept too. if cfg.LOGGER.get('NARRATION_CKPT', False): for _key, _abbr in (('Metrics/M2TMetrics_stp2mt_c/M2T_Bleu_1', 'M2T_Bleu_1'), ('Metrics/M2TMetrics_stp2mt_c/M2T_Bleu_4', 'M2T_Bleu_4')): callbacks.append(ModelCheckpoint( dirpath=os.path.join(cfg.FOLDER_EXP, 'checkpoints'), filename='max-' + _abbr + '-{epoch}', monitor=_key, mode='max', save_top_k=1, save_last=False, save_on_train_epoch_end=False, every_n_epochs=1, )) return callbacks class progressBar(RichProgressBar): def __init__(self, ): super().__init__() def get_metrics(self, trainer, model): # Don't show the version number items = super().get_metrics(trainer, model) items.pop("v_num", None) return items class progressLogger(Callback): def __init__(self, logger, metric_monitor: dict, precision: int = 3, log_every_n_steps: int = 1): # Metric to monitor self.logger = logger self.metric_monitor = metric_monitor self.precision = precision self.log_every_n_steps = log_every_n_steps def on_train_start(self, trainer: Trainer, pl_module: LightningModule, **kwargs) -> None: self.logger.info("Training started") def on_train_end(self, trainer: Trainer, pl_module: LightningModule, **kwargs) -> None: self.logger.info("Training done") def on_validation_epoch_end(self, trainer: Trainer, pl_module: LightningModule, **kwargs) -> None: if trainer.sanity_checking: self.logger.info("Sanity checking ok.") def on_train_epoch_end(self, trainer: Trainer, pl_module: LightningModule, padding=False, **kwargs) -> None: metric_format = f"{{:.{self.precision}e}}" line = f"Epoch {trainer.current_epoch}" if padding: line = f"{line:>{len('Epoch xxxx')}}" # Right padding if trainer.current_epoch % self.log_every_n_steps == 0: metrics_str = [] losses_dict = trainer.callback_metrics for metric_name, dico_name in self.metric_monitor.items(): if dico_name in losses_dict: metric = losses_dict[dico_name].item() metric = metric_format.format(metric) metric = f"{metric_name} {metric}" metrics_str.append(metric) line = line + ": " + " ".join(metrics_str) self.logger.info(line)