Team Ai
Apppublic

OpenMotionLab/MotionGPT

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