Team Ai
Apppublic

leowajda/diffusion_model

sourceHugging Faceagpl-3.0updated 3y agoView on Hugging Face
1likes
diffusion_sampler.py145 linesDownload Raw Back to root
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