OpenMotionLab/MotionGPT
118
1import os2import numpy as np3import torch4import logging5from pathlib import Path6from pytorch_lightning import LightningModule7from os.path import join as pjoin8from collections import OrderedDict9# from mGPT.metrics import BaseMetrics10from mGPT.config import get_obj_from_str11 12 13class BaseModel(LightningModule):14 def __init__(self, *args, **kwargs):15 super().__init__(*args, **kwargs)16 17 # self.configure_metrics()18 19 # Ablation20 self.test_step_outputs = []21 self.times = []22 self.rep_i = 023 24 def training_step(self, batch, batch_idx):25 return self.allsplit_step("train", batch, batch_idx)26 27 def validation_step(self, batch, batch_idx):28 return self.allsplit_step("val", batch, batch_idx)29 30 def test_step(self, batch, batch_idx):31 outputs = self.allsplit_step("test", batch, batch_idx)32 self.test_step_outputs.append(outputs)33 return outputs34 35 def predict_step(self, batch, batch_idx):36 return self.forward(batch)37 38 def on_train_epoch_end(self):39 # Log steps and losses40 dico = self.step_log_dict()41 # Log losses42 dico.update(self.loss_log_dict('train'))43 # Write to log only if not sanity check44 if not self.trainer.sanity_checking:45 self.log_dict(dico, sync_dist=True, rank_zero_only=True)46 47 def on_validation_epoch_end(self):48 # Log steps and losses49 dico = self.step_log_dict()50 # Log losses51 dico.update(self.loss_log_dict('train'))52 dico.update(self.loss_log_dict('val'))53 # Log metrics54 dico.update(self.metrics_log_dict())55 # Write to log only if not sanity check56 if not self.trainer.sanity_checking:57 self.log_dict(dico, sync_dist=True, rank_zero_only=True)58 59 def on_test_epoch_end(self):60 # Log metrics61 dico = self.metrics_log_dict()62 # Write to log only if not sanity check63 if not self.trainer.sanity_checking:64 self.log_dict(dico, sync_dist=True, rank_zero_only=True)65 self.save_npy(self.test_step_outputs)66 self.rep_i = self.rep_i + 167 # Free up the memory68 self.test_step_outputs.clear()69 70 def preprocess_state_dict(self, state_dict):71 new_state_dict = OrderedDict()72 73 # metric_state_dict = self.metrics.state_dict()74 loss_state_dict = self._losses.state_dict()75 76 # for k, v in metric_state_dict.items():77 # new_state_dict['metrics.' + k] = v78 79 for k, v in loss_state_dict.items():80 new_state_dict['_losses.' + k] = v81 82 for k, v in state_dict.items():83 if '_losses' not in k and 'Metrics' not in k:84 new_state_dict[k] = v85 86 return new_state_dict87 88 def load_state_dict(self, state_dict, strict=True):89 new_state_dict = self.preprocess_state_dict(state_dict)90 super().load_state_dict(new_state_dict, strict)91 92 def step_log_dict(self):93 return {94 "epoch": float(self.trainer.current_epoch),95 "step": float(self.trainer.current_epoch)96 }97 98 def loss_log_dict(self, split: str):99 losses = self._losses['losses_' + split]100 loss_dict = losses.compute(split)101 return loss_dict102 103 def metrics_log_dict(self):104 105 # For TM2TMetrics MM106 if self.trainer.datamodule.is_mm and "TM2TMetrics" in self.hparams.metrics_dict:107 metrics_dicts = ['MMMetrics']108 else:109 metrics_dicts = self.hparams.metrics_dict110 111 # Compute all metrics112 metrics_log_dict = {}113 for metric in metrics_dicts:114 metrics_dict = getattr(115 self.metrics,116 metric).compute(sanity_flag=self.trainer.sanity_checking)117 metrics_log_dict.update({118 f"Metrics/{metric}": value.item()119 for metric, value in metrics_dict.items()120 })121 122 return metrics_log_dict123 124 def configure_optimizers(self):125 # Optimizer126 optim_target = self.hparams.cfg.TRAIN.OPTIM.target127 if len(optim_target.split('.')) == 1:128 optim_target = 'torch.optim.' + optim_target129 optimizer = get_obj_from_str(optim_target)(130 params=self.parameters(), **self.hparams.cfg.TRAIN.OPTIM.params)131 132 # Scheduler133 scheduler_target = self.hparams.cfg.TRAIN.LR_SCHEDULER.target134 if len(scheduler_target.split('.')) == 1:135 scheduler_target = 'torch.optim.lr_scheduler.' + scheduler_target136 lr_scheduler = get_obj_from_str(scheduler_target)(137 optimizer=optimizer, **self.hparams.cfg.TRAIN.LR_SCHEDULER.params)138 139 return {'optimizer': optimizer, 'lr_scheduler': lr_scheduler}140 141 def configure_metrics(self):142 self.metrics = BaseMetrics(datamodule=self.datamodule, **self.hparams)143 144 def save_npy(self, outputs):145 cfg = self.hparams.cfg146 output_dir = Path(147 os.path.join(148 cfg.FOLDER,149 str(cfg.model.target.split('.')[-2].lower()),150 str(cfg.NAME),151 "samples_" + cfg.TIME,152 ))153 if cfg.TEST.SAVE_PREDICTIONS:154 lengths = [i[1] for i in outputs]155 outputs = [i[0] for i in outputs]156 157 if cfg.TEST.DATASETS[0].lower() in ["humanml3d", "kit"]:158 keyids = self.trainer.datamodule.test_dataset.name_list159 for i in range(len(outputs)):160 for bid in range(161 min(cfg.TEST.BATCH_SIZE, outputs[i].shape[0])):162 keyid = keyids[i * cfg.TEST.BATCH_SIZE + bid]163 data = self.trainer.datamodule.test_dataset.data_dict[164 keyid]165 166 motion = torch.tensor(data['motion'],167 device=outputs[i].device)168 motion = self.datamodule.normalize(motion)169 length = data['length']170 text_list = data['text']171 gen_joints = outputs[i][bid][:lengths[i][bid]].cpu(172 ).numpy()173 if cfg.TEST.REPLICATION_TIMES > 1:174 name = f"{keyid}.npy"175 else:176 name = f"{keyid}.npy"177 # save predictions results178 npypath = output_dir / name179 np.save(npypath, gen_joints)180 npypath = output_dir / f"{keyid}_gt.npy"181 joints = self.feats2joints(motion).cpu().numpy()182 np.save(npypath, joints)183 184 with open(output_dir / f"{keyid}.txt", "a") as f:185 for text in text_list:186 f.write(f"{text['caption']}\n")187 188 elif cfg.TEST.DATASETS[0].lower() in ["humanact12", "uestc"]:189 keyids = range(len(self.trainer.datamodule.test_dataset))190 for i in range(len(outputs)):191 for bid in range(192 min(cfg.TEST.BATCH_SIZE, outputs[i].shape[0])):193 keyid = keyids[i * cfg.TEST.BATCH_SIZE + bid]194 gen_joints = outputs[i][bid].cpu()195 gen_joints = gen_joints.permute(2, 0,196 1)[:lengths[i][bid],197 ...].numpy()198 if cfg.TEST.REPLICATION_TIMES > 1:199 name = f"{keyid}_{self.rep_i}"200 else:201 name = f"{keyid}.npy"202 # save predictions results203 npypath = output_dir / name204 np.save(npypath, gen_joints)205 