Spaces:
Runtime error
Runtime error
Download app.py from TInkybala/DDPM-Fashion: direct link, hf CLI and curl.
- Browser
- Download file 2.95 kB
-
https://huggingface.co/spaces/TInkybala/DDPM-Fashion/resolve/main/app.py
- Command line
-
hf download hf://spaces/TInkybala/DDPM-Fashion/app.py
-
curl -L -o app.py https://huggingface.co/spaces/TInkybala/DDPM-Fashion/resolve/main/app.py
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 | |
| 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() | |