CoolFace
Apppublic

cnywt/SyncTalk

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
sphere_harmonics.py87 linesDownload Raw Back to shencoder
1import numpy as np2 3import torch4import torch.nn as nn5from torch.autograd import Function6from torch.autograd.function import once_differentiable7from torch.cuda.amp import custom_bwd, custom_fwd 8 9try:10    import _shencoder as _backend11except ImportError:12    from .backend import _backend13 14class _sh_encoder(Function):15    @staticmethod16    @custom_fwd(cast_inputs=torch.float32) # force float32 for better precision17    def forward(ctx, inputs, degree, calc_grad_inputs=False):18        # inputs: [B, input_dim], float in [-1, 1]19        # RETURN: [B, F], float20 21        inputs = inputs.contiguous()22        B, input_dim = inputs.shape # batch size, coord dim23        output_dim = degree ** 224        25        outputs = torch.empty(B, output_dim, dtype=inputs.dtype, device=inputs.device)26 27        if calc_grad_inputs:28            dy_dx = torch.empty(B, input_dim * output_dim, dtype=inputs.dtype, device=inputs.device)29        else:30            dy_dx = None31 32        _backend.sh_encode_forward(inputs, outputs, B, input_dim, degree, dy_dx)33 34        ctx.save_for_backward(inputs, dy_dx)35        ctx.dims = [B, input_dim, degree]36 37        return outputs38    39    @staticmethod40    #@once_differentiable41    @custom_bwd42    def backward(ctx, grad):43        # grad: [B, C * C]44 45        inputs, dy_dx = ctx.saved_tensors46 47        if dy_dx is not None:48            grad = grad.contiguous()49            B, input_dim, degree = ctx.dims50            grad_inputs = torch.zeros_like(inputs)51            _backend.sh_encode_backward(grad, inputs, B, input_dim, degree, dy_dx, grad_inputs)52            return grad_inputs, None, None53        else:54            return None, None, None55 56 57 58sh_encode = _sh_encoder.apply59 60 61class SHEncoder(nn.Module):62    def __init__(self, input_dim=3, degree=4):63        super().__init__()64 65        self.input_dim = input_dim # coord dims, must be 366        self.degree = degree # 0 ~ 467        self.output_dim = degree ** 268 69        assert self.input_dim == 3, "SH encoder only support input dim == 3"70        assert self.degree > 0 and self.degree <= 8, "SH encoder only supports degree in [1, 8]"71        72    def __repr__(self):73        return f"SHEncoder: input_dim={self.input_dim} degree={self.degree}"74    75    def forward(self, inputs, size=1):76        # inputs: [..., input_dim], normalized real world positions in [-size, size]77        # return: [..., degree^2]78 79        inputs = inputs / size # [-1, 1]80 81        prefix_shape = list(inputs.shape[:-1])82        inputs = inputs.reshape(-1, self.input_dim)83 84        outputs = sh_encode(inputs, self.degree, inputs.requires_grad)85        outputs = outputs.reshape(prefix_shape + [self.output_dim])86 87        return outputs