OpenMotionLab/MotionGPT
118
1# -*- coding: utf-8 -*-2 3# Max-Planck-Gesellschaft zur Förderung der Wissenschaften e.V. (MPG) is4# holder of all proprietary rights on this computer program.5# You can only use this computer program if you have closed6# a license agreement with MPG or you get the right to use the computer7# program from someone who is authorized to grant you that right.8# Any use of the computer program without a valid license is prohibited and9# liable to prosecution.10#11# Copyright©2020 Max-Planck-Gesellschaft zur Förderung12# der Wissenschaften e.V. (MPG). acting on behalf of its Max Planck Institute13# for Intelligent Systems. All rights reserved.14#15# Contact: ps-license@tuebingen.mpg.de16 17from dataclasses import dataclass, fields18 19 20class Transform:21 22 def collate(self, lst_datastruct):23 from ..tools import collate_tensor_with_padding24 example = lst_datastruct[0]25 26 def collate_or_none(key):27 if example[key] is None:28 return None29 key_lst = [x[key] for x in lst_datastruct]30 return collate_tensor_with_padding(key_lst)31 32 kwargs = {key: collate_or_none(key) for key in example.datakeys}33 34 return self.Datastruct(**kwargs)35 36 37# Inspired from SMPLX library38# need to define "datakeys" and transforms39@dataclass40class Datastruct:41 42 def __getitem__(self, key):43 return getattr(self, key)44 45 def __setitem__(self, key, value):46 self.__dict__[key] = value47 48 def get(self, key, default=None):49 return getattr(self, key, default)50 51 def __iter__(self):52 return self.keys()53 54 def keys(self):55 keys = [t.name for t in fields(self)]56 return iter(keys)57 58 def values(self):59 values = [getattr(self, t.name) for t in fields(self)]60 return iter(values)61 62 def items(self):63 data = [(t.name, getattr(self, t.name)) for t in fields(self)]64 return iter(data)65 66 def to(self, *args, **kwargs):67 for key in self.datakeys:68 if self[key] is not None:69 self[key] = self[key].to(*args, **kwargs)70 return self71 72 @property73 def device(self):74 return self[self.datakeys[0]].device75 76 def detach(self):77 78 def detach_or_none(tensor):79 if tensor is not None:80 return tensor.detach()81 return None82 83 kwargs = {key: detach_or_none(self[key]) for key in self.datakeys}84 return self.transforms.Datastruct(**kwargs)85 