OpenMotionLab/MotionGPT
118
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 