File size: 2,857 Bytes
6dea0da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
import torch


class EDM():
    """
        Definition of most of the diffusion parameterization, following (Karras et al., "Elucidating...", 2022)
    """

    def __init__(self, args):
        self.args = args
        self.sigma_min = args.diff_params.sigma_min
        self.sigma_max = args.diff_params.sigma_max
        self.P_mean = args.diff_params.P_mean
        self.P_std = args.diff_params.P_std
        self.ro = args.diff_params.ro
        self.ro_train = args.diff_params.ro_train
        self.sigma_data = args.diff_params.sigma_data
        self.Schurn = args.diff_params.Schurn
        self.Stmin = args.diff_params.Stmin
        self.Stmax = args.diff_params.Stmax
        self.Snoise = args.diff_params.Snoise

    def get_gamma(self, t):
        """
        Get the parameter gamma that defines the stochasticity of the sampler
        Args
            t (Tensor): shape: (N_steps, ) Tensor of timesteps, from which we will compute gamma
        """
        N = t.shape[0]
        gamma = torch.zeros(t.shape).to(t.device)

        indexes = torch.logical_and(t > self.Stmin, t < self.Stmax)
        gamma[indexes] = gamma[indexes] + torch.min(torch.Tensor([self.Schurn / N, 2**(1 / 2) - 1]))

        return gamma

    def create_schedule(self, nb_steps):
        i = torch.arange(0, nb_steps + 1)
        t = (self.sigma_max**(1 / self.ro) + i / (nb_steps - 1) * (self.sigma_min**(1 / self.ro) - self.sigma_max**(1 / self.ro)))**self.ro
        t[-1] = 0
        return t

    def create_schedule_from_initial_t(self, initial_t, nb_steps):
        i = torch.arange(0, nb_steps + 1)
        t = (initial_t**(1 / self.ro) + i / (nb_steps - 1) * (self.sigma_min**(1 / self.ro) - initial_t**(1 / self.ro)))**self.ro
        t[-1] = 0
        return t

    def sample_prior(self, shape, sigma):
        n = torch.randn(shape).to(sigma.device) * sigma
        return n

    def cskip(self, sigma):
        return self.sigma_data**2 * (sigma**2 + self.sigma_data**2)**-1

    def cout(self, sigma):
        return sigma * self.sigma_data * (self.sigma_data**2 + sigma**2)**(-0.5)

    def cin(self, sigma):
        return (self.sigma_data**2 + sigma**2)**(-0.5)

    def cnoise(self, sigma):
        return (1 / 4) * torch.log(sigma)

    def denoiser(self, xn, net, sigma):
        """
        Applies the model and its EDM preconditioning to produce a denoised estimate.
        Args:
            xn (Tensor): shape (B,T) noisy latent to denoise
            net (nn.Module): the diffusion prior network
            sigma (Tensor): noise level (equal to timestep t)
        """
        if len(sigma.shape) == 1:
            sigma = sigma.unsqueeze(-1)
        cskip = self.cskip(sigma)
        cout = self.cout(sigma)
        cin = self.cin(sigma)
        cnoise = self.cnoise(sigma)

        return cskip * xn + cout * net(cin * xn, cnoise)