Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
smpl.py192 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 typing import Optional18from torch import Tensor19import smplx20 21from .base import Datastruct, dataclass, Transform22 23from .rots2rfeats import Rots2Rfeats24from .rots2joints import Rots2Joints25from .joints2jfeats import Joints2Jfeats26 27 28class SMPLTransform(Transform):29    def __init__(self, rots2rfeats: Rots2Rfeats,30                 rots2joints: Rots2Joints,31                 joints2jfeats: Joints2Jfeats,32                 **kwargs):33        self.rots2rfeats = rots2rfeats34        self.rots2joints = rots2joints35        self.joints2jfeats = joints2jfeats36 37    def Datastruct(self, **kwargs):38        return SMPLDatastruct(_rots2rfeats=self.rots2rfeats,39                              _rots2joints=self.rots2joints,40                              _joints2jfeats=self.joints2jfeats,41                              transforms=self,42                              **kwargs)43 44    def __repr__(self):45        return "SMPLTransform()"46 47 48class RotIdentityTransform(Transform):49    def __init__(self, **kwargs):50        return51 52    def Datastruct(self, **kwargs):53        return RotTransDatastruct(**kwargs)54 55    def __repr__(self):56        return "RotIdentityTransform()"57 58 59@dataclass60class RotTransDatastruct(Datastruct):61    rots: Tensor62    trans: Tensor63 64    transforms: RotIdentityTransform = RotIdentityTransform()65 66    def __post_init__(self):67        self.datakeys = ["rots", "trans"]68 69    def __len__(self):70        return len(self.rots)71 72 73@dataclass74class SMPLDatastruct(Datastruct):75    transforms: SMPLTransform76    _rots2rfeats: Rots2Rfeats77    _rots2joints: Rots2Joints78    _joints2jfeats: Joints2Jfeats79 80    features: Optional[Tensor] = None81    rots_: Optional[RotTransDatastruct] = None82    rfeats_: Optional[Tensor] = None83    joints_: Optional[Tensor] = None84    jfeats_: Optional[Tensor] = None85    vertices_: Optional[Tensor] = None86 87    def __post_init__(self):88        self.datakeys = ['features', 'rots_', 'rfeats_',89                         'joints_', 'jfeats_', 'vertices_']90        # starting point91        if self.features is not None and self.rfeats_ is None:92            self.rfeats_ = self.features93 94    @property95    def rots(self):96        # Cached value97        if self.rots_ is not None:98            return self.rots_99 100        # self.rfeats_ should be defined101        assert self.rfeats_ is not None102 103        self._rots2rfeats.to(self.rfeats.device)104        self.rots_ = self._rots2rfeats.inverse(self.rfeats)105        return self.rots_106 107    @property108    def rfeats(self):109        # Cached value110        if self.rfeats_ is not None:111            return self.rfeats_112 113        # self.rots_ should be defined114        assert self.rots_ is not None115 116        self._rots2rfeats.to(self.rots.device)117        self.rfeats_ = self._rots2rfeats(self.rots)118        return self.rfeats_119 120    @property121    def joints(self):122        # Cached value123        if self.joints_ is not None:124            return self.joints_125 126        self._rots2joints.to(self.rots.device)127        self.joints_ = self._rots2joints(self.rots)128        return self.joints_129 130    @property131    def jfeats(self):132        # Cached value133        if self.jfeats_ is not None:134            return self.jfeats_135 136        self._joints2jfeats.to(self.joints.device)137        self.jfeats_ = self._joints2jfeats(self.joints)138        return self.jfeats_139    140    @property141    def vertices(self):142        # Cached value143        if self.vertices_ is not None:144            return self.vertices_145 146        self._rots2joints.to(self.rots.device)147        self.vertices_ = self._rots2joints(self.rots, jointstype="vertices")148        return self.vertices_149    150    def __len__(self):151        return len(self.rfeats)152 153 154def get_body_model(model_type, gender, batch_size, device='cpu', ext='pkl'):155    '''156    type: smpl, smplx smplh and others. Refer to smplx tutorial157    gender: male, female, neutral158    batch_size: an positive integar159    '''160    mtype = model_type.upper()161    if gender != 'neutral':162        if not isinstance(gender, str):163            gender = str(gender.astype(str)).upper()164        else:165            gender = gender.upper()166    else:167        gender = gender.upper()168        ext = 'npz'169    body_model_path = f'data/smpl_models/{model_type}/{mtype}_{gender}.{ext}'170 171    body_model = smplx.create(body_model_path, model_type=type,172                              gender=gender, ext=ext,173                              use_pca=False,174                              num_pca_comps=12,175                              create_global_orient=True,176                              create_body_pose=True,177                              create_betas=True,178                              create_left_hand_pose=True,179                              create_right_hand_pose=True,180                              create_expression=True,181                              create_jaw_pose=True,182                              create_leye_pose=True,183                              create_reye_pose=True,184                              create_transl=True,185                              batch_size=batch_size)186    187    if device == 'cuda':188        return body_model.cuda()189    else:190        return body_model191 192