Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
base.py205 linesDownload Raw Back to models
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