import os
import sys
import signal
import time
import csv
import warnings
import random
import shutil
import subprocess
import platform
import glob as glob_mod
from pathlib import Path
import spaces # ZeroGPU
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import pprint
from loguru import logger
import smplx
import soundfile as sf
import librosa
from transformers import pipeline
# Add project root to sys.path so intra-repo imports work
BASE_DIR = Path(__file__).parent.resolve()
if str(BASE_DIR) not in sys.path:
sys.path.insert(0, str(BASE_DIR))
if platform.system() == "Linux":
os.environ['PYOPENGL_PLATFORM'] = 'egl'
from huggingface_hub import snapshot_download, hf_hub_download
# ---------------------------------------------------------------------------
# Paths and directories
# ---------------------------------------------------------------------------
CKPT_DIR = BASE_DIR / "ckpt"
MEAN_STD_DIR = BASE_DIR / "mean_std"
WEIGHTS_DIR = BASE_DIR / "weights"
SMPLX_DIR = BASE_DIR / "datasets" / "hub" / "smplx_models"
DATA_DIR = BASE_DIR / "datasets" / "BEAT_SMPL" / "beat_v2.0.0" / "beat_english_v2.0.0"
PRETRAINED_VQ_DIR = BASE_DIR / "datasets" / "hub" / "pretrained_vq"
OUTPUT_DIR = BASE_DIR / "outputs" / "audio2pose"
for d in [CKPT_DIR, MEAN_STD_DIR, WEIGHTS_DIR, SMPLX_DIR, DATA_DIR, PRETRAINED_VQ_DIR, OUTPUT_DIR]:
d.mkdir(parents=True, exist_ok=True)
# ---------------------------------------------------------------------------
# Download pretrained weights at startup
# ---------------------------------------------------------------------------
print("[GestureLSM] Downloading model weights from pliu23/GestureLSM...")
weights_cache = snapshot_download(
repo_id="pliu23/GestureLSM",
repo_type="model",
local_dir=str(BASE_DIR / "hf_weights_cache"),
allow_patterns=["*.pth", "*.bin"],
)
weights_cache = Path(weights_cache)
# Map weights to expected locations
weight_map = {
"new_540_shortcut.bin": CKPT_DIR / "new_540_shortcut.bin",
"net_300000_upper.pth": CKPT_DIR / "net_300000_upper.pth",
"net_300000_hands.pth": CKPT_DIR / "net_300000_hands.pth",
"net_300000_lower.pth": CKPT_DIR / "net_300000_lower.pth",
"net_300000_face.pth": CKPT_DIR / "net_300000_face.pth",
"AESKConv_240_100.bin": WEIGHTS_DIR / "AESKConv_240_100.bin",
}
for src_name, dst_path in weight_map.items():
src = weights_cache / src_name
if src.exists() and not dst_path.exists():
shutil.copy2(src, dst_path)
print(f" Copied {src_name} -> {dst_path}")
# Copy face VQ model
face_vq_src = weights_cache / "net_300000_face.pth"
face_vq_dst = PRETRAINED_VQ_DIR / "face_vertex_1layer_790.bin"
if face_vq_src.exists() and not face_vq_dst.exists():
shutil.copy2(face_vq_src, face_vq_dst)
print(f" Copied face VQ model -> {face_vq_dst}")
# Also check for AESKConv in the repo's weights dir
if not (WEIGHTS_DIR / "AESKConv_240_100.bin").exists():
# Try to find it in the repo
local_aesk = BASE_DIR / "weights" / "AESKConv_240_100.bin"
if local_aesk.exists():
shutil.copy2(local_aesk, WEIGHTS_DIR / "AESKConv_240_100.bin")
# ---------------------------------------------------------------------------
# Download SMPLX model
# ---------------------------------------------------------------------------
print("[GestureLSM] Setting up SMPLX model...")
smplx_model_path = SMPLX_DIR / "smplx" / "SMPLX_NEUTRAL_2020.npz"
if not smplx_model_path.exists():
smplx_model_path.parent.mkdir(parents=True, exist_ok=True)
try:
smplx_file = hf_hub_download(
repo_id="Tharun156/GestureLSM",
filename="datasets/hub/smplx_models/smplx/SMPLX_NEUTRAL_2020.npz",
repo_type="space",
)
shutil.copy2(smplx_file, smplx_model_path)
print(f" Copied SMPLX model -> {smplx_model_path}")
except Exception as e:
print(f" WARNING: Could not download SMPLX model from Tharun156: {e}")
# Try alternative source
try:
smplx_file = hf_hub_download(
repo_id="pliu23/GestureLSM",
filename="SMPLX_NEUTRAL_2020.npz",
repo_type="model",
)
shutil.copy2(smplx_file, smplx_model_path)
print(f" Copied SMPLX model from pliu23 -> {smplx_model_path}")
except Exception as e2:
print(f" WARNING: Could not download SMPLX model: {e2}")
# Create dummy train_test_split.csv (needed by CustomDataset)
csv_path = DATA_DIR / "train_test_split.csv"
if not csv_path.exists():
with open(csv_path, 'w', newline='') as f:
writer = csv.writer(f)
writer.writerow(['id', 'type'])
writer.writerow(['2_scott_0_1_1', 'test'])
# Create dummy data directories needed by the dataset loader
for subdir in ['smplxflame_30', 'textgrid', 'onset_amplitude', 'fasttext']:
(DATA_DIR / subdir).mkdir(parents=True, exist_ok=True)
print("[GestureLSM] Setup complete.")
# ---------------------------------------------------------------------------
# Import project modules
# ---------------------------------------------------------------------------
from utils import config as config_module, other_tools_hf, other_tools
from utils.joints import upper_body_mask, hands_body_mask, lower_body_mask
from dataloaders import data_tools
from dataloaders.build_vocab import Vocab
from dataloaders.data_tools import joints_list
from utils import rotation_conversions as rc
from models.vq.model import RVQVAE
from models.config import instantiate_from_config
device = "cuda" if torch.cuda.is_available() else "cpu"
# Load Whisper for ASR (replaces MFA)
print("[GestureLSM] Loading Whisper ASR model...")
whisper_pipe = pipeline(
"automatic-speech-recognition",
model="openai/whisper-tiny.en",
chunk_length_s=30,
device=device,
return_timestamps=True,
)
# ---------------------------------------------------------------------------
# Config loading (replaces config.parse_args)
# ---------------------------------------------------------------------------
def load_config():
"""Load config the same way demo.py does: configargparse + OmegaConf."""
cfg_path = str(BASE_DIR / "configs" / "shortcut_rvqvae_128_hf.yaml")
args, cfg = config_module.parse_args(cfg_path)
return args, cfg
# ---------------------------------------------------------------------------
# TextGrid creation from Whisper (replaces MFA)
# ---------------------------------------------------------------------------
def create_textgrid_from_whisper(audio_path, textgrid_path, audio_sr=16000):
"""Create a TextGrid file from Whisper word-level timestamps, replacing MFA."""
import textgrid as tg
result = whisper_pipe(audio_path, return_timestamps=True)
audio_data, sr = librosa.load(audio_path, sr=audio_sr)
audio_duration = len(audio_data) / sr
grid = tg.TextGrid()
word_tier = tg.IntervalTier(name="words", minTime=0)
grid.maxTime = audio_duration
word_tier.maxTime = audio_duration
if "chunks" in result:
for chunk in result["chunks"]:
text = chunk["text"].strip()
start_str, end_str = chunk["timestamp"]
start = float(start_str) if start_str is not None else 0.0
end = float(end_str) if end_str is not None else audio_duration
if text:
word_tier.add(minTime=start, maxTime=end, mark=text)
else:
word_tier.add(minTime=0, maxTime=audio_duration, mark=result.get("text", ""))
grid.append(word_tier)
grid.write(textgrid_path)
print(f"[GestureLSM] Created TextGrid: {textgrid_path}")
# ---------------------------------------------------------------------------
# GestureLSM Demo class (adapted from demo.py BaseTrainer)
# ---------------------------------------------------------------------------
class GestureLSMDemo:
def __init__(self, args, cfg):
self.args = args
self.cfg = cfg
self.rank = 0
self.ori_joint_list = joints_list[self.args.ori_joints]
self.tar_joint_list_face = joints_list["beat_smplx_face"]
self.tar_joint_list_upper = joints_list["beat_smplx_upper"]
self.tar_joint_list_hands = joints_list["beat_smplx_hands"]
self.tar_joint_list_lower = joints_list["beat_smplx_lower"]
self.joints = 55
self.joint_mask_face = np.zeros(len(list(self.ori_joint_list.keys())) * 3)
for joint_name in self.tar_joint_list_face:
self.joint_mask_face[self.ori_joint_list[joint_name][1] - self.ori_joint_list[joint_name][0]:self.ori_joint_list[joint_name][1]] = 1
self.joint_mask_upper = np.zeros(len(list(self.ori_joint_list.keys())) * 3)
for joint_name in self.tar_joint_list_upper:
self.joint_mask_upper[self.ori_joint_list[joint_name][1] - self.ori_joint_list[joint_name][0]:self.ori_joint_list[joint_name][1]] = 1
self.joint_mask_hands = np.zeros(len(list(self.ori_joint_list.keys())) * 3)
for joint_name in self.tar_joint_list_hands:
self.joint_mask_hands[self.ori_joint_list[joint_name][1] - self.ori_joint_list[joint_name][0]:self.ori_joint_list[joint_name][1]] = 1
self.joint_mask_lower = np.zeros(len(list(self.ori_joint_list.keys())) * 3)
for joint_name in self.tar_joint_list_lower:
self.joint_mask_lower[self.ori_joint_list[joint_name][1] - self.ori_joint_list[joint_name][0]:self.ori_joint_list[joint_name][1]] = 1
# Load SMPLX model
self.smplx = smplx.create(
self.args.data_path_1 + "smplx_models/",
model_type='smplx',
gender='NEUTRAL_2020',
use_face_contour=False,
num_betas=300,
num_expression_coeffs=100,
ext='npz',
use_pca=False,
).to(self.rank).eval()
# Load the main model
model_module = __import__(f"models.{cfg.model.model_name}", fromlist=["something"])
self.model = torch.nn.DataParallel(
getattr(model_module, cfg.model.g_name)(cfg), [0]
).cuda()
# Load VQ-VAE models
# Face VQ model: AESKConv_240_100.bin (not used in inference, just loaded for compatibility)
self.args.vae_layer = 2
self.args.vae_length = 240
self.args.vae_test_dim = 100
vq_model_module = __import__("models.motion_representation", fromlist=["something"])
self.vq_model_face = getattr(vq_model_module, "VQVAEConvZero")(self.args).to(self.rank)
try:
other_tools.load_checkpoints(self.vq_model_face, str(WEIGHTS_DIR / "AESKConv_240_100.bin"), self.args.e_name)
except Exception as e:
print(f"WARNING: Could not load face VQ model (not needed for inference): {e}")
self.vq_model_face.eval()
self.vq_model_upper = self._create_rvqvae_model(78, args.vqvae_upper_path)
self.vq_model_hands = self._create_rvqvae_model(180, args.vqvae_hands_path)
self.vq_model_lower = self._create_rvqvae_model(57, args.vqvae_lower_path)
self.vq_model_upper.eval().to(self.rank)
self.vq_model_hands.eval().to(self.rank)
self.vq_model_lower.eval().to(self.rank)
self.vqvae_latent_scale = self.args.vqvae_latent_scale
self.args.vae_length = 240
# Normalization
self.use_trans = self.args.use_trans
self.mean = np.load(args.mean_pose_path)
self.std = np.load(args.std_pose_path)
for part in ['upper', 'hands', 'lower']:
mask = globals()[f'{part}_body_mask']
setattr(self, f'mean_{part}', torch.from_numpy(self.mean[mask]).cuda())
setattr(self, f'std_{part}', torch.from_numpy(self.std[mask]).cuda())
if self.args.use_trans:
self.trans_mean = torch.from_numpy(np.load(self.args.mean_trans_path)).cuda()
self.trans_std = torch.from_numpy(np.load(self.args.std_trans_path)).cuda()
def _create_rvqvae_model(self, dim_pose, checkpoint_path):
args = self.args
model = RVQVAE(
args, dim_pose, args.nb_code, args.code_dim, args.code_dim,
args.down_t, args.stride_t, args.width, args.depth,
args.dilation_growth_rate, args.vq_act, args.vq_norm
)
model.load_state_dict(torch.load(checkpoint_path)['net'])
return model
def inverse_selection_tensor(self, filtered_t, selection_array, n):
selection_array = torch.from_numpy(selection_array).cuda()
original_shape_t = torch.zeros((n, 165)).cuda()
selected_indices = torch.where(selection_array == 1)[0]
for i in range(n):
original_shape_t[i, selected_indices] = filtered_t[i]
return original_shape_t
def _load_data(self, dict_data):
tar_pose_raw = dict_data["pose"]
tar_pose = tar_pose_raw[:, :, :165].to(self.rank)
tar_contact = tar_pose_raw[:, :, 165:169].to(self.rank)
tar_trans = dict_data["trans"].to(self.rank)
tar_trans_v = dict_data["trans_v"].to(self.rank)
tar_exps = dict_data["facial"].to(self.rank)
in_audio = dict_data["audio"].to(self.rank)
in_word = dict_data["word"].to(self.rank)
tar_beta = dict_data["beta"].to(self.rank)
tar_id = dict_data["id"].to(self.rank).long()
bs, n, j = tar_pose.shape[0], tar_pose.shape[1], self.joints
tar_pose_hands = tar_pose[:, :, 25*3:55*3]
tar_pose_hands = rc.axis_angle_to_matrix(tar_pose_hands.reshape(bs, n, 30, 3))
tar_pose_hands = rc.matrix_to_rotation_6d(tar_pose_hands).reshape(bs, n, 30*6)
tar_pose_upper = tar_pose[:, :, self.joint_mask_upper.astype(bool)]
tar_pose_upper = rc.axis_angle_to_matrix(tar_pose_upper.reshape(bs, n, 13, 3))
tar_pose_upper = rc.matrix_to_rotation_6d(tar_pose_upper).reshape(bs, n, 13*6)
tar_pose_leg = tar_pose[:, :, self.joint_mask_lower.astype(bool)]
tar_pose_leg = rc.axis_angle_to_matrix(tar_pose_leg.reshape(bs, n, 9, 3))
tar_pose_leg = rc.matrix_to_rotation_6d(tar_pose_leg).reshape(bs, n, 9*6)
tar_pose_lower = tar_pose_leg
if self.args.pose_norm:
tar_pose_upper = (tar_pose_upper - self.mean_upper) / self.std_upper
tar_pose_hands = (tar_pose_hands - self.mean_hands) / self.std_hands
tar_pose_lower = (tar_pose_lower - self.mean_lower) / self.std_lower
if self.use_trans:
tar_trans_v = (tar_trans_v - self.trans_mean) / self.trans_std
tar_pose_lower = torch.cat([tar_pose_lower, tar_trans_v], dim=-1)
latent_upper_top = self.vq_model_upper.map2latent(tar_pose_upper)
latent_hands_top = self.vq_model_hands.map2latent(tar_pose_hands)
latent_lower_top = self.vq_model_lower.map2latent(tar_pose_lower)
latent_in = torch.cat([latent_upper_top, latent_hands_top, latent_lower_top], dim=2) / self.args.vqvae_latent_scale
return {
"in_audio": in_audio,
"in_word": in_word,
"tar_trans": tar_trans,
"tar_exps": tar_exps,
"tar_beta": tar_beta,
"tar_pose": tar_pose,
"latent_in": latent_in,
"tar_id": tar_id,
"tar_contact": tar_contact,
"style_feature": None,
}
def _g_test(self, loaded_data):
bs, n, j = loaded_data["tar_pose"].shape[0], loaded_data["tar_pose"].shape[1], self.joints
tar_pose = loaded_data["tar_pose"]
tar_beta = loaded_data["tar_beta"]
tar_exps = loaded_data["tar_exps"]
tar_contact = loaded_data["tar_contact"]
tar_trans = loaded_data["tar_trans"]
in_word = loaded_data["in_word"]
in_audio = loaded_data["in_audio"]
in_x0 = loaded_data['latent_in']
in_seed = loaded_data['latent_in']
remain = n % 8
if remain != 0:
tar_pose = tar_pose[:, :-remain, :]
tar_beta = tar_beta[:, :-remain, :]
tar_trans = tar_trans[:, :-remain, :]
in_word = in_word[:, :-remain]
tar_exps = tar_exps[:, :-remain, :]
tar_contact = tar_contact[:, :-remain, :]
in_x0 = in_x0[:, :in_x0.shape[1] - (remain // self.args.vqvae_squeeze_scale), :]
in_seed = in_seed[:, :in_x0.shape[1] - (remain // self.args.vqvae_squeeze_scale), :]
n = n - remain
rec_all_upper = []
rec_all_lower = []
rec_all_hands = []
vqvae_squeeze_scale = self.args.vqvae_squeeze_scale
roundt = (n - self.args.pre_frames * vqvae_squeeze_scale) // (self.args.pose_length - self.args.pre_frames * vqvae_squeeze_scale)
remain = (n - self.args.pre_frames * vqvae_squeeze_scale) % (self.args.pose_length - self.args.pre_frames * vqvae_squeeze_scale)
round_l = self.args.pose_length - self.args.pre_frames * vqvae_squeeze_scale
for i in range(0, roundt):
in_word_tmp = in_word[:, i*(round_l):(i+1)*(round_l)+self.args.pre_frames * vqvae_squeeze_scale]
in_audio_tmp = in_audio[:, i*(16000//30*round_l):(i+1)*(16000//30*round_l)+16000//30*self.args.pre_frames * vqvae_squeeze_scale]
in_id_tmp = loaded_data['tar_id'][:, i*(round_l):(i+1)*(round_l)+self.args.pre_frames]
in_seed_tmp = in_seed[:, i*(round_l)//vqvae_squeeze_scale:(i+1)*(round_l)//vqvae_squeeze_scale+self.args.pre_frames]
in_x0_tmp = in_x0[:, i*(round_l)//vqvae_squeeze_scale:(i+1)*(round_l)//vqvae_squeeze_scale+self.args.pre_frames]
if i == 0:
in_seed_tmp = in_seed_tmp[:, :self.args.pre_frames, :]
else:
in_seed_tmp = last_sample[:, -self.args.pre_frames:, :]
cond_ = {'y': {}}
cond_['y']['audio_onset'] = in_audio_tmp
cond_['y']['word'] = in_word_tmp
cond_['y']['id'] = in_id_tmp
cond_['y']['seed'] = in_seed_tmp
cond_['y']['mask'] = (torch.zeros([self.args.batch_size, 1, 1, self.args.pose_length]) < 1).cuda()
cond_['y']['style_feature'] = torch.zeros([bs, 512]).cuda()
sample = self.model(cond_)['latents']
sample = sample.squeeze().permute(1, 0).unsqueeze(0)
last_sample = sample.clone()
rec_latent_upper = sample[..., :128]
rec_latent_hands = sample[..., 128:2*128]
rec_latent_lower = sample[..., 2*128:]
if i == 0:
rec_all_upper.append(rec_latent_upper)
rec_all_hands.append(rec_latent_hands)
rec_all_lower.append(rec_latent_lower)
else:
rec_all_upper.append(rec_latent_upper[:, self.args.pre_frames:])
rec_all_hands.append(rec_latent_hands[:, self.args.pre_frames:])
rec_all_lower.append(rec_latent_lower[:, self.args.pre_frames:])
rec_all_upper = torch.cat(rec_all_upper, dim=1) * self.vqvae_latent_scale
rec_all_hands = torch.cat(rec_all_hands, dim=1) * self.vqvae_latent_scale
rec_all_lower = torch.cat(rec_all_lower, dim=1) * self.vqvae_latent_scale
rec_upper = self.vq_model_upper.latent2origin(rec_all_upper)[0]
rec_hands = self.vq_model_hands.latent2origin(rec_all_hands)[0]
rec_lower = self.vq_model_lower.latent2origin(rec_all_lower)[0]
if self.use_trans:
rec_trans_v = rec_lower[..., -3:]
rec_trans_v = rec_trans_v * self.trans_std + self.trans_mean
rec_trans = torch.zeros_like(rec_trans_v)
rec_trans = torch.cumsum(rec_trans_v, dim=-2)
rec_trans[..., 1] = rec_trans_v[..., 1]
rec_lower = rec_lower[..., :-3]
if self.args.pose_norm:
rec_upper = rec_upper * self.std_upper + self.mean_upper
rec_hands = rec_hands * self.std_hands + self.mean_hands
rec_lower = rec_lower * self.std_lower + self.mean_lower
n = n - remain
tar_pose = tar_pose[:, :n, :]
tar_exps = tar_exps[:, :n, :]
tar_trans = tar_trans[:, :n, :]
tar_beta = tar_beta[:, :n, :]
rec_exps = tar_exps
rec_pose_legs = rec_lower[:, :, :54]
bs, n = rec_pose_legs.shape[0], rec_pose_legs.shape[1]
rec_pose_upper = rec_upper.reshape(bs, n, 13, 6)
rec_pose_upper = rc.rotation_6d_to_matrix(rec_pose_upper)
rec_pose_upper = rc.matrix_to_axis_angle(rec_pose_upper).reshape(bs*n, 13*3)
rec_pose_upper_recover = self.inverse_selection_tensor(rec_pose_upper, self.joint_mask_upper, bs*n)
rec_pose_lower = rec_pose_legs.reshape(bs, n, 9, 6)
rec_pose_lower = rc.rotation_6d_to_matrix(rec_pose_lower)
rec_pose_lower = rc.matrix_to_axis_angle(rec_pose_lower).reshape(bs*n, 9*3)
rec_pose_lower_recover = self.inverse_selection_tensor(rec_pose_lower, self.joint_mask_lower, bs*n)
rec_pose_hands = rec_hands.reshape(bs, n, 30, 6)
rec_pose_hands = rc.rotation_6d_to_matrix(rec_pose_hands)
rec_pose_hands = rc.matrix_to_axis_angle(rec_pose_hands).reshape(bs*n, 30*3)
rec_pose_hands_recover = self.inverse_selection_tensor(rec_pose_hands, self.joint_mask_hands, bs*n)
rec_pose = rec_pose_upper_recover + rec_pose_lower_recover + rec_pose_hands_recover
rec_pose[:, 66:69] = tar_pose.reshape(bs*n, 55*3)[:, 66:69]
rec_pose = rc.axis_angle_to_matrix(rec_pose.reshape(bs*n, j, 3))
rec_pose = rc.matrix_to_rotation_6d(rec_pose).reshape(bs, n, j*6)
tar_pose = rc.axis_angle_to_matrix(tar_pose.reshape(bs*n, j, 3))
tar_pose = rc.matrix_to_rotation_6d(tar_pose).reshape(bs, n, j*6)
return {
'rec_pose': rec_pose,
'rec_trans': rec_trans,
'tar_pose': tar_pose,
'tar_exps': tar_exps,
'tar_beta': tar_beta,
'tar_trans': tar_trans,
'rec_exps': rec_exps,
}
def test_demo(self, epoch):
results_save_path = self.checkpoint_path + f"/{epoch}/"
if os.path.exists(results_save_path):
shutil.rmtree(results_save_path)
os.makedirs(results_save_path)
start_time = time.time()
total_length = 0
self.model.eval()
self.smplx.eval()
with torch.no_grad():
for its, batch_data in enumerate(self.test_loader):
loaded_data = self._load_data(batch_data)
net_out = self._g_test(loaded_data)
tar_pose = net_out['tar_pose']
rec_pose = net_out['rec_pose']
tar_exps = net_out['tar_exps']
tar_beta = net_out['tar_beta']
rec_trans = net_out['rec_trans']
tar_trans = net_out['tar_trans']
rec_exps = net_out['rec_exps']
bs, n, j = tar_pose.shape[0], tar_pose.shape[1], self.joints
if (30 / self.args.pose_fps) != 1:
assert 30 % self.args.pose_fps == 0
n *= int(30 / self.args.pose_fps)
tar_pose = torch.nn.functional.interpolate(tar_pose.permute(0, 2, 1), scale_factor=30 / self.args.pose_fps, mode='linear').permute(0, 2, 1)
rec_pose = torch.nn.functional.interpolate(rec_pose.permute(0, 2, 1), scale_factor=30 / self.args.pose_fps, mode='linear').permute(0, 2, 1)
rec_pose = rc.rotation_6d_to_matrix(rec_pose.reshape(bs*n, j, 6))
rec_pose = rc.matrix_to_axis_angle(rec_pose).reshape(bs*n, j*3)
rec_pose_np = rec_pose.detach().cpu().numpy()
rec_trans_np = rec_trans.detach().cpu().numpy().reshape(bs*n, 3)
rec_exp_np = rec_exps.detach().cpu().numpy().reshape(bs*n, 100)
gt_npz = np.load(str(BASE_DIR / "demo" / "examples" / "2_scott_0_1_1.npz"), allow_pickle=True)
results_npz_file_save_path = results_save_path + f"result_{self.time_name_expend}" + '.npz'
np.savez(results_npz_file_save_path,
betas=gt_npz["betas"],
poses=rec_pose_np,
expressions=rec_exp_np,
trans=rec_trans_np,
model='smplx2020',
gender='neutral',
mocap_frame_rate=30,
)
total_length += n
render_vid_path = self._render_video(
results_npz_file_save_path,
results_save_path,
self.audio_path,
)
end_time = time.time() - start_time
logger.info(f"total inference time: {int(end_time)} s for {int(total_length/self.args.pose_fps)} s motion")
return render_vid_path, results_npz_file_save_path
def _render_video(self, res_npz_path, output_dir, audio_path):
"""Render the generated motion to a video with audio."""
import trimesh
import pyrender
import imageio
data_np_body = np.load(res_npz_path, allow_pickle=True)
if not os.path.exists(output_dir):
os.makedirs(output_dir)
faces = np.load(str(SMPLX_DIR / "smplx" / "SMPLX_NEUTRAL_2020.npz"), allow_pickle=True)["f"]
n = data_np_body["poses"].shape[0]
beta = torch.from_numpy(data_np_body["betas"]).to(torch.float32).unsqueeze(0).cuda()
beta = beta.repeat(n, 1)
expression = torch.from_numpy(data_np_body["expressions"][:n]).to(torch.float32).cuda()
jaw_pose = torch.from_numpy(data_np_body["poses"][:n, 66:69]).to(torch.float32).cuda()
pose = torch.from_numpy(data_np_body["poses"][:n]).to(torch.float32).cuda()
transl = torch.from_numpy(data_np_body["trans"][:n]).to(torch.float32).cuda()
with torch.no_grad():
output = self.smplx(
betas=beta, transl=transl, expression=expression, jaw_pose=jaw_pose,
global_orient=pose[:, :3], body_pose=pose[:, 3:21*3+3],
left_hand_pose=pose[:, 25*3:40*3], right_hand_pose=pose[:, 40*3:55*3],
leye_pose=pose[:, 69:72], reye_pose=pose[:, 72:75],
return_verts=True
)
vertices_all = output["vertices"].cpu().detach().numpy()
render_video_fps = 30
fig_resolution = (500, 500)
renderer = pyrender.OffscreenRenderer(*fig_resolution)
uniform_color = [220, 220, 220, 255]
angle_rad = np.deg2rad(-2)
pose_camera = np.array([
[1.0, 0.0, 0.0, 0.0],
[0.0, np.cos(angle_rad), -np.sin(angle_rad), 1.0],
[0.0, np.sin(angle_rad), np.cos(angle_rad), 5.0],
[0.0, 0.0, 0.0, 1.0]
])
angle_rad = np.deg2rad(-30)
pose_light = np.array([
[1.0, 0.0, 0.0, 0.0],
[0.0, np.cos(angle_rad), -np.sin(angle_rad), 0.0],
[0.0, np.sin(angle_rad), np.cos(angle_rad), 3.0],
[0.0, 0.0, 0.0, 1.0]
])
output_frames_dir = os.path.join(output_dir, "frames/")
os.makedirs(output_frames_dir, exist_ok=True)
num_frames = vertices_all.shape[0]
for i in range(num_frames):
if i % 100 == 0:
print(f"Rendering frame {i}/{num_frames}")
vertices = vertices_all[i]
trimesh_mesh = trimesh.Trimesh(
vertices=vertices, faces=faces,
vertex_colors=uniform_color
)
mesh = pyrender.Mesh.from_trimesh(trimesh_mesh, smooth=True)
scene = pyrender.Scene()
scene.add(mesh)
camera = pyrender.OrthographicCamera(xmag=1.0, ymag=1.0)
scene.add(camera, pose=pose_camera)
light = pyrender.DirectionalLight(color=[1.0, 1.0, 1.0], intensity=4.0)
scene.add(light, pose=pose_light)
fig, _ = renderer.render(scene)
imageio.imwrite(os.path.join(output_frames_dir, f"frame_{i:06d}.png"), fig)
renderer.delete()
# Create video from frames
silent_video = os.path.join(output_dir, "silence_video.mp4")
cmd = [
'ffmpeg', '-y', '-framerate', str(render_video_fps),
'-i', os.path.join(output_frames_dir, 'frame_%06d.png'),
'-c:v', 'libx264', '-pix_fmt', 'yuv420p',
silent_video
]
subprocess.run(cmd, check=True, capture_output=True)
# Clean up frames
for f in glob_mod.glob(os.path.join(output_frames_dir, "*.png")):
os.remove(f)
os.rmdir(output_frames_dir)
# Add audio to video
final_clip = os.path.join(output_dir, "result.mp4")
cmd = [
'ffmpeg', '-y',
'-i', silent_video, '-i', audio_path,
'-map', '0:v', '-map', '1:a',
'-c:v', 'copy', '-shortest',
final_clip
]
subprocess.run(cmd, check=True, capture_output=True)
os.remove(silent_video)
return final_clip
# ---------------------------------------------------------------------------
# Main inference function
# ---------------------------------------------------------------------------
@spaces.GPU(duration=180)
def gesturelsm(audio_path):
"""Main inference function for the Gradio demo."""
args, cfg = load_config()
if not sys.warnoptions:
warnings.simplefilter("ignore")
other_tools_hf.set_random_seed(args)
# Prepare audio and textgrid
tmp_dir = os.path.join(args.out_path, "custom", "hf_demo/")
os.makedirs(tmp_dir + "/", exist_ok=True)
time_local = time.localtime()
time_name_expend = "%02d%02d_%02d%02d%02d_" % (time_local[1], time_local[2], time_local[3], time_local[4], time_local[5])
# Copy uploaded audio
saved_audio_path = os.path.join(tmp_dir, "tmp.wav")
audio_data, sr = librosa.load(audio_path, sr=args.audio_sr)
sf.write(saved_audio_path, audio_data, args.audio_sr)
# Create TextGrid using Whisper (replaces MFA)
textgrid_path = os.path.join(tmp_dir, "tmp.TextGrid")
create_textgrid_from_whisper(saved_audio_path, textgrid_path, audio_sr=args.audio_sr)
args.textgrid_file_path = textgrid_path
args.audio_file_path = saved_audio_path
# Create trainer instance
trainer = GestureLSMDemo(args, cfg)
trainer.audio_path = saved_audio_path
trainer.checkpoint_path = tmp_dir
trainer.time_name_expend = time_name_expend
args.tmp_dir = tmp_dir
# Build test data
test_data = __import__(f"dataloaders.{args.dataset}", fromlist=["something"]).CustomDataset(args, "test")
trainer.test_loader = torch.utils.data.DataLoader(
test_data, batch_size=1, shuffle=False, num_workers=0, drop_last=False
)
# Load model checkpoint
other_tools.load_checkpoints(trainer.model, args.test_ckpt, args.g_name)
result = trainer.test_demo(999)
return result
import gradio as gr
examples = [
["demo/examples/2_scott_0_1_1.wav"],
["demo/examples/2_scott_0_2_2.wav"],
["demo/examples/2_scott_0_3_3.wav"],
["demo/examples/2_scott_0_4_4.wav"],
["demo/examples/2_scott_0_5_5.wav"],
]
CSS = """
.dark .gradio-container { color: var(--body-text-color); }
"""
demo = gr.Interface(css=CSS,
fn=gesturelsm,
inputs=[
gr.Audio(type="filepath", label="Upload Audio"),
],
outputs=[
gr.Video(format="mp4", visible=True, label="Generated Gesture Video"),
gr.File(label="Download motion (visualize in Blender)"),
],
title="GestureLSM: Latent Shortcut based Co-Speech Gesture Generation with Spatial-Temporal Modeling",
description="1. Upload your audio.
"
"2. Wait for the rendering to happen (1-4 minutes).
"
"3. View the generated gesture video.
"
"4. The face animation is fixed; only body motion is generated.
",
article="Project: [GestureLSM](https://github.com/andypinxinliu/GestureLSM) | "
"Paper: [arXiv:2501.18898](https://arxiv.org/abs/2501.18898)",
examples=examples,
theme=gr.themes.Citrus(),
)
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860)