File size: 3,828 Bytes
4947683 | 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 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | import torch
import torch.nn.functional as F
import torch.nn as nn
import numpy as np
import math
def cosine_beta_schedule(timesteps, s=0.008):
"""
cosine schedule as proposed in https://arxiv.org/abs/2102.09672
"""
steps = timesteps + 1
x = torch.linspace(0, timesteps, steps)
alphas_cumprod = torch.cos(((x / timesteps) + s) / (1 + s) * math.pi * 0.5) ** 2
alphas_cumprod = alphas_cumprod / alphas_cumprod[0]
betas = 1 - (alphas_cumprod[1:] / alphas_cumprod[:-1])
return torch.clip(betas, 0.0001, 0.9999)
def linear_beta_schedule(timesteps, beta_start, beta_end):
return torch.linspace(beta_start, beta_end, timesteps)
def quadratic_beta_schedule(timesteps, beta_start, beta_end):
return torch.linspace(beta_start**0.5, beta_end**0.5, timesteps) ** 2
def sigmoid_beta_schedule(timesteps, beta_start, beta_end):
betas = torch.linspace(-6, 6, timesteps)
return torch.sigmoid(betas) * (beta_end - beta_start) + beta_start
def p_wrapped_normal(x, sigma, N=10, T=1.0):
p_ = 0
for i in range(-N, N + 1):
p_ += torch.exp(-(x + T * i) ** 2 / 2 / sigma ** 2)
return p_
def d_log_p_wrapped_normal(x, sigma, N=10, T=1.0):
p_ = 0
for i in range(-N, N + 1):
p_ += (x + T * i) / sigma ** 2 * torch.exp(-(x + T * i) ** 2 / 2 / sigma ** 2)
return p_ / p_wrapped_normal(x, sigma, N, T)
def sigma_norm(sigma, T=1.0, sn = 10000):
sigmas = sigma[None, :].repeat(sn, 1)
x_sample = sigma * torch.randn_like(sigmas)
x_sample = x_sample % T
normal_ = d_log_p_wrapped_normal(x_sample, sigmas, T = T)
return (normal_ ** 2).mean(dim = 0)
class BetaScheduler(nn.Module):
def __init__(
self,
timesteps,
scheduler_mode,
beta_start = 0.0001,
beta_end = 0.02
):
super(BetaScheduler, self).__init__()
self.timesteps = timesteps
if scheduler_mode == 'cosine':
betas = cosine_beta_schedule(timesteps)
elif scheduler_mode == 'linear':
betas = linear_beta_schedule(timesteps, beta_start, beta_end)
elif scheduler_mode == 'quadratic':
betas = quadratic_beta_schedule(timesteps, beta_start, beta_end)
elif scheduler_mode == 'sigmoid':
betas = sigmoid_beta_schedule(timesteps, beta_start, beta_end)
betas = torch.cat([torch.zeros([1]), betas], dim=0)
alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, axis=0)
sigmas = torch.zeros_like(betas)
sigmas[1:] = betas[1:] * (1. - alphas_cumprod[:-1]) / (1. - alphas_cumprod[1:])
sigmas = torch.sqrt(sigmas)
self.register_buffer('betas', betas)
self.register_buffer('alphas', alphas)
self.register_buffer('alphas_cumprod', alphas_cumprod)
self.register_buffer('sigmas', sigmas)
def uniform_sample_t(self, batch_size, device):
ts = np.random.choice(np.arange(1, self.timesteps+1), batch_size)
return torch.from_numpy(ts).to(device)
class SigmaScheduler(nn.Module):
def __init__(
self,
timesteps,
sigma_begin = 0.01,
sigma_end = 1.0
):
super(SigmaScheduler, self).__init__()
self.timesteps = timesteps
self.sigma_begin = sigma_begin
self.sigma_end = sigma_end
sigmas = torch.FloatTensor(np.exp(np.linspace(np.log(sigma_begin), np.log(sigma_end), timesteps)))
sigmas_norm_ = sigma_norm(sigmas)
self.register_buffer('sigmas', torch.cat([torch.zeros([1]), sigmas], dim=0))
self.register_buffer('sigmas_norm', torch.cat([torch.ones([1]), sigmas_norm_], dim=0))
def uniform_sample_t(self, batch_size, device):
ts = np.random.choice(np.arange(1, self.timesteps+1), batch_size)
return torch.from_numpy(ts).to(device)
|