leowajda/diffusion_model
1
1import tensorflow as tf2import math3from tensorflow import keras4from keras.models import load_model5 6 7def as_float32(t: tf.Tensor) -> tf.Tensor:8 return tf.cast(t, dtype=tf.float32)9 10 11def batch_reshape(t: tf.Tensor, x: tf.Tensor) -> tf.Tensor:12 def inner_function(coeff: tf.Tensor) -> tf.Tensor:13 batch_dim = tf.shape(x)[0]14 return tf.reshape(tf.gather(coeff, t), [batch_dim, 1, 1, 1])15 16 return inner_function17 18 19class DiffusionSampler:20 def __init__(21 self,22 model: keras.Model | str,23 ema_model: keras.Model | str,24 timesteps: int | None = 1_000,25 beta_start: float | None = 1e-4,26 beta_end: float | None = 0.02,27 noise_scheduler: str = "linear",28 ema: float = 0.999,29 ):30 self.noise_predictor = load_model(filepath=model, safe_mode=False) if isinstance(model, str) else model31 self.ema_noise_predictor = load_model(filepath=ema_model, safe_mode=False) if isinstance(model, str) else ema_model32 self.ema = ema33 self.beta_start = beta_start34 self.beta_end = beta_end35 self.timesteps = timesteps36 37 betas = self.noise_scheduler(noise_scheduler)38 alphas = 1.0 - betas39 alphas_cum_prod = tf.math.cumprod(alphas, axis=0)40 alphas_cum_prod_prev = tf.concat([tf.constant([1.0], dtype=tf.float64), alphas_cum_prod[:-1]], axis=0)41 posterior_variances = betas * (1.0 - alphas_cum_prod_prev) / (1.0 - alphas_cum_prod)42 43 self.betas = as_float32(betas)44 self.posterior_variances = as_float32(posterior_variances)45 self.alphas_cum_prod_prev = as_float32(alphas_cum_prod_prev)46 self.one_minus_alphas_cum_prod = as_float32(1.0 - alphas_cum_prod)47 self.one_minus_alphas_cum_prod_prev = as_float32(1.0 - alphas_cum_prod_prev)48 49 self.sqrt_one_minus_alphas_cum_prod = as_float32(tf.sqrt(1.0 - alphas_cum_prod))50 self.sqrt_alphas_cum_prod_prev = as_float32(tf.sqrt(alphas_cum_prod_prev))51 self.sqrt_alphas_cum_prod = as_float32(tf.sqrt(alphas_cum_prod))52 53 self.rev_sqrt_alphas_cum_prod = as_float32(1.0 / tf.sqrt(alphas_cum_prod))54 self.rev_sqrt_alphas = as_float32(tf.sqrt(1.0 / alphas))55 56 def ddpm_sample(self, pred_noise: tf.Tensor, x_t: tf.Tensor, t: tf.Tensor) -> tf.Tensor:57 batch_dim = tf.shape(x_t)[0]58 at_timestep = batch_reshape(t, x_t)59 60 beta = at_timestep(self.betas)61 rev_sqrt_alpha = at_timestep(self.rev_sqrt_alphas)62 sqrt_one_minus_alpha_cum_prod = at_timestep(self.sqrt_one_minus_alphas_cum_prod)63 posterior_variance = at_timestep(self.posterior_variances)64 65 mean = rev_sqrt_alpha * (66 x_t - (beta / sqrt_one_minus_alpha_cum_prod) * pred_noise67 )68 69 nonzero_mask = tf.reshape(70 1 - tf.cast(tf.equal(t, 0), dtype=tf.float32), [batch_dim, 1, 1, 1]71 )72 73 random_noise = tf.random.normal(shape=x_t.shape, dtype=x_t.dtype)74 return mean + nonzero_mask * tf.sqrt(posterior_variance) * random_noise75 76 def ddim_sample(self, pred_noise: tf.Tensor, x_t: tf.Tensor, t: tf.Tensor, eta: float = 0.0) -> tf.Tensor:77 at_timestep = batch_reshape(t, x_t)78 79 sqrt_alpha_cum_prod_prev = at_timestep(self.sqrt_alphas_cum_prod_prev)80 rev_sqrt_alpha_cum_prod = at_timestep(self.rev_sqrt_alphas_cum_prod)81 sqrt_one_minus_alpha_cum_prod = at_timestep(self.sqrt_one_minus_alphas_cum_prod)82 alpha_cum_prod_prev = at_timestep(self.alphas_cum_prod_prev)83 one_minus_alpha_cum_prod = at_timestep(self.one_minus_alphas_cum_prod)84 one_minus_alpha_cum_prod_prev = at_timestep(self.one_minus_alphas_cum_prod_prev)85 86 x0_t = (87 (x_t - (sqrt_one_minus_alpha_cum_prod * pred_noise)) * rev_sqrt_alpha_cum_prod88 )89 c1 = eta * tf.sqrt(90 (one_minus_alpha_cum_prod_prev / one_minus_alpha_cum_prod) * (91 one_minus_alpha_cum_prod / alpha_cum_prod_prev)92 )93 94 x_t_dir = tf.sqrt(one_minus_alpha_cum_prod_prev - tf.square(c1))95 random_noise = tf.random.normal(shape=x_t.shape, dtype=x_t.dtype)96 return sqrt_alpha_cum_prod_prev * x0_t + x_t_dir * pred_noise + c1 * random_noise97 98 def noise_scheduler(self, scheduler: str, max_beta: int = 0.02) -> tf.Tensor:99 pi, T = [tf.constant(num, dtype=tf.float64) for num in (math.pi, self.timesteps)]100 101 alpha_bar = lambda i: tf.math.cos((i + 0.008) / 1.008 * pi / 2) ** 2102 cosine_scheduler = lambda t: tf.minimum(1 - alpha_bar((t + 1) / T) / alpha_bar(t / T), max_beta)103 104 if scheduler == "linear":105 x = tf.linspace(start=self.beta_start, stop=self.beta_end, num=self.timesteps)106 return tf.cast(x, dtype=tf.float64)107 108 elif scheduler == "cosine":109 x = tf.vectorized_map(fn=cosine_scheduler, elems=tf.range(self.timesteps, dtype=tf.float64))110 return tf.cast(x, dtype=tf.float64)111 112 def x_t(self, x_start: tf.Tensor, t: tf.Tensor, noise: tf.Tensor) -> tf.Tensor:113 at_timestep = batch_reshape(t, x_start)114 115 sqrt_alpha_cum_prod = at_timestep(self.sqrt_alphas_cum_prod)116 sqrt_one_minus_alpha_cum_prod = at_timestep(self.sqrt_one_minus_alphas_cum_prod)117 return sqrt_alpha_cum_prod * x_start + sqrt_one_minus_alpha_cum_prod * noise118 119 @tf.function120 def generate_images(121 self,122 num_images: int,123 steps: int,124 sample_strategy: str = "ddim",125 step_strategy: str = "uniform",126 ema: bool = True,127 ):128 sampling_stategies = {129 ("ddpm", "linear"): (self.ddpm_sample, tf.range(self.timesteps, dtype=tf.float64)),130 ("ddpm", "quadratic"): (self.ddpm_sample, tf.range(self.timesteps, dtype=tf.float64)),131 ("ddim", "linear"): (self.ddim_sample, tf.range(steps, dtype=tf.float64)),132 ("ddim", "quadratic"): (self.ddim_sample, tf.cast(tf.linspace(start=0.0, stop=tf.sqrt(self.timesteps * 0.8), num=steps) ** 2, dtype=tf.float64))133 }134 135 noise_predictor = self.ema_noise_predictor if ema else self.noise_predictor136 sampler, seq = sampling_stategies[(sample_strategy, step_strategy)]137 samples = tf.random.normal(shape=(num_images, 64, 64, 3), dtype=tf.float32)138 139 for t in tf.reverse(seq, axis=[0]):140 tt = tf.cast(tf.fill(dims=(num_images,), value=t), dtype=tf.int64)141 pred_noise = noise_predictor([samples, tt], training=False)142 samples = sampler(pred_noise, samples, tt, )143 144 return tf.clip_by_value(samples * 127.5 + 127.5, 0.0, 255.0)145 