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

Method forward

detrsmpl/models/body_models/flame.py:53–98  ·  view source on GitHub ↗

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

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

Source from the content-addressed store, hash-verified

51 self.num_joints = get_keypoint_num(convention=self.keypoint_dst)
52
53 def forward(self,
54 *args,
55 return_verts: bool = True,
56 return_full_pose: bool = False,
57 **kwargs) -> dict:
58 """Forward function.
59
60 Args:
61 *args: extra arguments for FLAME
62 return_verts: whether to return vertices
63 return_full_pose: whether to return full pose parameters
64 **kwargs: extra arguments for FLAME
65
66 Returns:
67 output: contains output parameters and attributes
68 """
69 flame_output = super(FLAME, self).forward(*args, **kwargs)
70 joints = flame_output.joints
71 joints, joint_mask = convert_kps(joints,
72 src=self.keypoint_src,
73 dst=self.keypoint_dst,
74 approximate=self.keypoint_approximate)
75 if isinstance(joint_mask, np.ndarray):
76 joint_mask = torch.tensor(joint_mask,
77 dtype=torch.uint8,
78 device=joints.device)
79
80 batch_size = joints.shape[0]
81 joint_mask = joint_mask.reshape(1, -1).expand(batch_size, -1)
82
83 output = dict(global_orient=flame_output.global_orient,
84 neck_pose=flame_output.neck_pose,
85 jaw_pose=flame_output.jaw_pose,
86 joints=joints,
87 joint_mask=joint_mask,
88 keypoints=torch.cat([joints, joint_mask[:, :, None]],
89 dim=-1),
90 betas=flame_output.betas,
91 expression=flame_output.expression)
92
93 if return_verts:
94 output['vertices'] = flame_output.vertices
95 if return_full_pose:
96 output['full_pose'] = flame_output.full_pose
97
98 return output
99
100
101class FLAMELayer(_FLAMELayer):

Callers 1

forwardMethod · 0.45

Calls 1

convert_kpsFunction · 0.90

Tested by

no test coverage detected