DDPM-Fashion / app.py
Michael Yip
Update space
8129674
Raw History Blame Contribute Delete
2.95 kB
import gradio as gr
import numpy as np
import random
class DDPMCB(TrainCB):
def __init__(self , n_steps, beta_min, beta_max):
super().__init__()
self.n_steps, self.beta_min, self.beta_max = n_steps, beta_min, beta_max
#schedule the beta with linearly
self.beta = torch.linspace(self.beta_min, self.beta_max, self.n_steps)
self.alpha = 1. - self.beta
self.alpha_bar = torch.cumprod(self.alpha, dim=0) #cumulative product causes alpha_bar to decrease over time
self.sigma = self.beta.sqrt() #used in the reverse processs (see 3.2 of paper)
def predict(self, learn):
learn.preds = learn.model(*learn.batch[0]).sample #model.sample() will do the reverse diffusion and predict the noise
def before_batch(self, learn): #Add gaussian noise to batch before passing in to the Unet, modify learn.batch
device = learn.batch[0].device
epsilon = torch.randn(learn.batch[0].shape, device=device)
self.x0 = learn.batch[0]
self.alpha_bar = self.alpha_bar.to(device)
n = self.x0.shape[0] #number of images
#uniformly select n timesteps from the range of 0 to n_steps
t = torch.randint(0, self.n_steps, (n,), device=device, dtype=torch.long)
alpha_bart = self.alpha_bar[t].reshape(-1,1,1,1).to(device) #reshaping for the alphabart to fit across the n_inputs axis
xt = alpha_bart.sqrt()*self.x0 + (1-alpha_bart).sqrt()*epsilon #formula to get xt
#update batch to noised images
learn.batch = ((xt, t), epsilon)
#PLSSS learn this part idk how the sampling part works
@torch.no_grad()
def sample(self, model, sz):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
#generate gaussian noise Xt
x_t = torch.randn(sz, device=device)
preds = []
# iterate from T to 0
for t in reversed(range(self.n_steps)):
z = torch.randn(sz, device=device) if t > 1 else torch.zeros(sz, device=device)
alpha_t = self.alpha[t]
alpha_bart = self.alpha_bar[t]
prediction = model(x_t, t).sample
sigma_t = self.sigma[t]
x_t = (1/alpha_t.sqrt()) * (x_t - ((1-alpha_t) / (1-alpha_bart).sqrt()) * prediction) + (sigma_t * z)
preds.append(x_t.cpu())
return preds
# import spaces #[uncomment to use ZeroGPU]
from diffusers import DiffusionPipeline
import torch
device = "cuda" if torch.cuda.is_available() else "cpu"
#load model from pkl
ddpm_cb = DDPMCB(n_steps=1000, beta_min=0.0001, beta_max=0.02)
model = torch.load('fashion_ddpm.pkl')
# @spaces.GPU #[uncomment to use ZeroGPU]
def generate():
return to_pil_image(ddpm_cb.sample(learn.model, (1,1,32,32))[-1])
interface = gr.Interface(
fn = generate,
outputs=gr.Image(type="pil"),
title= "fashion mnist generation"
)
if __name__ == "__main__":
interface.launch()