Shellbrady/LivePortrait5
0
1# coding: utf-82 3"""4functions for processing and transforming 3D facial keypoints5"""6 7import numpy as np8import torch9import torch.nn.functional as F10 11PI = np.pi12 13 14def headpose_pred_to_degree(pred):15 """16 pred: (bs, 66) or (bs, 1) or others17 """18 if pred.ndim > 1 and pred.shape[1] == 66:19 # NOTE: note that the average is modified to 97.520 device = pred.device21 idx_tensor = [idx for idx in range(0, 66)]22 idx_tensor = torch.FloatTensor(idx_tensor).to(device)23 pred = F.softmax(pred, dim=1)24 degree = torch.sum(pred*idx_tensor, axis=1) * 3 - 97.525 26 return degree27 28 return pred29 30 31def get_rotation_matrix(pitch_, yaw_, roll_):32 """ the input is in degree33 """34 # calculate the rotation matrix: vps @ rot35 36 # transform to radian37 pitch = pitch_ / 180 * PI38 yaw = yaw_ / 180 * PI39 roll = roll_ / 180 * PI40 41 device = pitch.device42 43 if pitch.ndim == 1:44 pitch = pitch.unsqueeze(1)45 if yaw.ndim == 1:46 yaw = yaw.unsqueeze(1)47 if roll.ndim == 1:48 roll = roll.unsqueeze(1)49 50 # calculate the euler matrix51 bs = pitch.shape[0]52 ones = torch.ones([bs, 1]).to(device)53 zeros = torch.zeros([bs, 1]).to(device)54 x, y, z = pitch, yaw, roll55 56 rot_x = torch.cat([57 ones, zeros, zeros,58 zeros, torch.cos(x), -torch.sin(x),59 zeros, torch.sin(x), torch.cos(x)60 ], dim=1).reshape([bs, 3, 3])61 62 rot_y = torch.cat([63 torch.cos(y), zeros, torch.sin(y),64 zeros, ones, zeros,65 -torch.sin(y), zeros, torch.cos(y)66 ], dim=1).reshape([bs, 3, 3])67 68 rot_z = torch.cat([69 torch.cos(z), -torch.sin(z), zeros,70 torch.sin(z), torch.cos(z), zeros,71 zeros, zeros, ones72 ], dim=1).reshape([bs, 3, 3])73 74 rot = rot_z @ rot_y @ rot_x75 return rot.permute(0, 2, 1) # transpose76 