diffusers/community-pipelines-mirror
Community Pipeline Examples For more information about community pipelines, please have a look at this issue. Community pipeline examples consist pipelines that have been added by the community. Please have a look at the following tables to get an overview of all community examples. Click on the Code Example to get a copy-and-paste ready code example that you can try out. If a community pipeline doesn't work as expected, please open an issue and ping the author on it. Please… See the full description on the dataset page: https://huggingface.co/datasets/diffusers/community-pipelines-mirror.
922k
1from typing import Union2 3import torch4from PIL import Image5from torchvision import transforms as tfms6from tqdm.auto import tqdm7from transformers import CLIPTextModel, CLIPTokenizer8 9from diffusers import (10 AutoencoderKL,11 DDIMScheduler,12 DiffusionPipeline,13 LMSDiscreteScheduler,14 PNDMScheduler,15 UNet2DConditionModel,16)17 18 19class MagicMixPipeline(DiffusionPipeline):20 def __init__(21 self,22 vae: AutoencoderKL,23 text_encoder: CLIPTextModel,24 tokenizer: CLIPTokenizer,25 unet: UNet2DConditionModel,26 scheduler: Union[PNDMScheduler, LMSDiscreteScheduler, DDIMScheduler],27 ):28 super().__init__()29 30 self.register_modules(vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, unet=unet, scheduler=scheduler)31 32 # convert PIL image to latents33 def encode(self, img):34 with torch.no_grad():35 latent = self.vae.encode(tfms.ToTensor()(img).unsqueeze(0).to(self.device) * 2 - 1)36 latent = 0.18215 * latent.latent_dist.sample()37 return latent38 39 # convert latents to PIL image40 def decode(self, latent):41 latent = (1 / 0.18215) * latent42 with torch.no_grad():43 img = self.vae.decode(latent).sample44 img = (img / 2 + 0.5).clamp(0, 1)45 img = img.detach().cpu().permute(0, 2, 3, 1).numpy()46 img = (img * 255).round().astype("uint8")47 return Image.fromarray(img[0])48 49 # convert prompt into text embeddings, also unconditional embeddings50 def prep_text(self, prompt):51 text_input = self.tokenizer(52 prompt,53 padding="max_length",54 max_length=self.tokenizer.model_max_length,55 truncation=True,56 return_tensors="pt",57 )58 59 text_embedding = self.text_encoder(text_input.input_ids.to(self.device))[0]60 61 uncond_input = self.tokenizer(62 "",63 padding="max_length",64 max_length=self.tokenizer.model_max_length,65 truncation=True,66 return_tensors="pt",67 )68 69 uncond_embedding = self.text_encoder(uncond_input.input_ids.to(self.device))[0]70 71 return torch.cat([uncond_embedding, text_embedding])72 73 def __call__(74 self,75 img: Image.Image,76 prompt: str,77 kmin: float = 0.3,78 kmax: float = 0.6,79 mix_factor: float = 0.5,80 seed: int = 42,81 steps: int = 50,82 guidance_scale: float = 7.5,83 ) -> Image.Image:84 tmin = steps - int(kmin * steps)85 tmax = steps - int(kmax * steps)86 87 text_embeddings = self.prep_text(prompt)88 89 self.scheduler.set_timesteps(steps)90 91 width, height = img.size92 encoded = self.encode(img)93 94 torch.manual_seed(seed)95 noise = torch.randn(96 (1, self.unet.config.in_channels, height // 8, width // 8),97 ).to(self.device)98 99 latents = self.scheduler.add_noise(100 encoded,101 noise,102 timesteps=self.scheduler.timesteps[tmax],103 )104 105 input = torch.cat([latents] * 2)106 107 input = self.scheduler.scale_model_input(input, self.scheduler.timesteps[tmax])108 109 with torch.no_grad():110 pred = self.unet(111 input,112 self.scheduler.timesteps[tmax],113 encoder_hidden_states=text_embeddings,114 ).sample115 116 pred_uncond, pred_text = pred.chunk(2)117 pred = pred_uncond + guidance_scale * (pred_text - pred_uncond)118 119 latents = self.scheduler.step(pred, self.scheduler.timesteps[tmax], latents).prev_sample120 121 for i, t in enumerate(tqdm(self.scheduler.timesteps)):122 if i > tmax:123 if i < tmin: # layout generation phase124 orig_latents = self.scheduler.add_noise(125 encoded,126 noise,127 timesteps=t,128 )129 130 input = (131 (mix_factor * latents) + (1 - mix_factor) * orig_latents132 ) # interpolating between layout noise and conditionally generated noise to preserve layout sematics133 input = torch.cat([input] * 2)134 135 else: # content generation phase136 input = torch.cat([latents] * 2)137 138 input = self.scheduler.scale_model_input(input, t)139 140 with torch.no_grad():141 pred = self.unet(142 input,143 t,144 encoder_hidden_states=text_embeddings,145 ).sample146 147 pred_uncond, pred_text = pred.chunk(2)148 pred = pred_uncond + guidance_scale * (pred_text - pred_uncond)149 150 latents = self.scheduler.step(pred, t, latents).prev_sample151 152 return self.decode(latents)153 