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
1# Copyright 2023 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15import inspect16from typing import Callable, List, Optional, Union17 18import torch19from packaging import version20from transformers import CLIPImageProcessor, CLIPTextModel, CLIPTokenizer21 22from diffusers import DiffusionPipeline23from diffusers.configuration_utils import FrozenDict24from diffusers.models import AutoencoderKL, UNet2DConditionModel25from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import StableDiffusionPipelineOutput26from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker27from diffusers.schedulers import (28 DDIMScheduler,29 DPMSolverMultistepScheduler,30 EulerAncestralDiscreteScheduler,31 EulerDiscreteScheduler,32 LMSDiscreteScheduler,33 PNDMScheduler,34)35from diffusers.utils import deprecate, is_accelerate_available, logging36 37 38logger = logging.get_logger(__name__) # pylint: disable=invalid-name39 40 41class ComposableStableDiffusionPipeline(DiffusionPipeline):42 r"""43 Pipeline for text-to-image generation using Stable Diffusion.44 45 This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods the46 library implements for all the pipelines (such as downloading or saving, running on a particular device, etc.)47 48 Args:49 vae ([`AutoencoderKL`]):50 Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations.51 text_encoder ([`CLIPTextModel`]):52 Frozen text-encoder. Stable Diffusion uses the text portion of53 [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically54 the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.55 tokenizer (`CLIPTokenizer`):56 Tokenizer of class57 [CLIPTokenizer](https://huggingface.co/docs/transformers/v4.21.0/en/model_doc/clip#transformers.CLIPTokenizer).58 unet ([`UNet2DConditionModel`]): Conditional U-Net architecture to denoise the encoded image latents.59 scheduler ([`SchedulerMixin`]):60 A scheduler to be used in combination with `unet` to denoise the encoded image latents. Can be one of61 [`DDIMScheduler`], [`LMSDiscreteScheduler`], or [`PNDMScheduler`].62 safety_checker ([`StableDiffusionSafetyChecker`]):63 Classification module that estimates whether generated images could be considered offensive or harmful.64 Please, refer to the [model card](https://huggingface.co/runwayml/stable-diffusion-v1-5) for details.65 feature_extractor ([`CLIPImageProcessor`]):66 Model that extracts features from generated images to be used as inputs for the `safety_checker`.67 """68 69 _optional_components = ["safety_checker", "feature_extractor"]70 71 def __init__(72 self,73 vae: AutoencoderKL,74 text_encoder: CLIPTextModel,75 tokenizer: CLIPTokenizer,76 unet: UNet2DConditionModel,77 scheduler: Union[78 DDIMScheduler,79 PNDMScheduler,80 LMSDiscreteScheduler,81 EulerDiscreteScheduler,82 EulerAncestralDiscreteScheduler,83 DPMSolverMultistepScheduler,84 ],85 safety_checker: StableDiffusionSafetyChecker,86 feature_extractor: CLIPImageProcessor,87 requires_safety_checker: bool = True,88 ):89 super().__init__()90 91 if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:92 deprecation_message = (93 f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"94 f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "95 "to update the config accordingly as leaving `steps_offset` might led to incorrect results"96 " in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"97 " it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"98 " file"99 )100 deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)101 new_config = dict(scheduler.config)102 new_config["steps_offset"] = 1103 scheduler._internal_dict = FrozenDict(new_config)104 105 if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:106 deprecation_message = (107 f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."108 " `clip_sample` should be set to False in the configuration file. Please make sure to update the"109 " config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"110 " future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"111 " nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"112 )113 deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)114 new_config = dict(scheduler.config)115 new_config["clip_sample"] = False116 scheduler._internal_dict = FrozenDict(new_config)117 118 if safety_checker is None and requires_safety_checker:119 logger.warning(120 f"You have disabled the safety checker for {self.__class__} by passing `safety_checker=None`. Ensure"121 " that you abide to the conditions of the Stable Diffusion license and do not expose unfiltered"122 " results in services or applications open to the public. Both the diffusers team and Hugging Face"123 " strongly recommend to keep the safety filter enabled in all public facing circumstances, disabling"124 " it only for use-cases that involve analyzing network behavior or auditing its results. For more"125 " information, please have a look at https://github.com/huggingface/diffusers/pull/254 ."126 )127 128 if safety_checker is not None and feature_extractor is None:129 raise ValueError(130 "Make sure to define a feature extractor when loading {self.__class__} if you want to use the safety"131 " checker. If you do not want to use the safety checker, you can pass `'safety_checker=None'` instead."132 )133 134 is_unet_version_less_0_9_0 = hasattr(unet.config, "_diffusers_version") and version.parse(135 version.parse(unet.config._diffusers_version).base_version136 ) < version.parse("0.9.0.dev0")137 is_unet_sample_size_less_64 = hasattr(unet.config, "sample_size") and unet.config.sample_size < 64138 if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:139 deprecation_message = (140 "The configuration file of the unet has set the default `sample_size` to smaller than"141 " 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"142 " following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"143 " CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"144 " \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"145 " configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"146 " in the config might lead to incorrect results in future versions. If you have downloaded this"147 " checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"148 " the `unet/config.json` file"149 )150 deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)151 new_config = dict(unet.config)152 new_config["sample_size"] = 64153 unet._internal_dict = FrozenDict(new_config)154 155 self.register_modules(156 vae=vae,157 text_encoder=text_encoder,158 tokenizer=tokenizer,159 unet=unet,160 scheduler=scheduler,161 safety_checker=safety_checker,162 feature_extractor=feature_extractor,163 )164 self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)165 self.register_to_config(requires_safety_checker=requires_safety_checker)166 167 def enable_vae_slicing(self):168 r"""169 Enable sliced VAE decoding.170 171 When this option is enabled, the VAE will split the input tensor in slices to compute decoding in several172 steps. This is useful to save some memory and allow larger batch sizes.173 """174 self.vae.enable_slicing()175 176 def disable_vae_slicing(self):177 r"""178 Disable sliced VAE decoding. If `enable_vae_slicing` was previously invoked, this method will go back to179 computing decoding in one step.180 """181 self.vae.disable_slicing()182 183 def enable_sequential_cpu_offload(self, gpu_id=0):184 r"""185 Offloads all models to CPU using accelerate, significantly reducing memory usage. When called, unet,186 text_encoder, vae and safety checker have their state dicts saved to CPU and then are moved to a187 `torch.device('meta') and loaded to GPU only when their specific submodule has its `forward` method called.188 """189 if is_accelerate_available():190 from accelerate import cpu_offload191 else:192 raise ImportError("Please install accelerate via `pip install accelerate`")193 194 device = torch.device(f"cuda:{gpu_id}")195 196 for cpu_offloaded_model in [self.unet, self.text_encoder, self.vae]:197 if cpu_offloaded_model is not None:198 cpu_offload(cpu_offloaded_model, device)199 200 if self.safety_checker is not None:201 # TODO(Patrick) - there is currently a bug with cpu offload of nn.Parameter in accelerate202 # fix by only offloading self.safety_checker for now203 cpu_offload(self.safety_checker.vision_model, device)204 205 @property206 def _execution_device(self):207 r"""208 Returns the device on which the pipeline's models will be executed. After calling209 `pipeline.enable_sequential_cpu_offload()` the execution device can only be inferred from Accelerate's module210 hooks.211 """212 if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):213 return self.device214 for module in self.unet.modules():215 if (216 hasattr(module, "_hf_hook")217 and hasattr(module._hf_hook, "execution_device")218 and module._hf_hook.execution_device is not None219 ):220 return torch.device(module._hf_hook.execution_device)221 return self.device222 223 def _encode_prompt(self, prompt, device, num_images_per_prompt, do_classifier_free_guidance, negative_prompt):224 r"""225 Encodes the prompt into text encoder hidden states.226 227 Args:228 prompt (`str` or `list(int)`):229 prompt to be encoded230 device: (`torch.device`):231 torch device232 num_images_per_prompt (`int`):233 number of images that should be generated per prompt234 do_classifier_free_guidance (`bool`):235 whether to use classifier free guidance or not236 negative_prompt (`str` or `List[str]`):237 The prompt or prompts not to guide the image generation. Ignored when not using guidance (i.e., ignored238 if `guidance_scale` is less than `1`).239 """240 batch_size = len(prompt) if isinstance(prompt, list) else 1241 242 text_inputs = self.tokenizer(243 prompt,244 padding="max_length",245 max_length=self.tokenizer.model_max_length,246 truncation=True,247 return_tensors="pt",248 )249 text_input_ids = text_inputs.input_ids250 untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids251 252 if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):253 removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1])254 logger.warning(255 "The following part of your input was truncated because CLIP can only handle sequences up to"256 f" {self.tokenizer.model_max_length} tokens: {removed_text}"257 )258 259 if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:260 attention_mask = text_inputs.attention_mask.to(device)261 else:262 attention_mask = None263 264 text_embeddings = self.text_encoder(265 text_input_ids.to(device),266 attention_mask=attention_mask,267 )268 text_embeddings = text_embeddings[0]269 270 # duplicate text embeddings for each generation per prompt, using mps friendly method271 bs_embed, seq_len, _ = text_embeddings.shape272 text_embeddings = text_embeddings.repeat(1, num_images_per_prompt, 1)273 text_embeddings = text_embeddings.view(bs_embed * num_images_per_prompt, seq_len, -1)274 275 # get unconditional embeddings for classifier free guidance276 if do_classifier_free_guidance:277 uncond_tokens: List[str]278 if negative_prompt is None:279 uncond_tokens = [""] * batch_size280 elif type(prompt) is not type(negative_prompt):281 raise TypeError(282 f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="283 f" {type(prompt)}."284 )285 elif isinstance(negative_prompt, str):286 uncond_tokens = [negative_prompt]287 elif batch_size != len(negative_prompt):288 raise ValueError(289 f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"290 f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"291 " the batch size of `prompt`."292 )293 else:294 uncond_tokens = negative_prompt295 296 max_length = text_input_ids.shape[-1]297 uncond_input = self.tokenizer(298 uncond_tokens,299 padding="max_length",300 max_length=max_length,301 truncation=True,302 return_tensors="pt",303 )304 305 if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:306 attention_mask = uncond_input.attention_mask.to(device)307 else:308 attention_mask = None309 310 uncond_embeddings = self.text_encoder(311 uncond_input.input_ids.to(device),312 attention_mask=attention_mask,313 )314 uncond_embeddings = uncond_embeddings[0]315 316 # duplicate unconditional embeddings for each generation per prompt, using mps friendly method317 seq_len = uncond_embeddings.shape[1]318 uncond_embeddings = uncond_embeddings.repeat(1, num_images_per_prompt, 1)319 uncond_embeddings = uncond_embeddings.view(batch_size * num_images_per_prompt, seq_len, -1)320 321 # For classifier free guidance, we need to do two forward passes.322 # Here we concatenate the unconditional and text embeddings into a single batch323 # to avoid doing two forward passes324 text_embeddings = torch.cat([uncond_embeddings, text_embeddings])325 326 return text_embeddings327 328 def run_safety_checker(self, image, device, dtype):329 if self.safety_checker is not None:330 safety_checker_input = self.feature_extractor(self.numpy_to_pil(image), return_tensors="pt").to(device)331 image, has_nsfw_concept = self.safety_checker(332 images=image, clip_input=safety_checker_input.pixel_values.to(dtype)333 )334 else:335 has_nsfw_concept = None336 return image, has_nsfw_concept337 338 def decode_latents(self, latents):339 latents = 1 / 0.18215 * latents340 image = self.vae.decode(latents).sample341 image = (image / 2 + 0.5).clamp(0, 1)342 # we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16343 image = image.cpu().permute(0, 2, 3, 1).float().numpy()344 return image345 346 def prepare_extra_step_kwargs(self, generator, eta):347 # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature348 # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.349 # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502350 # and should be between [0, 1]351 352 accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())353 extra_step_kwargs = {}354 if accepts_eta:355 extra_step_kwargs["eta"] = eta356 357 # check if the scheduler accepts generator358 accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())359 if accepts_generator:360 extra_step_kwargs["generator"] = generator361 return extra_step_kwargs362 363 def check_inputs(self, prompt, height, width, callback_steps):364 if not isinstance(prompt, str) and not isinstance(prompt, list):365 raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")366 367 if height % 8 != 0 or width % 8 != 0:368 raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")369 370 if (callback_steps is None) or (371 callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)372 ):373 raise ValueError(374 f"`callback_steps` has to be a positive integer but is {callback_steps} of type"375 f" {type(callback_steps)}."376 )377 378 def prepare_latents(self, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None):379 shape = (batch_size, num_channels_latents, height // self.vae_scale_factor, width // self.vae_scale_factor)380 if latents is None:381 if device.type == "mps":382 # randn does not work reproducibly on mps383 latents = torch.randn(shape, generator=generator, device="cpu", dtype=dtype).to(device)384 else:385 latents = torch.randn(shape, generator=generator, device=device, dtype=dtype)386 else:387 if latents.shape != shape:388 raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")389 latents = latents.to(device)390 391 # scale the initial noise by the standard deviation required by the scheduler392 latents = latents * self.scheduler.init_noise_sigma393 return latents394 395 @torch.no_grad()396 def __call__(397 self,398 prompt: Union[str, List[str]],399 height: Optional[int] = None,400 width: Optional[int] = None,401 num_inference_steps: int = 50,402 guidance_scale: float = 7.5,403 negative_prompt: Optional[Union[str, List[str]]] = None,404 num_images_per_prompt: Optional[int] = 1,405 eta: float = 0.0,406 generator: Optional[torch.Generator] = None,407 latents: Optional[torch.FloatTensor] = None,408 output_type: Optional[str] = "pil",409 return_dict: bool = True,410 callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,411 callback_steps: int = 1,412 weights: Optional[str] = "",413 ):414 r"""415 Function invoked when calling the pipeline for generation.416 417 Args:418 prompt (`str` or `List[str]`):419 The prompt or prompts to guide the image generation.420 height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):421 The height in pixels of the generated image.422 width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor):423 The width in pixels of the generated image.424 num_inference_steps (`int`, *optional*, defaults to 50):425 The number of denoising steps. More denoising steps usually lead to a higher quality image at the426 expense of slower inference.427 guidance_scale (`float`, *optional*, defaults to 5.0):428 Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).429 `guidance_scale` is defined as `w` of equation 2. of [Imagen430 Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >431 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,432 usually at the expense of lower image quality.433 negative_prompt (`str` or `List[str]`, *optional*):434 The prompt or prompts not to guide the image generation. Ignored when not using guidance (i.e., ignored435 if `guidance_scale` is less than `1`).436 num_images_per_prompt (`int`, *optional*, defaults to 1):437 The number of images to generate per prompt.438 eta (`float`, *optional*, defaults to 0.0):439 Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to440 [`schedulers.DDIMScheduler`], will be ignored for others.441 generator (`torch.Generator`, *optional*):442 A [torch generator](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make generation443 deterministic.444 latents (`torch.FloatTensor`, *optional*):445 Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image446 generation. Can be used to tweak the same generation with different prompts. If not provided, a latents447 tensor will ge generated by sampling using the supplied random `generator`.448 output_type (`str`, *optional*, defaults to `"pil"`):449 The output format of the generate image. Choose between450 [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`.451 return_dict (`bool`, *optional*, defaults to `True`):452 Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a453 plain tuple.454 callback (`Callable`, *optional*):455 A function that will be called every `callback_steps` steps during inference. The function will be456 called with the following arguments: `callback(step: int, timestep: int, latents: torch.FloatTensor)`.457 callback_steps (`int`, *optional*, defaults to 1):458 The frequency at which the `callback` function will be called. If not specified, the callback will be459 called at every step.460 461 Returns:462 [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`:463 [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] if `return_dict` is True, otherwise a `tuple.464 When returning a tuple, the first element is a list with the generated images, and the second element is a465 list of `bool`s denoting whether the corresponding generated image likely represents "not-safe-for-work"466 (nsfw) content, according to the `safety_checker`.467 """468 # 0. Default height and width to unet469 height = height or self.unet.config.sample_size * self.vae_scale_factor470 width = width or self.unet.config.sample_size * self.vae_scale_factor471 472 # 1. Check inputs. Raise error if not correct473 self.check_inputs(prompt, height, width, callback_steps)474 475 # 2. Define call parameters476 batch_size = 1 if isinstance(prompt, str) else len(prompt)477 device = self._execution_device478 # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)479 # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`480 # corresponds to doing no classifier free guidance.481 do_classifier_free_guidance = guidance_scale > 1.0482 483 if "|" in prompt:484 prompt = [x.strip() for x in prompt.split("|")]485 print(f"composing {prompt}...")486 487 if not weights:488 # specify weights for prompts (excluding the unconditional score)489 print("using equal positive weights (conjunction) for all prompts...")490 weights = torch.tensor([guidance_scale] * len(prompt), device=self.device).reshape(-1, 1, 1, 1)491 else:492 # set prompt weight for each493 num_prompts = len(prompt) if isinstance(prompt, list) else 1494 weights = [float(w.strip()) for w in weights.split("|")]495 # guidance scale as the default496 if len(weights) < num_prompts:497 weights.append(guidance_scale)498 else:499 weights = weights[:num_prompts]500 assert len(weights) == len(prompt), "weights specified are not equal to the number of prompts"501 weights = torch.tensor(weights, device=self.device).reshape(-1, 1, 1, 1)502 else:503 weights = guidance_scale504 505 # 3. Encode input prompt506 text_embeddings = self._encode_prompt(507 prompt, device, num_images_per_prompt, do_classifier_free_guidance, negative_prompt508 )509 510 # 4. Prepare timesteps511 self.scheduler.set_timesteps(num_inference_steps, device=device)512 timesteps = self.scheduler.timesteps513 514 # 5. Prepare latent variables515 num_channels_latents = self.unet.config.in_channels516 latents = self.prepare_latents(517 batch_size * num_images_per_prompt,518 num_channels_latents,519 height,520 width,521 text_embeddings.dtype,522 device,523 generator,524 latents,525 )526 527 # composable diffusion528 if isinstance(prompt, list) and batch_size == 1:529 # remove extra unconditional embedding530 # N = one unconditional embed + conditional embeds531 text_embeddings = text_embeddings[len(prompt) - 1 :]532 533 # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline534 extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)535 536 # 7. Denoising loop537 num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order538 with self.progress_bar(total=num_inference_steps) as progress_bar:539 for i, t in enumerate(timesteps):540 # expand the latents if we are doing classifier free guidance541 latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents542 latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)543 544 # predict the noise residual545 noise_pred = []546 for j in range(text_embeddings.shape[0]):547 noise_pred.append(548 self.unet(latent_model_input[:1], t, encoder_hidden_states=text_embeddings[j : j + 1]).sample549 )550 noise_pred = torch.cat(noise_pred, dim=0)551 552 # perform guidance553 if do_classifier_free_guidance:554 noise_pred_uncond, noise_pred_text = noise_pred[:1], noise_pred[1:]555 noise_pred = noise_pred_uncond + (weights * (noise_pred_text - noise_pred_uncond)).sum(556 dim=0, keepdims=True557 )558 559 # compute the previous noisy sample x_t -> x_t-1560 latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample561 562 # call the callback, if provided563 if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):564 progress_bar.update()565 if callback is not None and i % callback_steps == 0:566 step_idx = i // getattr(self.scheduler, "order", 1)567 callback(step_idx, t, latents)568 569 # 8. Post-processing570 image = self.decode_latents(latents)571 572 # 9. Run safety checker573 image, has_nsfw_concept = self.run_safety_checker(image, device, text_embeddings.dtype)574 575 # 10. Convert to PIL576 if output_type == "pil":577 image = self.numpy_to_pil(image)578 579 if not return_dict:580 return (image, has_nsfw_concept)581 582 return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)583 