multimodalart/EchoMimic-zero
8
1import torch2 3tensor_interpolation = None4 5 6def get_tensor_interpolation_method():7 return tensor_interpolation8 9 10def set_tensor_interpolation_method(is_slerp):11 global tensor_interpolation12 tensor_interpolation = slerp if is_slerp else linear13 14 15def linear(v1, v2, t):16 return (1.0 - t) * v1 + t * v217 18 19def slerp(20 v0: torch.Tensor, v1: torch.Tensor, t: float, DOT_THRESHOLD: float = 0.999521) -> torch.Tensor:22 u0 = v0 / v0.norm()23 u1 = v1 / v1.norm()24 dot = (u0 * u1).sum()25 if dot.abs() > DOT_THRESHOLD:26 # logger.info(f'warning: v0 and v1 close to parallel, using linear interpolation instead.')27 return (1.0 - t) * v0 + t * v128 omega = dot.acos()29 return (((1.0 - t) * omega).sin() * v0 + (t * omega).sin() * v1) / omega.sin()30 