Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
utils.py82 linesDownload Raw Back to data
1import torch2import rich3import pickle4import numpy as np5 6 7def lengths_to_mask(lengths):8    max_len = max(lengths)9    mask = torch.arange(max_len, device=lengths.device).expand(10        len(lengths), max_len) < lengths.unsqueeze(1)11    return mask12 13 14# padding to max length in one batch15def collate_tensors(batch):16    if isinstance(batch[0], np.ndarray):17        batch = [torch.tensor(b).float() for b in batch]18 19    dims = batch[0].dim()20    max_size = [max([b.size(i) for b in batch]) for i in range(dims)]21    size = (len(batch), ) + tuple(max_size)22    canvas = batch[0].new_zeros(size=size)23    for i, b in enumerate(batch):24        sub_tensor = canvas[i]25        for d in range(dims):26            sub_tensor = sub_tensor.narrow(d, 0, b.size(d))27        sub_tensor.add_(b)28    return canvas29 30def humanml3d_collate(batch):31    notnone_batches = [b for b in batch if b is not None]32    EvalFlag = False if notnone_batches[0][5] is None else True33 34    # Sort by text length35    if EvalFlag:36        notnone_batches.sort(key=lambda x: x[5], reverse=True)37 38    # Motion only39    adapted_batch = {40        "motion":41        collate_tensors([torch.tensor(b[1]).float() for b in notnone_batches]),42        "length": [b[2] for b in notnone_batches],43    }44 45    # Text and motion46    if notnone_batches[0][0] is not None:47        adapted_batch.update({48            "text": [b[0] for b in notnone_batches],49            "all_captions": [b[7] for b in notnone_batches],50        })51 52    # Evaluation related53    if EvalFlag:54        adapted_batch.update({55            "text": [b[0] for b in notnone_batches],56            "word_embs":57            collate_tensors(58                [torch.tensor(b[3]).float() for b in notnone_batches]),59            "pos_ohot":60            collate_tensors(61                [torch.tensor(b[4]).float() for b in notnone_batches]),62            "text_len":63            collate_tensors([torch.tensor(b[5]) for b in notnone_batches]),64            "tokens": [b[6] for b in notnone_batches],65        })66 67    # Tasks68    if len(notnone_batches[0]) == 9:69        adapted_batch.update({"tasks": [b[8] for b in notnone_batches]})70 71    return adapted_batch72 73 74def load_pkl(path, description=None, progressBar=False):75    if progressBar:76        with rich.progress.open(path, 'rb', description=description) as file:77            data = pickle.load(file)78    else:79        with open(path, 'rb') as file:80            data = pickle.load(file)81    return data82