"""ChromaInpaintPipeline implements a text-guided image inpainting pipeline for the lodestones/Chroma1-HD model,based on the ChromaPipeline from Hugging Face Diffusers:contentReference[oaicite:0]{index=0} and the Stable Diffusion inpainting approach:contentReference[oaicite:1]{index=1}."""importgcimporttimeimporttorchfromtypingimportList, Union, Optionalfromdiffusers.modelsimportAutoencoderKLfromdiffusers.schedulersimportFlowMatchEulerDiscreteSchedulerfromdiffusers.models.transformersimportChromaTransformer2DModelfromtransformersimportT5EncoderModel, T5TokenizerFastfromdiffusers.pipelines.pipeline_utilsimportDiffusionPipeline, ImagePipelineOutputfromdiffusers.utils.torch_utilsimportrandn_tensorfromPILimportImagefromtransformersimport (
CLIPImageProcessor,
CLIPVisionModelWithProjection,
)
importnumpyasnpimporttorch.nn.functionalasFfromtqdm.autoimporttqdm# Add progress bar library# Copyright 2025 Black Forest Labs and The HuggingFace Team. All rights reserved.## Licensed under the Apache License, Version 2.0 (the "License");# you may not use this file except in compliance with the License.# You may obtain a copy of the License at## http://www.apache.org/licenses/LICENSE-2.0## Unless required by applicable law or agreed to in writing, software# distributed under the License is distributed on an "AS IS" BASIS,# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.# See the License for the specific language governing permissions and# limitations under the License.importgcimportinspectfromtypingimportAny, Callable, Dict, List, Optional, UnionimportnumpyasnpimportPIL.Imageimporttorchfromtransformersimport (
CLIPImageProcessor,
CLIPVisionModelWithProjection,
T5EncoderModel,
T5TokenizerFast,
)
fromdiffusers.image_processorimportPipelineImageInput, VaeImageProcessorfromdiffusers.loadersimportFluxIPAdapterMixin, FluxLoraLoaderMixin, TextualInversionLoaderMixinfromdiffusers.models.autoencodersimportAutoencoderKLfromdiffusers.models.transformersimportChromaTransformer2DModelfromdiffusers.schedulersimportDDIMSchedulerfromdiffusers.utilsimport (
USE_PEFT_BACKEND,
is_torch_xla_available,
logging,
replace_example_docstring,
scale_lora_layers,
unscale_lora_layers,
)
fromdiffusers.utils.torch_utilsimportrandn_tensorfromdiffusers.pipelines.pipeline_utilsimportDiffusionPipelinefromdiffusers.pipelines.chroma.pipeline_outputimportChromaPipelineOutputfromdiffusers.schedulersimportDDIMSchedulerfromdiffusers.models.transformersimportChromaTransformer2DModelifis_torch_xla_available():
importtorch_xla.core.xla_modelasxmXLA_AVAILABLE=Trueelse:
XLA_AVAILABLE=Falselogger=logging.get_logger(__name__) # pylint: disable=invalid-nameEXAMPLE_DOC_STRING=""" Examples: ```py >>> import torch >>> from diffusers import FluxInpaintPipeline >>> from diffusers.utils import load_image >>> pipe = FluxInpaintPipeline.from_pretrained("black-forest-labs/FLUX.1-schnell", torch_dtype=torch.bfloat16) >>> pipe.to("cuda") >>> prompt = "Face of a yellow cat, high resolution, sitting on a park bench" >>> img_url = "https://raw.githubusercontent.com/CompVis/latent-diffusion/main/data/inpainting_examples/overture-creations-5sI6fQgYIuo.png" >>> mask_url = "https://raw.githubusercontent.com/CompVis/latent-diffusion/main/data/inpainting_examples/overture-creations-5sI6fQgYIuo_mask.png" >>> source = load_image(img_url) >>> mask = load_image(mask_url) >>> image = pipe(prompt=prompt, image=source, mask_image=mask).images[0] >>> image.save("flux_inpainting.png") ```"""# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shiftdefcalculate_shift(
image_seq_len,
base_seq_len: int=256,
max_seq_len: int=4096,
base_shift: float=0.5,
max_shift: float=1.15,
):
m= (max_shift-base_shift) / (max_seq_len-base_seq_len)
b=base_shift-m*base_seq_lenmu=image_seq_len*m+breturnmu# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latentsdefretrieve_latents(
encoder_output: torch.Tensor, generator: Optional[torch.Generator] =None, sample_mode: str="sample"
):
ifhasattr(encoder_output, "latent_dist") andsample_mode=="sample":
returnencoder_output.latent_dist.sample(generator)
elifhasattr(encoder_output, "latent_dist") andsample_mode=="argmax":
returnencoder_output.latent_dist.mode()
elifhasattr(encoder_output, "latents"):
returnencoder_output.latentselse:
raiseAttributeError("Could not access latents of provided encoder_output")
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timestepsdefretrieve_timesteps(
scheduler,
num_inference_steps: Optional[int] =None,
device: Optional[Union[str, torch.device]] =None,
timesteps: Optional[List[int]] =None,
sigmas: Optional[List[float]] =None,
**kwargs,
):
r""" Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. Args: scheduler (`SchedulerMixin`): The scheduler to get timesteps from. num_inference_steps (`int`): The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` must be `None`. device (`str` or `torch.device`, *optional*): The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. timesteps (`List[int]`, *optional*): Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed, `num_inference_steps` and `sigmas` must be `None`. sigmas (`List[float]`, *optional*): Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed, `num_inference_steps` and `timesteps` must be `None`. Returns: `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the second element is the number of inference steps. """iftimestepsisnotNoneandsigmasisnotNone:
raiseValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
iftimestepsisnotNone:
accepts_timesteps="timesteps"inset(inspect.signature(scheduler.set_timesteps).parameters.keys())
ifnotaccepts_timesteps:
raiseValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"f" timestep schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
timesteps=scheduler.timestepsnum_inference_steps=len(timesteps)
elifsigmasisnotNone:
accept_sigmas="sigmas"inset(inspect.signature(scheduler.set_timesteps).parameters.keys())
ifnotaccept_sigmas:
raiseValueError(
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"f" sigmas schedules. Please check whether you are using the correct scheduler."
)
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
timesteps=scheduler.timestepsnum_inference_steps=len(timesteps)
else:
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
timesteps=scheduler.timestepsreturntimesteps, num_inference_stepsclassChromaInpaintPipeline(DiffusionPipeline, FluxLoraLoaderMixin, FluxIPAdapterMixin):
r""" The Flux pipeline for image inpainting. Reference: https://blackforestlabs.ai/announcing-black-forest-labs/ Args: transformer ([`ChromaTransformer2DModel`]): Conditional Transformer (MMDiT) architecture to denoise the encoded image latents. scheduler ([`DDIMScheduler`]): A scheduler to be used in combination with `transformer` to denoise the encoded image latents. vae ([`AutoencoderKL`]): Variational Auto-Encoder (VAE) Model to encode and decode images to and from latent representations. text_encoder ([`CLIPTextModel`]): [CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant. text_encoder_2 ([`T5EncoderModel`]): [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically the [google/t5-v1_1-xxl](https://huggingface.co/google/t5-v1_1-xxl) variant. tokenizer (`CLIPTokenizer`): Tokenizer of class [CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer). tokenizer_2 (`T5TokenizerFast`): Second Tokenizer of class [T5TokenizerFast](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5TokenizerFast). """model_cpu_offload_seq="text_encoder->text_encoder_2->image_encoder->transformer->vae"_optional_components= ["image_encoder", "feature_extractor"]
_callback_tensor_inputs= ["latents", "prompt_embeds"]
def__init__(
self,
scheduler: FlowMatchEulerDiscreteScheduler,
vae: AutoencoderKL,
text_encoder: T5EncoderModel,
tokenizer: T5TokenizerFast,
transformer: ChromaTransformer2DModel,
image_encoder: CLIPVisionModelWithProjection=None,
feature_extractor: CLIPImageProcessor=None,
):
super().__init__()
self.register_modules(
vae=vae,
text_encoder=text_encoder,
tokenizer=tokenizer,
transformer=transformer,
scheduler=scheduler,
image_encoder=image_encoder,
feature_extractor=feature_extractor,
)
self.vae_scale_factor=2** (len(self.vae.config.block_out_channels) -1) ifgetattr(self, "vae", None) else8self.latent_channels=self.vae.config.latent_channelsifgetattr(self, "vae", None) else16# Flux latents are turned into 2x2 patches and packed. This means the latent width and height has to be divisible# by the patch size. So the vae scale factor is multiplied by the patch size to account for thisself.image_processor=VaeImageProcessor(vae_scale_factor=self.vae_scale_factor*2)
self.default_sample_size=128self.mask_processor=VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor*2,
vae_latent_channels=self.latent_channels,
do_normalize=False,
do_binarize=True,
do_convert_grayscale=True,
)
def_get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] =None,
num_images_per_prompt: int=1,
max_sequence_length: int=512,
device: Optional[torch.device] =None,
dtype: Optional[torch.dtype] =None,
):
device=deviceorself._execution_devicedtype=dtypeorself.text_encoder.dtypeprompt= [prompt] ifisinstance(prompt, str) elsepromptbatch_size=len(prompt)
ifisinstance(self, TextualInversionLoaderMixin):
prompt=self.maybe_convert_prompt(prompt, self.tokenizer)
text_inputs=self.tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_length=False,
return_overflowing_tokens=False,
return_tensors="pt",
)
text_input_ids=text_inputs.input_idstokenizer_mask=text_inputs.attention_masktokenizer_mask_device=tokenizer_mask.to(device)
prompt_embeds=self.text_encoder(
text_input_ids.to(device),
output_hidden_states=False,
attention_mask=tokenizer_mask_device,
)[0]
prompt_embeds=prompt_embeds.to(dtype=dtype, device=device)
seq_lengths=tokenizer_mask_device.sum(dim=1)
mask_indices=torch.arange(tokenizer_mask_device.size(1), device=device).unsqueeze(0).expand(batch_size, -1)
attention_mask= (mask_indices<=seq_lengths.unsqueeze(1)).to(dtype=dtype, device=device)
_, seq_len, _=prompt_embeds.shape# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly methodprompt_embeds=prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds=prompt_embeds.view(batch_size*num_images_per_prompt, seq_len, -1)
attention_mask=attention_mask.repeat(1, num_images_per_prompt)
attention_mask=attention_mask.view(batch_size*num_images_per_prompt, seq_len)
returnprompt_embeds, attention_maskdefencode_prompt(
self,
prompt: Union[str, List[str]],
negative_prompt: Union[str, List[str]] =None,
device: Optional[torch.device] =None,
num_images_per_prompt: int=1,
prompt_embeds: Optional[torch.Tensor] =None,
negative_prompt_embeds: Optional[torch.Tensor] =None,
prompt_attention_mask: Optional[torch.Tensor] =None,
negative_prompt_attention_mask: Optional[torch.Tensor] =None,
do_classifier_free_guidance: bool=True,
max_sequence_length: int=256,
lora_scale: Optional[float] =None,
):
r""" Args: prompt (`str` or `List[str]`, *optional*): prompt to be encoded negative_prompt (`str` or `List[str]`, *optional*): The prompt not to guide the image generation. If not defined, one has to pass `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is less than `1`). device: (`torch.device`): torch device num_images_per_prompt (`int`): number of images that should be generated per prompt prompt_embeds (`torch.Tensor`, *optional*): Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not provided, text embeddings will be generated from `prompt` input argument. lora_scale (`float`, *optional*): A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded. """device=deviceorself._execution_device# set lora scale so that monkey patched LoRA# function of text encoder can correctly access itiflora_scaleisnotNoneandisinstance(self, FluxLoraLoaderMixin):
self._lora_scale=lora_scale# dynamically adjust the LoRA scaleifself.text_encoderisnotNoneandUSE_PEFT_BACKEND:
scale_lora_layers(self.text_encoder, lora_scale)
prompt= [prompt] ifisinstance(prompt, str) elsepromptifpromptisnotNone:
batch_size=len(prompt)
else:
batch_size=prompt_embeds.shape[0]
ifprompt_embedsisNone:
prompt_embeds, prompt_attention_mask=self._get_t5_prompt_embeds(
prompt=prompt,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
)
dtype=self.text_encoder.dtypeifself.text_encoderisnotNoneelseself.transformer.dtypetext_ids=torch.zeros(prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
negative_text_ids=Noneifdo_classifier_free_guidance:
ifnegative_prompt_embedsisNone:
negative_prompt=negative_promptor""negative_prompt= (
batch_size* [negative_prompt] ifisinstance(negative_prompt, str) elsenegative_prompt
)
ifpromptisnotNoneandtype(prompt) isnottype(negative_prompt):
raiseTypeError(
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="f" {type(prompt)}."
)
elifbatch_size!=len(negative_prompt):
raiseValueError(
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"" the batch size of `prompt`."
)
negative_prompt_embeds, negative_prompt_attention_mask=self._get_t5_prompt_embeds(
prompt=negative_prompt,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
device=device,
)
negative_text_ids=torch.zeros(negative_prompt_embeds.shape[1], 3).to(device=device, dtype=dtype)
ifself.text_encoderisnotNone:
ifisinstance(self, FluxLoraLoaderMixin) andUSE_PEFT_BACKEND:
# Retrieve the original scale by scaling back the LoRA layersunscale_lora_layers(self.text_encoder, lora_scale)
return (
prompt_embeds,
text_ids,
prompt_attention_mask,
negative_prompt_embeds,
negative_text_ids,
negative_prompt_attention_mask,
)
# Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline.encode_imagedefencode_image(self, image, device, num_images_per_prompt):
dtype=next(self.image_encoder.parameters()).dtypeifnotisinstance(image, torch.Tensor):
image=self.feature_extractor(image, return_tensors="pt").pixel_valuesimage=image.to(device=device, dtype=dtype)
image_embeds=self.image_encoder(image).image_embedsimage_embeds=image_embeds.repeat_interleave(num_images_per_prompt, dim=0)
returnimage_embeds# Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline.prepare_ip_adapter_image_embedsdefprepare_ip_adapter_image_embeds(
self, ip_adapter_image, ip_adapter_image_embeds, device, num_images_per_prompt
):
image_embeds= []
ifip_adapter_image_embedsisNone:
ifnotisinstance(ip_adapter_image, list):
ip_adapter_image= [ip_adapter_image]
iflen(ip_adapter_image) !=self.transformer.encoder_hid_proj.num_ip_adapters:
raiseValueError(
f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {self.transformer.encoder_hid_proj.num_ip_adapters} IP Adapters."
)
forsingle_ip_adapter_imageinip_adapter_image:
single_image_embeds=self.encode_image(single_ip_adapter_image, device, 1)
image_embeds.append(single_image_embeds[None, :])
else:
ifnotisinstance(ip_adapter_image_embeds, list):
ip_adapter_image_embeds= [ip_adapter_image_embeds]
iflen(ip_adapter_image_embeds) !=self.transformer.encoder_hid_proj.num_ip_adapters:
raiseValueError(
f"`ip_adapter_image_embeds` must have same length as the number of IP Adapters. Got {len(ip_adapter_image_embeds)} image embeds and {self.transformer.encoder_hid_proj.num_ip_adapters} IP Adapters."
)
forsingle_image_embedsinip_adapter_image_embeds:
image_embeds.append(single_image_embeds)
ip_adapter_image_embeds= []
forsingle_image_embedsinimage_embeds:
single_image_embeds=torch.cat([single_image_embeds] *num_images_per_prompt, dim=0)
single_image_embeds=single_image_embeds.to(device=device)
ip_adapter_image_embeds.append(single_image_embeds)
returnip_adapter_image_embeds# Copied from diffusers.pipelines.stable_diffusion_3.pipeline_stable_diffusion_3_inpaint.StableDiffusion3InpaintPipeline._encode_vae_imagedef_encode_vae_image(self, image: torch.Tensor, generator: torch.Generator):
ifisinstance(generator, list):
image_latents= [
retrieve_latents(self.vae.encode(image[i : i+1]), generator=generator[i])
foriinrange(image.shape[0])
]
image_latents=torch.cat(image_latents, dim=0)
else:
image_latents=retrieve_latents(self.vae.encode(image), generator=generator)
image_latents= (image_latents-self.vae.config.shift_factor) *self.vae.config.scaling_factorreturnimage_latents# Copied from diffusers.pipelines.stable_diffusion_3.pipeline_stable_diffusion_3_img2img.StableDiffusion3Img2ImgPipeline.get_timestepsdefget_timesteps(self, num_inference_steps, strength, device):
# get the original timestep using init_timestepinit_timestep=min(num_inference_steps*strength, num_inference_steps)
t_start=int(max(num_inference_steps-init_timestep, 0))
timesteps=self.scheduler.timesteps[t_start*self.scheduler.order :]
ifhasattr(self.scheduler, "set_begin_index"):
self.scheduler.set_begin_index(t_start*self.scheduler.order)
returntimesteps, num_inference_steps-t_startdefcheck_inputs(
self,
prompt,
height,
width,
strength,
negative_prompt=None,
prompt_embeds=None,
negative_prompt_embeds=None,
prompt_attention_mask=None,
negative_prompt_attention_mask=None,
callback_on_step_end_tensor_inputs=None,
max_sequence_length=None,
image=None,
mask_image=None,
padding_mask_crop=None,
):
ifstrength<0orstrength>1:
raiseValueError(f"The value of strength should in [0.0, 1.0] but is {strength}")
ifheight% (self.vae_scale_factor*2) !=0orwidth% (self.vae_scale_factor*2) !=0:
logger.warning(
f"`height` and `width` have to be divisible by {self.vae_scale_factor*2} but are {height} and {width}. Dimensions will be resized accordingly"
)
ifcallback_on_step_end_tensor_inputsisnotNoneandnotall(
kinself._callback_tensor_inputsforkincallback_on_step_end_tensor_inputs
):
raiseValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[kforkincallback_on_step_end_tensor_inputsifknotinself._callback_tensor_inputs]}"
)
ifpromptisnotNoneandprompt_embedsisnotNone:
raiseValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"" only forward one of the two."
)
elifpromptisNoneandprompt_embedsisNone:
raiseValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elifpromptisnotNoneand (notisinstance(prompt, str) andnotisinstance(prompt, list)):
raiseValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
ifnegative_promptisnotNoneandnegative_prompt_embedsisnotNone:
raiseValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
ifprompt_embedsisnotNoneandprompt_attention_maskisNone:
raiseValueError("Cannot provide `prompt_embeds` without also providing `prompt_attention_mask")
ifnegative_prompt_embedsisnotNoneandnegative_prompt_attention_maskisNone:
raiseValueError(
"Cannot provide `negative_prompt_embeds` without also providing `negative_prompt_attention_mask"
)
ifmax_sequence_lengthisnotNoneandmax_sequence_length>512:
raiseValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}")
ifpadding_mask_cropisnotNone:
ifnotisinstance(image, PIL.Image.Image):
raiseValueError(
f"The image should be a PIL image when inpainting mask crop, but is of type {type(image)}."
)
ifnotisinstance(mask_image, PIL.Image.Image):
raiseValueError(
f"The mask image should be a PIL image when inpainting mask crop, but is of type"f" {type(mask_image)}."
)
defcheck_inputs(
self,
prompt,
image,
mask_image,
strength,
height,
width,
negative_prompt=None,
prompt_embeds=None,
negative_prompt_embeds=None,
pooled_prompt_embeds=None,
negative_pooled_prompt_embeds=None,
callback_on_step_end_tensor_inputs=None,
padding_mask_crop=None,
max_sequence_length=None,
prompt_attention_mask=None,
negative_prompt_attention_mask=None,
):
ifstrength<0orstrength>1:
raiseValueError(f"The value of strength should in [0.0, 1.0] but is {strength}")
ifheight% (self.vae_scale_factor*2) !=0orwidth% (self.vae_scale_factor*2) !=0:
logger.warning(
f"`height` and `width` have to be divisible by {self.vae_scale_factor*2} but are {height} and {width}. Dimensions will be resized accordingly"
)
ifcallback_on_step_end_tensor_inputsisnotNoneandnotall(
kinself._callback_tensor_inputsforkincallback_on_step_end_tensor_inputs
):
raiseValueError(
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[kforkincallback_on_step_end_tensor_inputsifknotinself._callback_tensor_inputs]}"
)
ifpromptisnotNoneandprompt_embedsisnotNone:
raiseValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"" only forward one of the two."
)
elifpromptisNoneandprompt_embedsisNone:
raiseValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elifpromptisnotNoneand (notisinstance(prompt, str) andnotisinstance(prompt, list)):
raiseValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
ifnegative_promptisnotNoneandnegative_prompt_embedsisnotNone:
raiseValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
ifprompt_embedsisnotNoneandnegative_prompt_embedsisnotNone:
ifprompt_embeds.shape!=negative_prompt_embeds.shape:
raiseValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"f" {negative_prompt_embeds.shape}."
)
ifprompt_embedsisnotNoneandpooled_prompt_embedsisNone:
raiseValueError(
"If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`."
)
ifnegative_prompt_embedsisnotNoneandnegative_pooled_prompt_embedsisNone:
raiseValueError(
"If `negative_prompt_embeds` are provided, `negative_pooled_prompt_embeds` also have to be passed. Make sure to generate `negative_pooled_prompt_embeds` from the same text encoder that was used to generate `negative_prompt_embeds`."
)
ifprompt_embedsisnotNoneandprompt_attention_maskisNone:
raiseValueError("Cannot provide `prompt_embeds` without also providing `prompt_attention_mask")
ifnegative_prompt_embedsisnotNoneandnegative_prompt_attention_maskisNone:
raiseValueError(
"Cannot provide `negative_prompt_embeds` without also providing `negative_prompt_attention_mask"
)
ifpadding_mask_cropisnotNone:
ifnotisinstance(image, PIL.Image.Image):
raiseValueError(
f"The image should be a PIL image when inpainting mask crop, but is of type {type(image)}."
)
ifnotisinstance(mask_image, PIL.Image.Image):
raiseValueError(
f"The mask image should be a PIL image when inpainting mask crop, but is of type"f" {type(mask_image)}."
)
ifoutput_type!="pil":
raiseValueError(f"The output type should be PIL when inpainting mask crop, but is {output_type}.")
ifmax_sequence_lengthisnotNoneandmax_sequence_length>512:
raiseValueError(f"`max_sequence_length` cannot be greater than 512 but is {max_sequence_length}")
@staticmethod# Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline._prepare_latent_image_idsdef_prepare_latent_image_ids(batch_size, height, width, device, dtype):
latent_image_ids=torch.zeros(height, width, 3)
latent_image_ids[..., 1] =latent_image_ids[..., 1] +torch.arange(height)[:, None]
latent_image_ids[..., 2] =latent_image_ids[..., 2] +torch.arange(width)[None, :]
latent_image_id_height, latent_image_id_width, latent_image_id_channels=latent_image_ids.shapelatent_image_ids=latent_image_ids.reshape(
latent_image_id_height*latent_image_id_width, latent_image_id_channels
)
returnlatent_image_ids.to(device=device, dtype=dtype)
@staticmethod# Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline._pack_latentsdef_pack_latents(latents, batch_size, num_channels_latents, height, width):
latents=latents.view(batch_size, num_channels_latents, height//2, 2, width//2, 2)
latents=latents.permute(0, 2, 4, 1, 3, 5)
latents=latents.reshape(batch_size, (height//2) * (width//2), num_channels_latents*4)
returnlatents@staticmethod# Copied from diffusers.pipelines.flux.pipeline_flux.FluxPipeline._unpack_latentsdef_unpack_latents(latents, height, width, vae_scale_factor):
batch_size, num_patches, channels=latents.shape# VAE applies 8x compression on images but we must also account for packing which requires# latent height and width to be divisible by 2.height=2* (int(height) // (vae_scale_factor*2))
width=2* (int(width) // (vae_scale_factor*2))
latents=latents.view(batch_size, height//2, width//2, channels//4, 2, 2)
latents=latents.permute(0, 3, 1, 4, 2, 5)
latents=latents.reshape(batch_size, channels// (2*2), height, width)
returnlatentsdefprepare_latents(
self,
image,
timestep,
batch_size,
num_channels_latents,
height,
width,
dtype,
device,
generator,
latents=None,
):
ifisinstance(generator, list) andlen(generator) !=batch_size:
raiseValueError(
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"f" size of {batch_size}. Make sure the batch size matches the length of the generators."
)
# VAE applies 8x compression on images but we must also account for packing which requires# latent height and width to be divisible by 2.height=2* (int(height) // (self.vae_scale_factor*2))
width=2* (int(width) // (self.vae_scale_factor*2))
shape= (batch_size, num_channels_latents, height, width)
latent_image_ids=self._prepare_latent_image_ids(batch_size, height//2, width//2, device, dtype)
image=image.to(device=device, dtype=dtype)
ifimage.shape[1] !=self.latent_channels:
image_latents=self._encode_vae_image(image=image, generator=generator)
else:
image_latents=imageifbatch_size>image_latents.shape[0] andbatch_size%image_latents.shape[0] ==0:
# expand init_latents for batch_sizeadditional_image_per_prompt=batch_size//image_latents.shape[0]
image_latents=torch.cat([image_latents] *additional_image_per_prompt, dim=0)
elifbatch_size>image_latents.shape[0] andbatch_size%image_latents.shape[0] !=0:
raiseValueError(
f"Cannot duplicate `image` of batch size {image_latents.shape[0]} to {batch_size} text prompts."
)
else:
image_latents=torch.cat([image_latents], dim=0)
iflatentsisNone:
noise=randn_tensor(shape, generator=generator, device=device, dtype=dtype)
latents=self.scheduler.scale_noise(image_latents, timestep, noise)
else:
noise=latents.to(device)
latents=noisenoise=self._pack_latents(noise, batch_size, num_channels_latents, height, width)
image_latents=self._pack_latents(image_latents, batch_size, num_channels_latents, height, width)
latents=self._pack_latents(latents, batch_size, num_channels_latents, height, width)
returnlatents, noise, image_latents, latent_image_idsdefprepare_mask_latents(
self,
mask,
masked_image,
batch_size,
num_channels_latents,
num_images_per_prompt,
height,
width,
dtype,
device,
generator,
):
# VAE applies 8x compression on images but we must also account for packing which requires# latent height and width to be divisible by 2.height=2* (int(height) // (self.vae_scale_factor*2))
width=2* (int(width) // (self.vae_scale_factor*2))
# resize the mask to latents shape as we concatenate the mask to the latents# we do that before converting to dtype to avoid breaking in case we're using cpu_offload# and half precisionmask=torch.nn.functional.interpolate(mask, size=(height, width))
mask=mask.to(device=device, dtype=dtype)
batch_size=batch_size*num_images_per_promptmasked_image=masked_image.to(device=device, dtype=dtype)
ifmasked_image.shape[1] ==16:
masked_image_latents=masked_imageelse:
masked_image_latents=retrieve_latents(self.vae.encode(masked_image), generator=generator)
masked_image_latents= (
masked_image_latents-self.vae.config.shift_factor
) *self.vae.config.scaling_factor# duplicate mask and masked_image_latents for each generation per prompt, using mps friendly methodifmask.shape[0] <batch_size:
ifnotbatch_size%mask.shape[0] ==0:
raiseValueError(
"The passed mask and the required batch size don't match. Masks are supposed to be duplicated to"f" a total batch size of {batch_size}, but {mask.shape[0]} masks were passed. Make sure the number"" of masks that you pass is divisible by the total requested batch size."
)
mask=mask.repeat(batch_size//mask.shape[0], 1, 1, 1)
ifmasked_image_latents.shape[0] <batch_size:
ifnotbatch_size%masked_image_latents.shape[0] ==0:
raiseValueError(
"The passed images and the required batch size don't match. Images are supposed to be duplicated"f" to a total batch size of {batch_size}, but {masked_image_latents.shape[0]} images were passed."" Make sure the number of images that you pass is divisible by the total requested batch size."
)
masked_image_latents=masked_image_latents.repeat(batch_size//masked_image_latents.shape[0], 1, 1, 1)
# aligning device to prevent device errors when concating it with the latent model inputmasked_image_latents=masked_image_latents.to(device=device, dtype=dtype)
masked_image_latents=self._pack_latents(
masked_image_latents,
batch_size,
num_channels_latents,
height,
width,
)
mask=self._pack_latents(
mask.repeat(1, num_channels_latents, 1, 1),
batch_size,
num_channels_latents,
height,
width,
)
returnmask, masked_image_latents@propertydefguidance_scale(self):
returnself._guidance_scale@propertydefjoint_attention_kwargs(self):
returnself._joint_attention_kwargs@propertydefnum_timesteps(self):
returnself._num_timesteps@propertydefinterrupt(self):
returnself._interruptdef_prepare_attention_mask(
self,
batch_size,
sequence_length,
dtype,
attention_mask=None,
):
ifattention_maskisNone:
returnattention_mask# Extend the prompt attention mask to account for image tokens in the final sequenceattention_mask=torch.cat(
[attention_mask, torch.ones(batch_size, sequence_length, device=attention_mask.device)],
dim=1,
)
attention_mask=attention_mask.to(dtype)
returnattention_mask@propertydefdo_classifier_free_guidance(self):
returnself._guidance_scale>1@replace_example_docstring(EXAMPLE_DOC_STRING)@torch.no_grad()def__call__(
self,
prompt: Union[str, List[str]] =None,
negative_prompt: Union[str, List[str]] =None,
true_cfg_scale: float=1.0,
image: PipelineImageInput=None,
mask_image: PipelineImageInput=None,
masked_image_latents: PipelineImageInput=None,
height: Optional[int] =None,
width: Optional[int] =None,
padding_mask_crop: Optional[int] =None,
strength: float=0.6,
num_inference_steps: int=28,
sigmas: Optional[List[float]] =None,
guidance_scale: float=7.0,
num_images_per_prompt: Optional[int] =1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] =None,
latents: Optional[torch.FloatTensor] =None,
prompt_embeds: Optional[torch.FloatTensor] =None,
pooled_prompt_embeds: Optional[torch.FloatTensor] =None,
ip_adapter_image: Optional[PipelineImageInput] =None,
ip_adapter_image_embeds: Optional[List[torch.Tensor]] =None,
negative_ip_adapter_image: Optional[PipelineImageInput] =None,
negative_ip_adapter_image_embeds: Optional[List[torch.Tensor]] =None,
negative_prompt_embeds: Optional[torch.FloatTensor] =None,
negative_pooled_prompt_embeds: Optional[torch.FloatTensor] =None,
output_type: Optional[str] ="pil",
return_dict: bool=True,
joint_attention_kwargs: Optional[Dict[str, Any]] =None,
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] =None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
max_sequence_length: int=256,
prompt_attention_mask: Optional[torch.Tensor] =None,
negative_prompt_attention_mask: Optional[torch.Tensor] =None,
):
r""" Function invoked when calling the pipeline for generation. Args: prompt (`str` or `List[str]`, *optional*): The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. instead. negative_prompt (`str` or `List[str]`, *optional*): The prompt or prompts not to guide the image generation. If not defined, one has to pass `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is not greater than `1`). height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): The height in pixels of the generated image. This is set to 1024 by default for the best results. width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): The width in pixels of the generated image. This is set to 1024 by default for the best results. num_inference_steps (`int`, *optional*, defaults to 35): The number of denoising steps. More denoising steps usually lead to a higher quality image at the expense of slower inference. sigmas (`List[float]`, *optional*): Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed will be used. guidance_scale (`float`, *optional*, defaults to 3.5): Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2. of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting `guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, usually at the expense of lower image quality. strength (`float, *optional*, defaults to 0.9): Conceptually, indicates how much to transform the reference image. Must be between 0 and 1. image will be used as a starting point, adding more noise to it the larger the strength. The number of denoising steps depends on the amount of noise initially added. When strength is 1, added noise will be maximum and the denoising process will run for the full number of iterations specified in num_inference_steps. A value of 1, therefore, essentially ignores image. num_images_per_prompt (`int`, *optional*, defaults to 1): The number of images to generate per prompt. generator (`torch.Generator` or `List[torch.Generator]`, *optional*): One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make generation deterministic. latents (`torch.Tensor`, *optional*): Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image generation. Can be used to tweak the same generation with different prompts. If not provided, a latents tensor will be generated by sampling using the supplied random `generator`. prompt_embeds (`torch.Tensor`, *optional*): Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not provided, text embeddings will be generated from `prompt` input argument. ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters. ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*): Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not provided, embeddings are computed from the `ip_adapter_image` input argument. negative_ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters. negative_ip_adapter_image_embeds (`List[torch.Tensor]`, *optional*): Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. If not provided, embeddings are computed from the `ip_adapter_image` input argument. negative_prompt_embeds (`torch.Tensor`, *optional*): Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input argument. prompt_attention_mask (torch.Tensor, *optional*): Attention mask for the prompt embeddings. Used to mask out padding tokens in the prompt sequence. Chroma requires a single padding token remain unmasked. Please refer to https://huggingface.co/lodestones/Chroma#tldr-masking-t5-padding-tokens-enhanced-fidelity-and-increased-stability-during-training negative_prompt_attention_mask (torch.Tensor, *optional*): Attention mask for the negative prompt embeddings. Used to mask out padding tokens in the negative prompt sequence. Chroma requires a single padding token remain unmasked. PLease refer to https://huggingface.co/lodestones/Chroma#tldr-masking-t5-padding-tokens-enhanced-fidelity-and-increased-stability-during-training output_type (`str`, *optional*, defaults to `"pil"`): The output format of the generate image. Choose between [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. return_dict (`bool`, *optional*, defaults to `True`): Whether or not to return a [`~pipelines.flux.ChromaPipelineOutput`] instead of a plain tuple. joint_attention_kwargs (`dict`, *optional*): A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under `self.processor` in [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). callback_on_step_end (`Callable`, *optional*): A function that calls at the end of each denoising steps during the inference. The function is called with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by `callback_on_step_end_tensor_inputs`. callback_on_step_end_tensor_inputs (`List`, *optional*): The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the `._callback_tensor_inputs` attribute of your pipeline class. max_sequence_length (`int` defaults to 512): Maximum sequence length to use with the `prompt`. Examples: Returns: [`~pipelines.chroma.ChromaPipelineOutput`] or `tuple`: [`~pipelines.chroma.ChromaPipelineOutput`] if `return_dict` is True, otherwise a `tuple`. When returning a tuple, the first element is a list with the generated images. """height=heightorself.default_sample_size*self.vae_scale_factorwidth=widthorself.default_sample_size*self.vae_scale_factor# 1. Check inputs. Raise error if not correctself.check_inputs(
prompt=prompt,
height=height,
width=width,
strength=strength,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
max_sequence_length=max_sequence_length,
image=image,
mask_image=mask_image,
padding_mask_crop=padding_mask_crop,
)
self._guidance_scale=guidance_scaleself._joint_attention_kwargs=joint_attention_kwargsself._current_timestep=Noneself._interrupt=False# 2. Preprocess mask and imageifpadding_mask_cropisnotNone:
crops_coords=self.mask_processor.get_crop_region(mask_image, width, height, pad=padding_mask_crop)
resize_mode="fill"else:
crops_coords=Noneresize_mode="default"original_image=imageinit_image=self.image_processor.preprocess(
image, height=height, width=width, crops_coords=crops_coords, resize_mode=resize_mode
)
init_image=init_image.to(dtype=torch.float32)
# 3. Define call parametersifpromptisnotNoneandisinstance(prompt, str):
batch_size=1elifpromptisnotNoneandisinstance(prompt, list):
batch_size=len(prompt)
else:
batch_size=prompt_embeds.shape[0]
device=self._execution_devicelora_scale= (
self.joint_attention_kwargs.get("scale", None) ifself.joint_attention_kwargsisnotNoneelseNone
)
self.vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
self.text_encoder.to(device)
(
prompt_embeds,
text_ids,
prompt_attention_mask,
negative_prompt_embeds,
negative_text_ids,
negative_prompt_attention_mask,
) =self.encode_prompt(
prompt=prompt,
negative_prompt=negative_prompt,
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
prompt_attention_mask=prompt_attention_mask,
negative_prompt_attention_mask=negative_prompt_attention_mask,
do_classifier_free_guidance=self.do_classifier_free_guidance,
device=device,
num_images_per_prompt=num_images_per_prompt,
max_sequence_length=max_sequence_length,
lora_scale=lora_scale,
)
self.text_encoder.to('cpu')
gc.collect()
torch.cuda.empty_cache()
self.vae.to(device)
# 4. Prepare timestepssigmas=np.linspace(1.0, 1/num_inference_steps, num_inference_steps) ifsigmasisNoneelsesigmasimage_seq_len= (int(height) //self.vae_scale_factor//2) * (int(width) //self.vae_scale_factor//2)
mu=calculate_shift(
image_seq_len,
self.scheduler.config.get("base_image_seq_len", 256),
self.scheduler.config.get("max_image_seq_len", 4096),
self.scheduler.config.get("base_shift", 0.5),
self.scheduler.config.get("max_shift", 1.15),
)
timesteps, num_inference_steps=retrieve_timesteps(
self.scheduler,
num_inference_steps,
device,
sigmas=sigmas,
mu=mu,
)
timesteps, num_inference_steps=self.get_timesteps(num_inference_steps, strength, device)
num_warmup_steps=max(len(timesteps) -num_inference_steps*self.scheduler.order, 0)
self._num_timesteps=len(timesteps)
ifnum_inference_steps<1:
raiseValueError(
f"After adjusting the num_inference_steps by strength parameter: {strength}, the number of pipeline"f"steps is {num_inference_steps} which is < 1 and not appropriate for this pipeline."
)
latent_timestep=timesteps[:1].repeat(batch_size*num_images_per_prompt)
# 5. Prepare latent variablesnum_channels_latents=self.transformer.config.in_channels//4num_channels_transformer=self.transformer.config.in_channelslatents, noise, image_latents, latent_image_ids=self.prepare_latents(
init_image,
latent_timestep,
batch_size*num_images_per_prompt,
num_channels_latents,
height,
width,
prompt_embeds.dtype,
device,
generator,
latents,
)
mask_condition=self.mask_processor.preprocess(
mask_image, height=height, width=width, resize_mode=resize_mode, crops_coords=crops_coords
)
ifmasked_image_latentsisNone:
masked_image=init_image* (mask_condition<0.5)
else:
masked_image=masked_image_latentsmask, masked_image_latents=self.prepare_mask_latents(
mask_condition,
masked_image,
batch_size,
num_channels_latents,
num_images_per_prompt,
height,
width,
prompt_embeds.dtype,
device,
generator,
)
num_warmup_steps=max(len(timesteps) -num_inference_steps*self.scheduler.order, 0)
self._num_timesteps=len(timesteps)
# handle guidanceifself.transformer.config.guidance_embeds:
guidance=torch.full([1], guidance_scale, device=device, dtype=torch.float32)
guidance=guidance.expand(latents.shape[0])
else:
guidance=Noneif (ip_adapter_imageisnotNoneorip_adapter_image_embedsisnotNone) and (
negative_ip_adapter_imageisNoneandnegative_ip_adapter_image_embedsisNone
):
negative_ip_adapter_image=np.zeros((width, height, 3), dtype=np.uint8)
elif (ip_adapter_imageisNoneandip_adapter_image_embedsisNone) and (
negative_ip_adapter_imageisnotNoneornegative_ip_adapter_image_embedsisnotNone
):
ip_adapter_image=np.zeros((width, height, 3), dtype=np.uint8)
ifself.joint_attention_kwargsisNone:
self._joint_attention_kwargs= {}
image_embeds=Nonenegative_image_embeds=Noneifip_adapter_imageisnotNoneorip_adapter_image_embedsisnotNone:
image_embeds=self.prepare_ip_adapter_image_embeds(
ip_adapter_image,
ip_adapter_image_embeds,
device,
batch_size*num_images_per_prompt,
)
ifnegative_ip_adapter_imageisnotNoneornegative_ip_adapter_image_embedsisnotNone:
negative_image_embeds=self.prepare_ip_adapter_image_embeds(
negative_ip_adapter_image,
negative_ip_adapter_image_embeds,
device,
batch_size*num_images_per_prompt,
)
self.vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
attention_mask=self._prepare_attention_mask(
batch_size=latents.shape[0],
sequence_length=image_seq_len,
dtype=latents.dtype,
attention_mask=prompt_attention_mask,
)
negative_attention_mask=self._prepare_attention_mask(
batch_size=latents.shape[0],
sequence_length=image_seq_len,
dtype=latents.dtype,
attention_mask=negative_prompt_attention_mask,
)
# 6. Denoising loopwithself.progress_bar(total=num_inference_steps) asprogress_bar:
fori, tinenumerate(timesteps):
ifself.interrupt:
continueself._current_timestep=t# broadcast to batch dimension in a way that's compatible with ONNX/Core MLtimestep=t.expand(latents.shape[0])
ifimage_embedsisnotNone:
self._joint_attention_kwargs["ip_adapter_image_embeds"] =image_embedsnoise_pred=self.transformer(
hidden_states=latents,
timestep=timestep/1000,
encoder_hidden_states=prompt_embeds,
txt_ids=text_ids,
img_ids=latent_image_ids,
attention_mask=attention_mask,
joint_attention_kwargs=self.joint_attention_kwargs,
return_dict=False,
)[0]
ifself.do_classifier_free_guidance:
ifnegative_image_embedsisnotNone:
self._joint_attention_kwargs["ip_adapter_image_embeds"] =negative_image_embedsnoise_pred_uncond=self.transformer(
hidden_states=latents,
timestep=timestep/1000,
encoder_hidden_states=negative_prompt_embeds,
txt_ids=negative_text_ids,
img_ids=latent_image_ids,
attention_mask=negative_attention_mask,
joint_attention_kwargs=self.joint_attention_kwargs,
return_dict=False,
)[0]
noise_pred=noise_pred_uncond+guidance_scale* (noise_pred-noise_pred_uncond)
# compute the previous noisy sample x_t -> x_t-1latents_dtype=latents.dtypelatents=self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
# for 64 channel transformer only.init_latents_proper=image_latentsinit_mask=maskifi<len(timesteps) -1:
noise_timestep=timesteps[i+1]
init_latents_proper=self.scheduler.scale_noise(
init_latents_proper, torch.tensor([noise_timestep]), noise
)
latents= (1-init_mask) *init_latents_proper+init_mask*latentsiflatents.dtype!=latents_dtype:
iftorch.backends.mps.is_available():
# some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272latents=latents.to(latents_dtype)
ifcallback_on_step_endisnotNone:
callback_kwargs= {}
forkincallback_on_step_end_tensor_inputs:
callback_kwargs[k] =locals()[k]
callback_outputs=callback_on_step_end(self, i, t, callback_kwargs)
latents=callback_outputs.pop("latents", latents)
prompt_embeds=callback_outputs.pop("prompt_embeds", prompt_embeds)
# call the callback, if providedifi==len(timesteps) -1or ((i+1) >num_warmup_stepsand (i+1) %self.scheduler.order==0):
progress_bar.update()
ifXLA_AVAILABLE:
xm.mark_step()
self._current_timestep=Noneifoutput_type=="latent":
image=latentselse:
self.vae.to(device)
latents=self._unpack_latents(latents, height, width, self.vae_scale_factor)
latents= (latents/self.vae.config.scaling_factor) +self.vae.config.shift_factorimage=self.vae.decode(latents, return_dict=False)[0]
image=self.image_processor.postprocess(image, output_type=output_type)
self.vae.to('cpu')
gc.collect()
torch.cuda.empty_cache()
# Offload all modelsself.maybe_free_model_hooks()
ifnotreturn_dict:
return (image,)
returnChromaPipelineOutput(images=image)
I'd like someone to add an inpaiting pipeline for Chroma
I created a working one here: see the gist file
Working example: