po03087's picture
EgoLM baseline code (Ego3DLM snapshot, unmodified) + upload notes
3de4238 verified
Raw History Blame Contribute Delete
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)