Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
HumanML3D.py118 linesDownload Raw Back to data
1import numpy as np2import torch3from os.path import join as pjoin4from .humanml.utils.word_vectorizer import WordVectorizer5from .humanml.scripts.motion_process import (process_file, recover_from_ric)6from . import BASEDataModule7from .humanml import Text2MotionDatasetEval, Text2MotionDataset, Text2MotionDatasetCB, MotionDataset, MotionDatasetVQ, Text2MotionDatasetToken, Text2MotionDatasetM2T8from .utils import humanml3d_collate9 10 11class HumanML3DDataModule(BASEDataModule):12    def __init__(self, cfg, **kwargs):13 14        super().__init__(collate_fn=humanml3d_collate)15        self.cfg = cfg16        self.save_hyperparameters(logger=False)17        18        # Basic info of the dataset19        cfg.DATASET.JOINT_TYPE = 'humanml3d'20        self.name = "humanml3d"21        self.njoints = 2222        23        # Path to the dataset24        data_root = cfg.DATASET.HUMANML3D.ROOT25        self.hparams.data_root = data_root26        self.hparams.text_dir = pjoin(data_root, "texts")27        self.hparams.motion_dir = pjoin(data_root, 'new_joint_vecs')28        29        # Mean and std of the dataset30        self.hparams.mean = np.load(pjoin('assets/meta', "mean.npy"))31        self.hparams.std = np.load(pjoin('assets/meta', "std.npy"))32        33        # Mean and std for fair evaluation34        self.hparams.mean_eval = np.load(pjoin('assets/meta', "mean_eval.npy"))35        self.hparams.std_eval = np.load(pjoin('assets/meta', "std_eval.npy"))36        37        # Length of the dataset38        self.hparams.max_motion_length = cfg.DATASET.HUMANML3D.MAX_MOTION_LEN39        self.hparams.min_motion_length = cfg.DATASET.HUMANML3D.MIN_MOTION_LEN40        self.hparams.max_text_len = cfg.DATASET.HUMANML3D.MAX_TEXT_LEN41        self.hparams.unit_length = cfg.DATASET.HUMANML3D.UNIT_LEN42 43        # Additional parameters44        self.hparams.debug = cfg.DEBUG45        self.hparams.stage = cfg.TRAIN.STAGE46 47        # Dataset switch48        self.DatasetEval = Text2MotionDatasetEval49 50        if cfg.TRAIN.STAGE == "vae":51            if cfg.model.params.motion_vae.target.split('.')[-1].lower() == "vqvae":52                self.hparams.win_size = 6453                self.Dataset = MotionDatasetVQ54            else:55                self.Dataset = MotionDataset56        elif 'lm' in cfg.TRAIN.STAGE:57            self.hparams.code_path = cfg.DATASET.CODE_PATH58            self.hparams.task_path = cfg.DATASET.TASK_PATH59            self.hparams.std_text = cfg.DATASET.HUMANML3D.STD_TEXT60            self.Dataset = Text2MotionDatasetCB61        elif cfg.TRAIN.STAGE == "token":62            self.Dataset = Text2MotionDatasetToken63            self.DatasetEval = Text2MotionDatasetToken64        elif cfg.TRAIN.STAGE == "m2t":65            self.Dataset = Text2MotionDatasetM2T66            self.DatasetEval = Text2MotionDatasetM2T67        else:68            self.Dataset = Text2MotionDataset69 70        # Get additional info of the dataset71        self.nfeats = 26372        cfg.DATASET.NFEATS = self.nfeats73        74 75    def feats2joints(self, features):76        mean = torch.tensor(self.hparams.mean).to(features)77        std = torch.tensor(self.hparams.std).to(features)78        features = features * std + mean79        return recover_from_ric(features, self.njoints)80 81    def joints2feats(self, features):82        features = process_file(features, self.njoints)[0]83        return features84 85    def normalize(self, features):86        mean = torch.tensor(self.hparams.mean).to(features)87        std = torch.tensor(self.hparams.std).to(features)88        features = (features - mean) / std89        return features90 91    def denormalize(self, features):92        mean = torch.tensor(self.hparams.mean).to(features)93        std = torch.tensor(self.hparams.std).to(features)94        features = features * std + mean95        return features96 97    def renorm4t2m(self, features):98        # renorm to t2m norms for using t2m evaluators99        ori_mean = torch.tensor(self.hparams.mean).to(features)100        ori_std = torch.tensor(self.hparams.std).to(features)101        eval_mean = torch.tensor(self.hparams.mean_eval).to(features)102        eval_std = torch.tensor(self.hparams.std_eval).to(features)103        features = features * ori_std + ori_mean104        features = (features - eval_mean) / eval_std105        return features106 107    def mm_mode(self, mm_on=True):108        if mm_on:109            self.is_mm = True110            self.name_list = self.test_dataset.name_list111            self.mm_list = np.random.choice(self.name_list,112                                            self.cfg.METRIC.MM_NUM_SAMPLES,113                                            replace=False)114            self.test_dataset.name_list = self.mm_list115        else:116            self.is_mm = False117            self.test_dataset.name_list = self.name_list118