small refactor
This commit is contained in:
@@ -162,19 +162,20 @@ def render(motions):
|
||||
|
||||
def render_from_smpl(thetas, yfov, move_x, move_y, move_z, x_rot, y_rot, z_rot, draw_platform=True, depth_only=False, normals=False, smpl_model_path=None, shape_parameters=None):
|
||||
if shape_parameters is not None:
|
||||
betas_single = torch.tensor([shape_parameters], dtype=torch.float32)
|
||||
betas_tensor = torch.tensor([shape_parameters], dtype=torch.float32)
|
||||
batch_size = thetas.shape[3]
|
||||
betas_batch = betas_single.repeat(batch_size, 1) # Replicates the single sample across the batch
|
||||
betas_batch = betas_tensor.repeat(batch_size, 1) # Replicates the single sample across the batch
|
||||
betas_batch = betas_batch.to(device=get_torch_device())
|
||||
else:
|
||||
betas_batch = None
|
||||
rot2xyz = Rotation2xyz(device=get_torch_device(), smpl_model_path=smpl_model_path)
|
||||
|
||||
rot2xyz = Rotation2xyz(device=get_torch_device(), smpl_model_path=smpl_model_path, betas=betas_batch)
|
||||
faces = rot2xyz.smpl_model.faces
|
||||
|
||||
vertices = rot2xyz(thetas.clone().to(get_torch_device()).detach(), mask=None,
|
||||
pose_rep='rot6d', translation=True, glob=True,
|
||||
jointstype='vertices',
|
||||
vertstrans=True, betas=betas_batch)
|
||||
vertstrans=True)
|
||||
|
||||
frames = vertices.shape[3] # shape: 1, nb_frames, 3, nb_joints
|
||||
print (vertices.shape)
|
||||
|
||||
@@ -8,14 +8,15 @@ JOINTSTYPES = ["a2m", "a2mpl", "smpl", "vibe", "vertices"]
|
||||
|
||||
|
||||
class Rotation2xyz:
|
||||
def __init__(self, device, dataset='amass', smpl_model_path=None):
|
||||
def __init__(self, device, dataset='amass', smpl_model_path=None, betas=None):
|
||||
self.device = device
|
||||
self.dataset = dataset
|
||||
self.smpl_model = SMPL(smpl_model_path).eval().to(device)
|
||||
self.betas = betas
|
||||
|
||||
def __call__(self, x, mask, pose_rep, translation, glob,
|
||||
jointstype, vertstrans, beta=0,
|
||||
glob_rot=None, get_rotations_back=False, betas=None, **kwargs):
|
||||
glob_rot=None, get_rotations_back=False, **kwargs):
|
||||
if pose_rep == "xyz":
|
||||
return x
|
||||
|
||||
@@ -57,12 +58,12 @@ class Rotation2xyz:
|
||||
global_orient = rotations[:, 0]
|
||||
rotations = rotations[:, 1:]
|
||||
|
||||
if betas is None:
|
||||
betas = torch.zeros([rotations.shape[0], self.smpl_model.num_betas],
|
||||
if self.betas is None:
|
||||
self.betas = torch.zeros([rotations.shape[0], self.smpl_model.num_betas],
|
||||
dtype=rotations.dtype, device=rotations.device)
|
||||
betas[:, 1] = beta
|
||||
self.betas[:, 1] = beta
|
||||
# import ipdb; ipdb.set_trace()
|
||||
out = self.smpl_model(body_pose=rotations, global_orient=global_orient, betas=betas)
|
||||
out = self.smpl_model(body_pose=rotations, global_orient=global_orient, betas=self.betas)
|
||||
|
||||
# get the desirable joints
|
||||
joints = out[jointstype]
|
||||
|
||||
Reference in New Issue
Block a user