diff --git a/LCM/LCM_lora_inpaint.py b/LCM/LCM_lora_inpaint.py index b0d178d..6235b48 100644 --- a/LCM/LCM_lora_inpaint.py +++ b/LCM/LCM_lora_inpaint.py @@ -46,7 +46,8 @@ from diffusers.models.attention import BasicTransformerBlock from diffusers import StableDiffusionPipeline from diffusers.models.unet_2d_blocks import CrossAttnDownBlock2D, CrossAttnUpBlock2D, DownBlock2D, UpBlock2D from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import rescale_noise_cfg - +import io, base64, json +from urllib import request logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -2011,12 +2012,31 @@ class LCM_inpaint_final( # call the callback, if provided if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): progress_bar.update() - image = self.vae2.decode(latents / self.vae2.config.scaling_factor, return_dict=False, generator=generator)[0] - do_denormalize = [True] * image.shape[0] - image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) - image = image[0] - par = os.path.abspath(os.path.join(os.path.join(os.path.realpath(__file__), os.pardir), os.pardir)) - image.save(f"{par}/CanvasToolLone/taesd.png") + try: + image = self.vae2.decode(latents / self.vae2.config.scaling_factor, return_dict=False, generator=generator)[0] + do_denormalize = [True] * image.shape[0] + image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) + image = image[0] + width = image.size[0] + height = image.size[1] + buf = io.BytesIO() + image.save(buf, format='PNG') + byte_im = buf.getvalue() + byte_im = base64.b64encode(byte_im).decode('utf-8') + byte_im = f"data:image/png;base64,{byte_im}" + p = { + "data":{ + "img":byte_im, + "width":image.size[0], + "height":image.size[1] + } + } + data = json.dumps(p).encode('utf-8') + req = request.Request("http://localhost:5000/settaesd", data=data) + req.add_header("Content-Type", "application/json") + request.urlopen(req) + except: + pass if callback is not None and i % callback_steps == 0: step_idx = i // getattr(self.scheduler, "order", 1) callback(step_idx, t, latents) diff --git a/LCM/LCM_lora_inpaint_ipadapter.py b/LCM/LCM_lora_inpaint_ipadapter.py index 1b32f87..2daf9d1 100644 --- a/LCM/LCM_lora_inpaint_ipadapter.py +++ b/LCM/LCM_lora_inpaint_ipadapter.py @@ -47,6 +47,9 @@ from diffusers import StableDiffusionPipeline from diffusers.models.unet_2d_blocks import CrossAttnDownBlock2D, CrossAttnUpBlock2D, DownBlock2D, UpBlock2D from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import rescale_noise_cfg +import io, base64, json +from urllib import request + logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -1513,12 +1516,31 @@ class LCM_lora_inpaint_ipadapter( # call the callback, if provided if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): progress_bar.update() - image = self.vae2.decode(latents / self.vae2.config.scaling_factor, return_dict=False, generator=generator)[0] - do_denormalize = [True] * image.shape[0] - image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) - image = image[0] - par = os.path.abspath(os.path.join(os.path.join(os.path.realpath(__file__), os.pardir), os.pardir)) - image.save(f"{par}/CanvasToolLone/taesd.png") + try: + image = self.vae2.decode(latents / self.vae2.config.scaling_factor, return_dict=False, generator=generator)[0] + do_denormalize = [True] * image.shape[0] + image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) + image = image[0] + width = image.size[0] + height = image.size[1] + buf = io.BytesIO() + image.save(buf, format='PNG') + byte_im = buf.getvalue() + byte_im = base64.b64encode(byte_im).decode('utf-8') + byte_im = f"data:image/png;base64,{byte_im}" + p = { + "data":{ + "img":byte_im, + "width":image.size[0], + "height":image.size[1] + } + } + data = json.dumps(p).encode('utf-8') + req = request.Request("http://localhost:5000/settaesd", data=data) + req.add_header("Content-Type", "application/json") + request.urlopen(req) + except: + pass if callback is not None and i % callback_steps == 0: step_idx = i // getattr(self.scheduler, "order", 1) callback(step_idx, t, latents) diff --git a/LCM/pipeline_inpaint_cn_reference.py b/LCM/pipeline_inpaint_cn_reference.py index a13898e..62a4e7c 100644 --- a/LCM/pipeline_inpaint_cn_reference.py +++ b/LCM/pipeline_inpaint_cn_reference.py @@ -39,6 +39,8 @@ from diffusers.utils.torch_utils import randn_tensor, is_compiled_module import PIL.Image +import base64, io, json +from urllib import request logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -1096,8 +1098,24 @@ class LatentConsistencyModelPipeline_refinpaintcn(DiffusionPipeline): do_denormalize = [True] * image.shape[0] image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) image = image[0] - par = os.path.abspath(os.path.join(os.path.join(os.path.realpath(__file__), os.pardir), os.pardir)) - image.save(f"{par}/CanvasToolLone/taesd.png") + width = image.size[0] + height = image.size[1] + buf = io.BytesIO() + image.save(buf, format='PNG') + byte_im = buf.getvalue() + byte_im = base64.b64encode(byte_im).decode('utf-8') + byte_im = f"data:image/png;base64,{byte_im}" + p = { + "data":{ + "img":byte_im, + "width":image.size[0], + "height":image.size[1] + } + } + data = json.dumps(p).encode('utf-8') + req = request.Request("http://localhost:5000/settaesd", data=data) + req.add_header("Content-Type", "application/json") + request.urlopen(req) denoised = denoised.to(prompt_embeds.dtype) if hasattr(self, "final_offload_hook") and self.final_offload_hook is not None: diff --git a/LCM/stable_diffusion_reference_img2img_controlnet.py b/LCM/stable_diffusion_reference_img2img_controlnet.py index 2534a86..aefce0f 100644 --- a/LCM/stable_diffusion_reference_img2img_controlnet.py +++ b/LCM/stable_diffusion_reference_img2img_controlnet.py @@ -45,7 +45,8 @@ from diffusers import StableDiffusionPipeline from diffusers.models.attention import BasicTransformerBlock from diffusers.models.unet_2d_blocks import CrossAttnDownBlock2D, CrossAttnUpBlock2D, DownBlock2D, UpBlock2D from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion import rescale_noise_cfg - +import io, base64, json +from urllib import request logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -184,6 +185,7 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( def __init__( self, vae: AutoencoderKL, + vae2: AutoencoderKL, text_encoder: CLIPTextModel, tokenizer: CLIPTokenizer, unet: UNet2DConditionModel, @@ -217,6 +219,7 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( self.register_modules( vae=vae, + vae2 = vae2, text_encoder=text_encoder, tokenizer=tokenizer, unet=unet, @@ -987,6 +990,9 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( style_fidelity: float = 0.5, reference_attn: bool = True, reference_adain: bool = True, + ref_enabled: bool = True, + ip_enabled: bool = True, + cn_enabled: bool = True, **kwargs, ): r""" @@ -1096,22 +1102,23 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( "Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`", ) - controlnet = self.controlnet._orig_mod if is_compiled_module(self.controlnet) else self.controlnet + if cn_enabled: + controlnet = self.controlnet._orig_mod if is_compiled_module(self.controlnet) else self.controlnet - # align format for control guidance - if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list): - control_guidance_start = len(control_guidance_end) * [control_guidance_start] - elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list): - control_guidance_end = len(control_guidance_start) * [control_guidance_end] - elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list): - mult = len(controlnet.nets) if isinstance(controlnet, MultiControlNetModel) else 1 - control_guidance_start, control_guidance_end = ( - mult * [control_guidance_start], - mult * [control_guidance_end], - ) + # align format for control guidance + if not isinstance(control_guidance_start, list) and isinstance(control_guidance_end, list): + control_guidance_start = len(control_guidance_end) * [control_guidance_start] + elif not isinstance(control_guidance_end, list) and isinstance(control_guidance_start, list): + control_guidance_end = len(control_guidance_start) * [control_guidance_end] + elif not isinstance(control_guidance_start, list) and not isinstance(control_guidance_end, list): + mult = len(controlnet.nets) if isinstance(controlnet, MultiControlNetModel) else 1 + control_guidance_start, control_guidance_end = ( + mult * [control_guidance_start], + mult * [control_guidance_end], + ) # 1. Check inputs. Raise error if not correct - self.check_inputs( + '''self.check_inputs( prompt, control_image, callback_steps, @@ -1122,7 +1129,7 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( control_guidance_start, control_guidance_end, callback_on_step_end_tensor_inputs, - ) + )''' self._guidance_scale = guidance_scale self._clip_skip = clip_skip @@ -1138,15 +1145,16 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( device = self._execution_device - if isinstance(controlnet, MultiControlNetModel) and isinstance(controlnet_conditioning_scale, float): - controlnet_conditioning_scale = [controlnet_conditioning_scale] * len(controlnet.nets) + if cn_enabled: + if isinstance(controlnet, MultiControlNetModel) and isinstance(controlnet_conditioning_scale, float): + controlnet_conditioning_scale = [controlnet_conditioning_scale] * len(controlnet.nets) - global_pool_conditions = ( - controlnet.config.global_pool_conditions - if isinstance(controlnet, ControlNetModel) - else controlnet.nets[0].config.global_pool_conditions - ) - guess_mode = guess_mode or global_pool_conditions + global_pool_conditions = ( + controlnet.config.global_pool_conditions + if isinstance(controlnet, ControlNetModel) + else controlnet.nets[0].config.global_pool_conditions + ) + guess_mode = guess_mode or global_pool_conditions # 3. Encode input prompt text_encoder_lora_scale = ( @@ -1163,15 +1171,16 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( lora_scale=text_encoder_lora_scale, clip_skip=self.clip_skip, ) - ref_image = self.prepare_image( - image=ref_image, - width=width, - height=height, - batch_size=batch_size * num_images_per_prompt, - num_images_per_prompt=num_images_per_prompt, - device=device, - dtype=prompt_embeds.dtype, - ) + if ref_enabled: + ref_image = self.prepare_image( + image=ref_image, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=prompt_embeds.dtype, + ) # For classifier free guidance, we need to do two forward passes. # Here we concatenate the unconditional and text embeddings into a single batch # to avoid doing two forward passes @@ -1185,25 +1194,11 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( # 4. Prepare image image = self.image_processor.preprocess(image, height=height, width=width).to(dtype=torch.float32) - # 5. Prepare controlnet_conditioning_image - if isinstance(controlnet, ControlNetModel): - control_image = self.prepare_control_image( - image=control_image, - width=width, - height=height, - batch_size=batch_size * num_images_per_prompt, - num_images_per_prompt=num_images_per_prompt, - device=device, - dtype=controlnet.dtype, - do_classifier_free_guidance=self.do_classifier_free_guidance, - guess_mode=guess_mode, - ) - elif isinstance(controlnet, MultiControlNetModel): - control_images = [] - - for control_image_ in control_image: - control_image_ = self.prepare_control_image( - image=control_image_, + if cn_enabled: + # 5. Prepare controlnet_conditioning_image + if isinstance(controlnet, ControlNetModel): + control_image = self.prepare_control_image( + image=control_image, width=width, height=height, batch_size=batch_size * num_images_per_prompt, @@ -1213,12 +1208,27 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( do_classifier_free_guidance=self.do_classifier_free_guidance, guess_mode=guess_mode, ) + elif isinstance(controlnet, MultiControlNetModel): + control_images = [] - control_images.append(control_image_) + for control_image_ in control_image: + control_image_ = self.prepare_control_image( + image=control_image_, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=controlnet.dtype, + do_classifier_free_guidance=self.do_classifier_free_guidance, + guess_mode=guess_mode, + ) - control_image = control_images - else: - assert False + control_images.append(control_image_) + + control_image = control_images + else: + assert False # 5. Prepare timesteps self.scheduler.set_timesteps(num_inference_steps, device=device) @@ -1236,383 +1246,389 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( device, generator, ) - ref_image_latents = self.prepare_ref_latents( - ref_image, - batch_size * num_images_per_prompt, - prompt_embeds.dtype, - device, - generator, - self.do_classifier_free_guidance, - ) - - # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + if ref_enabled: + ref_image_latents = self.prepare_ref_latents( + ref_image, + batch_size * num_images_per_prompt, + prompt_embeds.dtype, + device, + generator, + self.do_classifier_free_guidance, + ) + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) - MODE = "write" - uc_mask = ( - torch.Tensor([1] * batch_size * num_images_per_prompt + [0] * batch_size * num_images_per_prompt) - .type_as(ref_image_latents) - .bool() - ) added_cond_kwargs = {"image_embeds": image_embeds} if ip_adapter_image is not None else None do_classifier_free_guidance = self.do_classifier_free_guidance - def hacked_basic_transformer_inner_forward( - self, - hidden_states: torch.FloatTensor, - attention_mask: Optional[torch.FloatTensor] = None, - encoder_hidden_states: Optional[torch.FloatTensor] = None, - encoder_attention_mask: Optional[torch.FloatTensor] = None, - timestep: Optional[torch.LongTensor] = None, - cross_attention_kwargs: Dict[str, Any] = None, - class_labels: Optional[torch.LongTensor] = None, - ): - if self.use_ada_layer_norm: - norm_hidden_states = self.norm1(hidden_states, timestep) - elif self.use_ada_layer_norm_zero: - norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( - hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype - ) - else: - norm_hidden_states = self.norm1(hidden_states) - # 1. Self-Attention - cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - if self.only_cross_attention: - attn_output = self.attn1( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, - attention_mask=attention_mask, - **cross_attention_kwargs, - ) - else: - if MODE == "write": - self.bank.append(norm_hidden_states.detach().clone()) + # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + if ref_enabled: + #extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + MODE = "write" + uc_mask = ( + torch.Tensor([1] * batch_size * num_images_per_prompt + [0] * batch_size * num_images_per_prompt) + .type_as(ref_image_latents) + .bool() + ) + #added_cond_kwargs = {"image_embeds": image_embeds} if ip_adapter_image is not None else None + #do_classifier_free_guidance = self.do_classifier_free_guidance + def hacked_basic_transformer_inner_forward( + self, + hidden_states: torch.FloatTensor, + attention_mask: Optional[torch.FloatTensor] = None, + encoder_hidden_states: Optional[torch.FloatTensor] = None, + encoder_attention_mask: Optional[torch.FloatTensor] = None, + timestep: Optional[torch.LongTensor] = None, + cross_attention_kwargs: Dict[str, Any] = None, + class_labels: Optional[torch.LongTensor] = None, + ): + if self.use_ada_layer_norm: + norm_hidden_states = self.norm1(hidden_states, timestep) + elif self.use_ada_layer_norm_zero: + norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( + hidden_states, timestep, class_labels, hidden_dtype=hidden_states.dtype + ) + else: + norm_hidden_states = self.norm1(hidden_states) + + # 1. Self-Attention + cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} + if self.only_cross_attention: attn_output = self.attn1( norm_hidden_states, encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, attention_mask=attention_mask, **cross_attention_kwargs, ) - if MODE == "read": - if attention_auto_machine_weight > self.attn_weight: - attn_output_uc = self.attn1( - norm_hidden_states, - encoder_hidden_states=torch.cat([norm_hidden_states] + self.bank, dim=1), - # attention_mask=attention_mask, - **cross_attention_kwargs, - ) - attn_output_c = attn_output_uc.clone() - if do_classifier_free_guidance and style_fidelity > 0: - attn_output_c[uc_mask] = self.attn1( - norm_hidden_states[uc_mask], - encoder_hidden_states=norm_hidden_states[uc_mask], - **cross_attention_kwargs, - ) - attn_output = style_fidelity * attn_output_c + (1.0 - style_fidelity) * attn_output_uc - self.bank.clear() - else: + else: + if MODE == "write": + self.bank.append(norm_hidden_states.detach().clone()) attn_output = self.attn1( norm_hidden_states, encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, attention_mask=attention_mask, **cross_attention_kwargs, ) - if self.use_ada_layer_norm_zero: - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = attn_output + hidden_states - - if self.attn2 is not None: - norm_hidden_states = ( - self.norm2(hidden_states, timestep) if self.use_ada_layer_norm else self.norm2(hidden_states) - ) - - # 2. Cross-Attention - attn_output = self.attn2( - norm_hidden_states, - encoder_hidden_states=encoder_hidden_states, - attention_mask=encoder_attention_mask, - **cross_attention_kwargs, - ) + if MODE == "read": + if attention_auto_machine_weight > self.attn_weight: + attn_output_uc = self.attn1( + norm_hidden_states, + encoder_hidden_states=torch.cat([norm_hidden_states] + self.bank, dim=1), + # attention_mask=attention_mask, + **cross_attention_kwargs, + ) + attn_output_c = attn_output_uc.clone() + if do_classifier_free_guidance and style_fidelity > 0: + attn_output_c[uc_mask] = self.attn1( + norm_hidden_states[uc_mask], + encoder_hidden_states=norm_hidden_states[uc_mask], + **cross_attention_kwargs, + ) + attn_output = style_fidelity * attn_output_c + (1.0 - style_fidelity) * attn_output_uc + self.bank.clear() + else: + attn_output = self.attn1( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states if self.only_cross_attention else None, + attention_mask=attention_mask, + **cross_attention_kwargs, + ) + if self.use_ada_layer_norm_zero: + attn_output = gate_msa.unsqueeze(1) * attn_output hidden_states = attn_output + hidden_states - # 3. Feed-forward - norm_hidden_states = self.norm3(hidden_states) + if self.attn2 is not None: + norm_hidden_states = ( + self.norm2(hidden_states, timestep) if self.use_ada_layer_norm else self.norm2(hidden_states) + ) - if self.use_ada_layer_norm_zero: - norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + # 2. Cross-Attention + attn_output = self.attn2( + norm_hidden_states, + encoder_hidden_states=encoder_hidden_states, + attention_mask=encoder_attention_mask, + **cross_attention_kwargs, + ) + hidden_states = attn_output + hidden_states - ff_output = self.ff(norm_hidden_states) + # 3. Feed-forward + norm_hidden_states = self.norm3(hidden_states) - if self.use_ada_layer_norm_zero: - ff_output = gate_mlp.unsqueeze(1) * ff_output + if self.use_ada_layer_norm_zero: + norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] - hidden_states = ff_output + hidden_states + ff_output = self.ff(norm_hidden_states) - return hidden_states + if self.use_ada_layer_norm_zero: + ff_output = gate_mlp.unsqueeze(1) * ff_output - def hacked_mid_forward(self, *args, **kwargs): - eps = 1e-6 - x = self.original_forward(*args, **kwargs) - if MODE == "write": - if gn_auto_machine_weight >= self.gn_weight: - var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) - self.mean_bank.append(mean) - self.var_bank.append(var) - if MODE == "read": - if len(self.mean_bank) > 0 and len(self.var_bank) > 0: - var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) - std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 - mean_acc = sum(self.mean_bank) / float(len(self.mean_bank)) - var_acc = sum(self.var_bank) / float(len(self.var_bank)) - std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - x_uc = (((x - mean) / std) * std_acc) + mean_acc - x_c = x_uc.clone() - if do_classifier_free_guidance and style_fidelity > 0: - x_c[uc_mask] = x[uc_mask] - x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc - self.mean_bank = [] - self.var_bank = [] - return x + hidden_states = ff_output + hidden_states - def hack_CrossAttnDownBlock2D_forward( - self, - hidden_states: torch.FloatTensor, - temb: Optional[torch.FloatTensor] = None, - encoder_hidden_states: Optional[torch.FloatTensor] = None, - attention_mask: Optional[torch.FloatTensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - encoder_attention_mask: Optional[torch.FloatTensor] = None, - ): - eps = 1e-6 + return hidden_states - # TODO(Patrick, William) - attention mask is not used - output_states = () - - for i, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)): - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] + def hacked_mid_forward(self, *args, **kwargs): + eps = 1e-6 + x = self.original_forward(*args, **kwargs) if MODE == "write": if gn_auto_machine_weight >= self.gn_weight: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - self.mean_bank.append([mean]) - self.var_bank.append([var]) + var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) + self.mean_bank.append(mean) + self.var_bank.append(var) if MODE == "read": if len(self.mean_bank) > 0 and len(self.var_bank) > 0: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 - mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) - var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) + mean_acc = sum(self.mean_bank) / float(len(self.mean_bank)) + var_acc = sum(self.var_bank) / float(len(self.var_bank)) std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc - hidden_states_c = hidden_states_uc.clone() + x_uc = (((x - mean) / std) * std_acc) + mean_acc + x_c = x_uc.clone() if do_classifier_free_guidance and style_fidelity > 0: - hidden_states_c[uc_mask] = hidden_states[uc_mask] - hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc + x_c[uc_mask] = x[uc_mask] + x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc + self.mean_bank = [] + self.var_bank = [] + return x - output_states = output_states + (hidden_states,) + def hack_CrossAttnDownBlock2D_forward( + self, + hidden_states: torch.FloatTensor, + temb: Optional[torch.FloatTensor] = None, + encoder_hidden_states: Optional[torch.FloatTensor] = None, + attention_mask: Optional[torch.FloatTensor] = None, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + encoder_attention_mask: Optional[torch.FloatTensor] = None, + ): + eps = 1e-6 - if MODE == "read": - self.mean_bank = [] - self.var_bank = [] + # TODO(Patrick, William) - attention mask is not used + output_states = () - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) + for i, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)): + hidden_states = resnet(hidden_states, temb) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + attention_mask=attention_mask, + encoder_attention_mask=encoder_attention_mask, + return_dict=False, + )[0] + if MODE == "write": + if gn_auto_machine_weight >= self.gn_weight: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + self.mean_bank.append([mean]) + self.var_bank.append([var]) + if MODE == "read": + if len(self.mean_bank) > 0 and len(self.var_bank) > 0: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) + var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) + std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 + hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc + hidden_states_c = hidden_states_uc.clone() + if do_classifier_free_guidance and style_fidelity > 0: + hidden_states_c[uc_mask] = hidden_states[uc_mask] + hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc - output_states = output_states + (hidden_states,) + output_states = output_states + (hidden_states,) - return hidden_states, output_states - - def hacked_DownBlock2D_forward(self, hidden_states, temb=None,**kwargs): - eps = 1e-6 - - output_states = () - - for i, resnet in enumerate(self.resnets): - hidden_states = resnet(hidden_states, temb) - - if MODE == "write": - if gn_auto_machine_weight >= self.gn_weight: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - self.mean_bank.append([mean]) - self.var_bank.append([var]) if MODE == "read": - if len(self.mean_bank) > 0 and len(self.var_bank) > 0: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 - mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) - var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) - std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc - hidden_states_c = hidden_states_uc.clone() - if do_classifier_free_guidance and style_fidelity > 0: - hidden_states_c[uc_mask] = hidden_states[uc_mask] - hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc + self.mean_bank = [] + self.var_bank = [] - output_states = output_states + (hidden_states,) + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) - if MODE == "read": - self.mean_bank = [] - self.var_bank = [] + output_states = output_states + (hidden_states,) - if self.downsamplers is not None: - for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states) + return hidden_states, output_states - output_states = output_states + (hidden_states,) + def hacked_DownBlock2D_forward(self, hidden_states, temb=None,**kwargs): + eps = 1e-6 - return hidden_states, output_states + output_states = () - def hacked_CrossAttnUpBlock2D_forward( - self, - hidden_states: torch.FloatTensor, - res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], - temb: Optional[torch.FloatTensor] = None, - encoder_hidden_states: Optional[torch.FloatTensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, - upsample_size: Optional[int] = None, - attention_mask: Optional[torch.FloatTensor] = None, - encoder_attention_mask: Optional[torch.FloatTensor] = None, - ): - eps = 1e-6 - # TODO(Patrick, William) - attention mask is not used - for i, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - hidden_states = resnet(hidden_states, temb) - hidden_states = attn( - hidden_states, - encoder_hidden_states=encoder_hidden_states, - cross_attention_kwargs=cross_attention_kwargs, - attention_mask=attention_mask, - encoder_attention_mask=encoder_attention_mask, - return_dict=False, - )[0] + for i, resnet in enumerate(self.resnets): + hidden_states = resnet(hidden_states, temb) + + if MODE == "write": + if gn_auto_machine_weight >= self.gn_weight: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + self.mean_bank.append([mean]) + self.var_bank.append([var]) + if MODE == "read": + if len(self.mean_bank) > 0 and len(self.var_bank) > 0: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) + var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) + std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 + hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc + hidden_states_c = hidden_states_uc.clone() + if do_classifier_free_guidance and style_fidelity > 0: + hidden_states_c[uc_mask] = hidden_states[uc_mask] + hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc + + output_states = output_states + (hidden_states,) - if MODE == "write": - if gn_auto_machine_weight >= self.gn_weight: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - self.mean_bank.append([mean]) - self.var_bank.append([var]) if MODE == "read": - if len(self.mean_bank) > 0 and len(self.var_bank) > 0: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 - mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) - var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) - std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc - hidden_states_c = hidden_states_uc.clone() - if do_classifier_free_guidance and style_fidelity > 0: - hidden_states_c[uc_mask] = hidden_states[uc_mask] - hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc + self.mean_bank = [] + self.var_bank = [] - if MODE == "read": - self.mean_bank = [] - self.var_bank = [] + if self.downsamplers is not None: + for downsampler in self.downsamplers: + hidden_states = downsampler(hidden_states) - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) + output_states = output_states + (hidden_states,) - return hidden_states + return hidden_states, output_states - def hacked_UpBlock2D_forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None,**kwargs): - eps = 1e-6 - for i, resnet in enumerate(self.resnets): - # pop res hidden states - res_hidden_states = res_hidden_states_tuple[-1] - res_hidden_states_tuple = res_hidden_states_tuple[:-1] - hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) - hidden_states = resnet(hidden_states, temb) + def hacked_CrossAttnUpBlock2D_forward( + self, + hidden_states: torch.FloatTensor, + res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], + temb: Optional[torch.FloatTensor] = None, + encoder_hidden_states: Optional[torch.FloatTensor] = None, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + upsample_size: Optional[int] = None, + attention_mask: Optional[torch.FloatTensor] = None, + encoder_attention_mask: Optional[torch.FloatTensor] = None, + ): + eps = 1e-6 + # TODO(Patrick, William) - attention mask is not used + for i, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)): + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + hidden_states = resnet(hidden_states, temb) + hidden_states = attn( + hidden_states, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + attention_mask=attention_mask, + encoder_attention_mask=encoder_attention_mask, + return_dict=False, + )[0] + + if MODE == "write": + if gn_auto_machine_weight >= self.gn_weight: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + self.mean_bank.append([mean]) + self.var_bank.append([var]) + if MODE == "read": + if len(self.mean_bank) > 0 and len(self.var_bank) > 0: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) + var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) + std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 + hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc + hidden_states_c = hidden_states_uc.clone() + if do_classifier_free_guidance and style_fidelity > 0: + hidden_states_c[uc_mask] = hidden_states[uc_mask] + hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc - if MODE == "write": - if gn_auto_machine_weight >= self.gn_weight: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - self.mean_bank.append([mean]) - self.var_bank.append([var]) if MODE == "read": - if len(self.mean_bank) > 0 and len(self.var_bank) > 0: - var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) - std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 - mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) - var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) - std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc - hidden_states_c = hidden_states_uc.clone() - if do_classifier_free_guidance and style_fidelity > 0: - hidden_states_c[uc_mask] = hidden_states[uc_mask] - hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc + self.mean_bank = [] + self.var_bank = [] - if MODE == "read": - self.mean_bank = [] - self.var_bank = [] + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, upsample_size) - if self.upsamplers is not None: - for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size) + return hidden_states - return hidden_states + def hacked_UpBlock2D_forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None,**kwargs): + eps = 1e-6 + for i, resnet in enumerate(self.resnets): + # pop res hidden states + res_hidden_states = res_hidden_states_tuple[-1] + res_hidden_states_tuple = res_hidden_states_tuple[:-1] + hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1) + hidden_states = resnet(hidden_states, temb) - if reference_attn: - attn_modules = [module for module in torch_dfs(self.unet) if isinstance(module, BasicTransformerBlock)] - attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) + if MODE == "write": + if gn_auto_machine_weight >= self.gn_weight: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + self.mean_bank.append([mean]) + self.var_bank.append([var]) + if MODE == "read": + if len(self.mean_bank) > 0 and len(self.var_bank) > 0: + var, mean = torch.var_mean(hidden_states, dim=(2, 3), keepdim=True, correction=0) + std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + mean_acc = sum(self.mean_bank[i]) / float(len(self.mean_bank[i])) + var_acc = sum(self.var_bank[i]) / float(len(self.var_bank[i])) + std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 + hidden_states_uc = (((hidden_states - mean) / std) * std_acc) + mean_acc + hidden_states_c = hidden_states_uc.clone() + if do_classifier_free_guidance and style_fidelity > 0: + hidden_states_c[uc_mask] = hidden_states[uc_mask] + hidden_states = style_fidelity * hidden_states_c + (1.0 - style_fidelity) * hidden_states_uc - for i, module in enumerate(attn_modules): - module._original_inner_forward = module.forward - module.forward = hacked_basic_transformer_inner_forward.__get__(module, BasicTransformerBlock) - module.bank = [] - module.attn_weight = float(i) / float(len(attn_modules)) + if MODE == "read": + self.mean_bank = [] + self.var_bank = [] - if reference_adain: - gn_modules = [self.unet.mid_block] - self.unet.mid_block.gn_weight = 0 + if self.upsamplers is not None: + for upsampler in self.upsamplers: + hidden_states = upsampler(hidden_states, upsample_size) - down_blocks = self.unet.down_blocks - for w, module in enumerate(down_blocks): - module.gn_weight = 1.0 - float(w) / float(len(down_blocks)) - gn_modules.append(module) + return hidden_states - up_blocks = self.unet.up_blocks - for w, module in enumerate(up_blocks): - module.gn_weight = float(w) / float(len(up_blocks)) - gn_modules.append(module) + if reference_attn: + attn_modules = [module for module in torch_dfs(self.unet) if isinstance(module, BasicTransformerBlock)] + attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0]) - for i, module in enumerate(gn_modules): - if getattr(module, "original_forward", None) is None: - module.original_forward = module.forward - if i == 0: - # mid_block - module.forward = hacked_mid_forward.__get__(module, torch.nn.Module) - elif isinstance(module, CrossAttnDownBlock2D): - module.forward = hack_CrossAttnDownBlock2D_forward.__get__(module, CrossAttnDownBlock2D) - elif isinstance(module, DownBlock2D): - module.forward = hacked_DownBlock2D_forward.__get__(module, DownBlock2D) - elif isinstance(module, CrossAttnUpBlock2D): - module.forward = hacked_CrossAttnUpBlock2D_forward.__get__(module, CrossAttnUpBlock2D) - elif isinstance(module, UpBlock2D): - module.forward = hacked_UpBlock2D_forward.__get__(module, UpBlock2D) - module.mean_bank = [] - module.var_bank = [] - module.gn_weight *= 2 + for i, module in enumerate(attn_modules): + module._original_inner_forward = module.forward + module.forward = hacked_basic_transformer_inner_forward.__get__(module, BasicTransformerBlock) + module.bank = [] + module.attn_weight = float(i) / float(len(attn_modules)) + if reference_adain: + gn_modules = [self.unet.mid_block] + self.unet.mid_block.gn_weight = 0 + down_blocks = self.unet.down_blocks + for w, module in enumerate(down_blocks): + module.gn_weight = 1.0 - float(w) / float(len(down_blocks)) + gn_modules.append(module) + + up_blocks = self.unet.up_blocks + for w, module in enumerate(up_blocks): + module.gn_weight = float(w) / float(len(up_blocks)) + gn_modules.append(module) + + for i, module in enumerate(gn_modules): + if getattr(module, "original_forward", None) is None: + module.original_forward = module.forward + if i == 0: + # mid_block + module.forward = hacked_mid_forward.__get__(module, torch.nn.Module) + elif isinstance(module, CrossAttnDownBlock2D): + module.forward = hack_CrossAttnDownBlock2D_forward.__get__(module, CrossAttnDownBlock2D) + elif isinstance(module, DownBlock2D): + module.forward = hacked_DownBlock2D_forward.__get__(module, DownBlock2D) + elif isinstance(module, CrossAttnUpBlock2D): + module.forward = hacked_CrossAttnUpBlock2D_forward.__get__(module, CrossAttnUpBlock2D) + elif isinstance(module, UpBlock2D): + module.forward = hacked_UpBlock2D_forward.__get__(module, UpBlock2D) + module.mean_bank = [] + module.var_bank = [] + module.gn_weight *= 2 + + if cn_enabled: # 7.1 Create tensor stating which controlnets to keep - controlnet_keep = [] - for i in range(len(timesteps)): - keeps = [ - 1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e) - for s, e in zip(control_guidance_start, control_guidance_end) - ] - controlnet_keep.append(keeps[0] if isinstance(controlnet, ControlNetModel) else keeps) + controlnet_keep = [] + for i in range(len(timesteps)): + keeps = [ + 1.0 - float(i / len(timesteps) < s or (i + 1) / len(timesteps) > e) + for s, e in zip(control_guidance_start, control_guidance_end) + ] + controlnet_keep.append(keeps[0] if isinstance(controlnet, ControlNetModel) else keeps) # 8. Denoising loop num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order @@ -1621,79 +1637,122 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( # expand the latents if we are doing classifier free guidance latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + if cn_enabled: + # controlnet(s) inference + if guess_mode and self.do_classifier_free_guidance: + # Infer ControlNet only for the conditional batch. + control_model_input = latents + control_model_input = self.scheduler.scale_model_input(control_model_input, t) + controlnet_prompt_embeds = prompt_embeds.chunk(2)[1] + else: + control_model_input = latent_model_input + controlnet_prompt_embeds = prompt_embeds - # controlnet(s) inference - if guess_mode and self.do_classifier_free_guidance: - # Infer ControlNet only for the conditional batch. - control_model_input = latents - control_model_input = self.scheduler.scale_model_input(control_model_input, t) - controlnet_prompt_embeds = prompt_embeds.chunk(2)[1] + if isinstance(controlnet_keep[i], list): + cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])] + else: + controlnet_cond_scale = controlnet_conditioning_scale + if isinstance(controlnet_cond_scale, list): + controlnet_cond_scale = controlnet_cond_scale[0] + cond_scale = controlnet_cond_scale * controlnet_keep[i] + + down_block_res_samples, mid_block_res_sample = self.controlnet( + control_model_input, + t, + encoder_hidden_states=controlnet_prompt_embeds, + controlnet_cond=control_image, + conditioning_scale=cond_scale, + guess_mode=guess_mode, + return_dict=False, + ) + + if guess_mode and self.do_classifier_free_guidance: + # Infered ControlNet only for the conditional batch. + # To apply the output of ControlNet to both the unconditional and conditional batches, + # add 0 to the unconditional batch to keep it unchanged. + down_block_res_samples = [torch.cat([torch.zeros_like(d), d]) for d in down_block_res_samples] + mid_block_res_sample = torch.cat([torch.zeros_like(mid_block_res_sample), mid_block_res_sample]) + + if ref_enabled: + noise = randn_tensor( + ref_image_latents.shape, generator=generator, device=device, dtype=ref_image_latents.dtype + ) + ref_xt = self.scheduler.add_noise( + ref_image_latents, + noise, + t.reshape( + 1, + ), + ) + ref_xt = torch.cat([ref_xt] * 2) if self.do_classifier_free_guidance else ref_xt + ref_xt = self.scheduler.scale_model_input(ref_xt, t) + + MODE = "write" + if cn_enabled: + self.unet( + ref_xt, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + down_block_additional_residuals=down_block_res_samples, + mid_block_additional_residual=mid_block_res_sample, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + ) + else: + self.unet( + ref_xt, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + ) + MODE = "read" + # predict the noise residual + if cn_enabled: + noise_pred = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + down_block_additional_residuals=down_block_res_samples, + mid_block_additional_residual=mid_block_res_sample, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + )[0] + else: + noise_pred = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + )[0] else: - control_model_input = latent_model_input - controlnet_prompt_embeds = prompt_embeds + if cn_enabled: + noise_pred = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + down_block_additional_residuals=down_block_res_samples, + mid_block_additional_residual=mid_block_res_sample, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + )[0] + else: + noise_pred = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + cross_attention_kwargs=self.cross_attention_kwargs, + return_dict=False, + added_cond_kwargs = added_cond_kwargs + )[0] - if isinstance(controlnet_keep[i], list): - cond_scale = [c * s for c, s in zip(controlnet_conditioning_scale, controlnet_keep[i])] - else: - controlnet_cond_scale = controlnet_conditioning_scale - if isinstance(controlnet_cond_scale, list): - controlnet_cond_scale = controlnet_cond_scale[0] - cond_scale = controlnet_cond_scale * controlnet_keep[i] - - down_block_res_samples, mid_block_res_sample = self.controlnet( - control_model_input, - t, - encoder_hidden_states=controlnet_prompt_embeds, - controlnet_cond=control_image, - conditioning_scale=cond_scale, - guess_mode=guess_mode, - return_dict=False, - ) - - if guess_mode and self.do_classifier_free_guidance: - # Infered ControlNet only for the conditional batch. - # To apply the output of ControlNet to both the unconditional and conditional batches, - # add 0 to the unconditional batch to keep it unchanged. - down_block_res_samples = [torch.cat([torch.zeros_like(d), d]) for d in down_block_res_samples] - mid_block_res_sample = torch.cat([torch.zeros_like(mid_block_res_sample), mid_block_res_sample]) - - - noise = randn_tensor( - ref_image_latents.shape, generator=generator, device=device, dtype=ref_image_latents.dtype - ) - ref_xt = self.scheduler.add_noise( - ref_image_latents, - noise, - t.reshape( - 1, - ), - ) - ref_xt = torch.cat([ref_xt] * 2) if self.do_classifier_free_guidance else ref_xt - ref_xt = self.scheduler.scale_model_input(ref_xt, t) - - MODE = "write" - self.unet( - ref_xt, - t, - encoder_hidden_states=prompt_embeds, - cross_attention_kwargs=self.cross_attention_kwargs, - down_block_additional_residuals=down_block_res_samples, - mid_block_additional_residual=mid_block_res_sample, - return_dict=False, - added_cond_kwargs = added_cond_kwargs - ) - MODE = "read" - # predict the noise residual - noise_pred = self.unet( - latent_model_input, - t, - encoder_hidden_states=prompt_embeds, - cross_attention_kwargs=self.cross_attention_kwargs, - down_block_additional_residuals=down_block_res_samples, - mid_block_additional_residual=mid_block_res_sample, - return_dict=False, - added_cond_kwargs = added_cond_kwargs - )[0] # perform guidance if self.do_classifier_free_guidance: @@ -1716,6 +1775,31 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( # call the callback, if provided if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): progress_bar.update() + try: + image = self.vae2.decode(latents / self.vae2.config.scaling_factor, return_dict=False, generator=generator)[0] + do_denormalize = [True] * image.shape[0] + image = self.image_processor.postprocess(image, output_type=output_type, do_denormalize=do_denormalize) + image = image[0] + width = image.size[0] + height = image.size[1] + buf = io.BytesIO() + image.save(buf, format='PNG') + byte_im = buf.getvalue() + byte_im = base64.b64encode(byte_im).decode('utf-8') + byte_im = f"data:image/png;base64,{byte_im}" + p = { + "data":{ + "img":byte_im, + "width":image.size[0], + "height":image.size[1] + } + } + data = json.dumps(p).encode('utf-8') + req = request.Request("http://localhost:5000/settaesd", data=data) + req.add_header("Content-Type", "application/json") + request.urlopen(req) + except: + pass if callback is not None and i % callback_steps == 0: step_idx = i // getattr(self.scheduler, "order", 1) callback(step_idx, t, latents) @@ -1749,4 +1833,4 @@ class StableDiffusionControlNetImg2ImgPipeline_ref( if not return_dict: return (image, has_nsfw_concept) - return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept) \ No newline at end of file + return StableDiffusionPipelineOutput(images=image, nsfw_content_detected=has_nsfw_concept)