Download mGPT/callback.py from po03087/egolm-protocol-v2-code: direct link, hf CLI and curl.
- Browser
- Download file 19.4 kB
-
https://huggingface.co/po03087/egolm-protocol-v2-code/resolve/main/mGPT/callback.py
- Command line
-
hf download hf://po03087/egolm-protocol-v2-code/mGPT/callback.py
-
curl -L -o callback.py https://huggingface.co/po03087/egolm-protocol-v2-code/resolve/main/mGPT/callback.py
19.4 kB
| 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) | |