small refactor

This commit is contained in:
kijai
2024-03-27 19:47:24 +02:00
parent ebb403c14e
commit cd8a28816b
2 changed files with 12 additions and 10 deletions
+5 -4
View File
@@ -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)
+7 -6
View File
@@ -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]