MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / forward

Method forward

detrsmpl/models/body_models/smpl.py:85–141  ·  view source on GitHub ↗

Forward function. Args: *args: extra arguments for SMPL return_verts: whether to return vertices return_full_pose: whether to return full pose parameters **kwargs: extra arguments for SMPL Returns: output: contains output

(self,
                *args,
                return_verts: bool = True,
                return_full_pose: bool = False,
                **kwargs)

Source from the content-addressed store, hash-verified

83 self.body_part_segmentation = body_segmentation('smpl')
84
85 def forward(self,
86 *args,
87 return_verts: bool = True,
88 return_full_pose: bool = False,
89 **kwargs) -> dict:
90 """Forward function.
91
92 Args:
93 *args: extra arguments for SMPL
94 return_verts: whether to return vertices
95 return_full_pose: whether to return full pose parameters
96 **kwargs: extra arguments for SMPL
97
98 Returns:
99 output: contains output parameters and attributes
100 """
101
102 kwargs['get_skin'] = True
103 smpl_output = super(SMPL, self).forward(*args, **kwargs)
104
105 if not hasattr(self, 'joints_regressor'):
106 joints = smpl_output.joints
107 else:
108 joints = vertices2joints(self.joints_regressor,
109 smpl_output.vertices)
110
111 if hasattr(self, 'joints_regressor_extra'):
112 extra_joints = vertices2joints(self.joints_regressor_extra,
113 smpl_output.vertices)
114 joints = torch.cat([joints, extra_joints], dim=1)
115
116 joints, joint_mask = convert_kps(joints,
117 src=self.keypoint_src,
118 dst=self.keypoint_dst,
119 approximate=self.keypoint_approximate)
120 if isinstance(joint_mask, np.ndarray):
121 joint_mask = torch.tensor(joint_mask,
122 dtype=torch.uint8,
123 device=joints.device)
124
125 batch_size = joints.shape[0]
126 joint_mask = joint_mask.reshape(1, -1).expand(batch_size, -1)
127
128 output = dict(global_orient=smpl_output.global_orient,
129 body_pose=smpl_output.body_pose,
130 joints=joints,
131 joint_mask=joint_mask,
132 keypoints=torch.cat([joints, joint_mask[:, :, None]],
133 dim=-1),
134 betas=smpl_output.betas)
135
136 if return_verts:
137 output['vertices'] = smpl_output.vertices
138 if return_full_pose:
139 output['full_pose'] = smpl_output.full_pose
140
141 return output
142

Callers

nothing calls this directly

Calls 2

vertices2jointsFunction · 0.90
convert_kpsFunction · 0.90

Tested by

no test coverage detected