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 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 