From cd8a28816b4d274b90feddccea31a468385f2ddd Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Mar 2024 19:47:24 +0200 Subject: [PATCH] small refactor --- mogen/smpl/render_mesh.py | 9 +++++---- mogen/smpl/rotation2xyz.py | 13 +++++++------ 2 files changed, 12 insertions(+), 10 deletions(-) diff --git a/mogen/smpl/render_mesh.py b/mogen/smpl/render_mesh.py index bec8c44..7690f6b 100644 --- a/mogen/smpl/render_mesh.py +++ b/mogen/smpl/render_mesh.py @@ -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) diff --git a/mogen/smpl/rotation2xyz.py b/mogen/smpl/rotation2xyz.py index ed774d3..70d24f9 100644 --- a/mogen/smpl/rotation2xyz.py +++ b/mogen/smpl/rotation2xyz.py @@ -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]