diff --git a/brushnet/brushnet.py b/brushnet/brushnet.py index 208492c..4cc32eb 100644 --- a/brushnet/brushnet.py +++ b/brushnet/brushnet.py @@ -138,7 +138,10 @@ class CrossAttnUpBlock2D(nn.Module): return_res_samples: Optional[bool]=False, up_block_add_samples: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 + if cross_attention_kwargs is not None: + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") + is_freeu_enabled = ( getattr(self, "s1", None) and getattr(self, "s2", None) @@ -194,7 +197,7 @@ class CrossAttnUpBlock2D(nn.Module): return_dict=False, )[0] else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -210,7 +213,7 @@ class CrossAttnUpBlock2D(nn.Module): if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size, scale=lora_scale) + hidden_states = upsampler(hidden_states, upsample_size) if return_res_samples: output_states = output_states + (hidden_states,) if up_block_add_samples is not None: @@ -409,8 +412,7 @@ class MidBlock2D(nn.Module): hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - lora_scale = 1.0 - hidden_states = self.resnets[0](hidden_states, temb, scale=lora_scale) + hidden_states = self.resnets[0](hidden_states, temb) for resnet in self.resnets[1:]: if self.training and self.gradient_checkpointing: @@ -431,7 +433,7 @@ class MidBlock2D(nn.Module): **ckpt_kwargs, ) else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) return hidden_states diff --git a/brushnet/pipeline_brushnet.py b/brushnet/pipeline_brushnet.py index 8e00e97..76c354d 100644 --- a/brushnet/pipeline_brushnet.py +++ b/brushnet/pipeline_brushnet.py @@ -27,7 +27,7 @@ from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusio from diffusers.pipelines.stable_diffusion.pipeline_output import StableDiffusionPipelineOutput from diffusers.pipelines.stable_diffusion.safety_checker import StableDiffusionSafetyChecker -from comfy.utils import ProgressBar as comfy_pbar +from comfy.utils import ProgressBar logger = logging.get_logger(__name__) # pylint: disable=invalid-name @@ -1164,7 +1164,7 @@ class StableDiffusionBrushNetPipeline( is_unet_compiled = is_compiled_module(self.unet) is_brushnet_compiled = is_compiled_module(self.brushnet) is_torch_higher_equal_2_1 = is_torch_version(">=", "2.1") - comfy_pbar(num_inference_steps) + comfy_pbar = ProgressBar(num_inference_steps) with self.progress_bar(total=num_inference_steps) as progress_bar: for i, t in enumerate(timesteps): # Relevant thread: @@ -1246,7 +1246,7 @@ class StableDiffusionBrushNetPipeline( # 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() - comfy_pbar(1) + comfy_pbar.update(1) 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/brushnet/unet_2d_blocks.py b/brushnet/unet_2d_blocks.py index b01824e..02620b0 100644 --- a/brushnet/unet_2d_blocks.py +++ b/brushnet/unet_2d_blocks.py @@ -18,7 +18,7 @@ import torch import torch.nn.functional as F from torch import nn -from diffusers.utils import is_torch_version, logging +from diffusers.utils import is_torch_version, logging, deprecate from diffusers.utils.torch_utils import apply_freeu from diffusers.models.activations import get_activation from diffusers.models.attention_processor import Attention, AttnAddedKVProcessor, AttnAddedKVProcessor2_0 @@ -856,8 +856,11 @@ class UNetMidBlock2DCrossAttn(nn.Module): cross_attention_kwargs: Optional[Dict[str, Any]] = None, encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 - hidden_states = self.resnets[0](hidden_states, temb, scale=lora_scale) + if cross_attention_kwargs is not None: + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") + + hidden_states = self.resnets[0](hidden_states, temb) for attn, resnet in zip(self.attentions, self.resnets[1:]): if self.training and self.gradient_checkpointing: @@ -894,7 +897,7 @@ class UNetMidBlock2DCrossAttn(nn.Module): encoder_attention_mask=encoder_attention_mask, return_dict=False, )[0] - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) return hidden_states @@ -994,7 +997,8 @@ class UNetMidBlock2DSimpleCrossAttn(nn.Module): encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - lora_scale = cross_attention_kwargs.get("scale", 1.0) + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") if attention_mask is None: # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. @@ -1007,7 +1011,7 @@ class UNetMidBlock2DSimpleCrossAttn(nn.Module): # mask = attention_mask if encoder_hidden_states is None else encoder_attention_mask mask = attention_mask - hidden_states = self.resnets[0](hidden_states, temb, scale=lora_scale) + hidden_states = self.resnets[0](hidden_states, temb) for attn, resnet in zip(self.attentions, self.resnets[1:]): # attn hidden_states = attn( @@ -1018,7 +1022,7 @@ class UNetMidBlock2DSimpleCrossAttn(nn.Module): ) # resnet - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) return hidden_states @@ -1211,23 +1215,22 @@ class AttnDownBlock2D(nn.Module): cross_attention_kwargs: Optional[Dict[str, Any]] = None, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} - - lora_scale = cross_attention_kwargs.get("scale", 1.0) + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") output_states = () for resnet, attn in zip(self.resnets, self.attentions): - cross_attention_kwargs.update({"scale": lora_scale}) - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn(hidden_states, **cross_attention_kwargs) output_states = output_states + (hidden_states,) if self.downsamplers is not None: for downsampler in self.downsamplers: if self.downsample_type == "resnet": - hidden_states = downsampler(hidden_states, temb=temb, scale=lora_scale) + hidden_states = downsampler(hidden_states, temb=temb) else: - hidden_states = downsampler(hidden_states, scale=lora_scale) + hidden_states = downsampler(hidden_states) output_states += (hidden_states,) @@ -1337,9 +1340,11 @@ class CrossAttnDownBlock2D(nn.Module): additional_residuals: Optional[torch.FloatTensor] = None, down_block_add_samples: Optional[torch.FloatTensor] = None, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: - output_states = () + if cross_attention_kwargs is not None: + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 + output_states = () blocks = list(zip(self.resnets, self.attentions)) @@ -1371,7 +1376,7 @@ class CrossAttnDownBlock2D(nn.Module): return_dict=False, )[0] else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -1392,7 +1397,7 @@ class CrossAttnDownBlock2D(nn.Module): if self.downsamplers is not None: for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, scale=lora_scale) + hidden_states = downsampler(hidden_states) if down_block_add_samples is not None: hidden_states = hidden_states + down_block_add_samples.pop(0) # todo: add before or after @@ -1455,9 +1460,13 @@ class DownBlock2D(nn.Module): self.gradient_checkpointing = False def forward( - self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0, - down_block_add_samples: Optional[torch.FloatTensor] = None, + self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, + down_block_add_samples: Optional[torch.FloatTensor] = None, *args, **kwargs ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + output_states = () for resnet in self.resnets: @@ -1478,7 +1487,7 @@ class DownBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale=scale) + hidden_states = resnet(hidden_states, temb) if down_block_add_samples is not None: hidden_states = hidden_states + down_block_add_samples.pop(0) @@ -1487,7 +1496,7 @@ class DownBlock2D(nn.Module): if self.downsamplers is not None: for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, scale=scale) + hidden_states = downsampler(hidden_states) if down_block_add_samples is not None: hidden_states = hidden_states + down_block_add_samples.pop(0) # todo: add before or after @@ -1659,15 +1668,18 @@ class AttnDownEncoderBlock2D(nn.Module): else: self.downsamplers = None - def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor: + def forward(self, hidden_states: torch.FloatTensor, *args, **kwargs) -> torch.FloatTensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb=None, scale=scale) - cross_attention_kwargs = {"scale": scale} - hidden_states = attn(hidden_states, **cross_attention_kwargs) + hidden_states = resnet(hidden_states, temb=None) + hidden_states = attn(hidden_states) if self.downsamplers is not None: for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, scale) + hidden_states = downsampler(hidden_states) return hidden_states @@ -1758,18 +1770,18 @@ class AttnSkipDownBlock2D(nn.Module): hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, skip_sample: Optional[torch.FloatTensor] = None, - scale: float = 1.0, + *args, + **kwargs, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...], torch.FloatTensor]: output_states = () for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb, scale=scale) - cross_attention_kwargs = {"scale": scale} - hidden_states = attn(hidden_states, **cross_attention_kwargs) + hidden_states = resnet(hidden_states, temb) + hidden_states = attn(hidden_states) output_states += (hidden_states,) if self.downsamplers is not None: - hidden_states = self.resnet_down(hidden_states, temb, scale=scale) + hidden_states = self.resnet_down(hidden_states, temb) for downsampler in self.downsamplers: skip_sample = downsampler(skip_sample) @@ -1845,16 +1857,21 @@ class SkipDownBlock2D(nn.Module): hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, skip_sample: Optional[torch.FloatTensor] = None, - scale: float = 1.0, + *args, + **kwargs, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...], torch.FloatTensor]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + output_states = () for resnet in self.resnets: - hidden_states = resnet(hidden_states, temb, scale) + hidden_states = resnet(hidden_states, temb) output_states += (hidden_states,) if self.downsamplers is not None: - hidden_states = self.resnet_down(hidden_states, temb, scale) + hidden_states = self.resnet_down(hidden_states, temb) for downsampler in self.downsamplers: skip_sample = downsampler(skip_sample) @@ -1930,8 +1947,12 @@ class ResnetDownsampleBlock2D(nn.Module): self.gradient_checkpointing = False def forward( - self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0 + self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, *args, **kwargs ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + output_states = () for resnet in self.resnets: @@ -1952,13 +1973,13 @@ class ResnetDownsampleBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale) + hidden_states = resnet(hidden_states, temb) output_states = output_states + (hidden_states,) if self.downsamplers is not None: for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, temb, scale) + hidden_states = downsampler(hidden_states, temb) output_states = output_states + (hidden_states,) @@ -2069,10 +2090,11 @@ class SimpleCrossAttnDownBlock2D(nn.Module): cross_attention_kwargs: Optional[Dict[str, Any]] = None, encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: - output_states = () cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - lora_scale = cross_attention_kwargs.get("scale", 1.0) + output_states = () if attention_mask is None: # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. @@ -2105,7 +2127,7 @@ class SimpleCrossAttnDownBlock2D(nn.Module): **cross_attention_kwargs, ) else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, @@ -2118,7 +2140,7 @@ class SimpleCrossAttnDownBlock2D(nn.Module): if self.downsamplers is not None: for downsampler in self.downsamplers: - hidden_states = downsampler(hidden_states, temb, scale=lora_scale) + hidden_states = downsampler(hidden_states, temb) output_states = output_states + (hidden_states,) @@ -2172,8 +2194,12 @@ class KDownBlock2D(nn.Module): self.gradient_checkpointing = False def forward( - self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0 + self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, *args, **kwargs ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + output_states = () for resnet in self.resnets: @@ -2194,7 +2220,7 @@ class KDownBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale) + hidden_states = resnet(hidden_states, temb) output_states += (hidden_states,) @@ -2279,8 +2305,11 @@ class KCrossAttnDownBlock2D(nn.Module): cross_attention_kwargs: Optional[Dict[str, Any]] = None, encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> Tuple[torch.FloatTensor, Tuple[torch.FloatTensor, ...]]: + cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") + output_states = () - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 for resnet, attn in zip(self.resnets, self.attentions): if self.training and self.gradient_checkpointing: @@ -2310,7 +2339,7 @@ class KCrossAttnDownBlock2D(nn.Module): encoder_attention_mask=encoder_attention_mask, ) else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -2430,24 +2459,28 @@ class AttnUpBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, upsample_size: Optional[int] = None, - scale: float = 1.0, + *args, + **kwargs, ) -> torch.FloatTensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for resnet, attn in 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, scale=scale) - cross_attention_kwargs = {"scale": scale} - hidden_states = attn(hidden_states, **cross_attention_kwargs) + hidden_states = resnet(hidden_states, temb) + hidden_states = attn(hidden_states) if self.upsamplers is not None: for upsampler in self.upsamplers: if self.upsample_type == "resnet": - hidden_states = upsampler(hidden_states, temb=temb, scale=scale) + hidden_states = upsampler(hidden_states, temb=temb) else: - hidden_states = upsampler(hidden_states, scale=scale) + hidden_states = upsampler(hidden_states) return hidden_states @@ -2556,7 +2589,10 @@ class CrossAttnUpBlock2D(nn.Module): return_res_samples: Optional[bool]=False, up_block_add_samples: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 + if cross_attention_kwargs is not None: + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") + is_freeu_enabled = ( getattr(self, "s1", None) and getattr(self, "s2", None) @@ -2612,7 +2648,7 @@ class CrossAttnUpBlock2D(nn.Module): return_dict=False, )[0] else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -2628,7 +2664,7 @@ class CrossAttnUpBlock2D(nn.Module): if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size, scale=lora_scale) + hidden_states = upsampler(hidden_states, upsample_size) if return_res_samples: output_states = output_states + (hidden_states,) if up_block_add_samples is not None: @@ -2695,10 +2731,15 @@ class UpBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, upsample_size: Optional[int] = None, - scale: float = 1.0, return_res_samples: Optional[bool]=False, up_block_add_samples: Optional[torch.FloatTensor] = None, + *args, + **kwargs, ) -> torch.FloatTensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + is_freeu_enabled = ( getattr(self, "s1", None) and getattr(self, "s2", None) @@ -2744,7 +2785,7 @@ class UpBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale=scale) + hidden_states = resnet(hidden_states, temb) if return_res_samples: output_states = output_states + (hidden_states,) @@ -2753,7 +2794,7 @@ class UpBlock2D(nn.Module): if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, upsample_size, scale=scale) + hidden_states = upsampler(hidden_states, upsample_size) if return_res_samples: output_states = output_states + (hidden_states,) @@ -2829,11 +2870,9 @@ class UpDecoderBlock2D(nn.Module): self.resolution_idx = resolution_idx - def forward( - self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0 - ) -> torch.FloatTensor: + def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor: for resnet in self.resnets: - hidden_states = resnet(hidden_states, temb=temb, scale=scale) + hidden_states = resnet(hidden_states, temb=temb) if self.upsamplers is not None: for upsampler in self.upsamplers: @@ -2929,17 +2968,14 @@ class AttnUpDecoderBlock2D(nn.Module): self.resolution_idx = resolution_idx - def forward( - self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None, scale: float = 1.0 - ) -> torch.FloatTensor: + def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor: for resnet, attn in zip(self.resnets, self.attentions): - hidden_states = resnet(hidden_states, temb=temb, scale=scale) - cross_attention_kwargs = {"scale": scale} - hidden_states = attn(hidden_states, temb=temb, **cross_attention_kwargs) + hidden_states = resnet(hidden_states, temb=temb) + hidden_states = attn(hidden_states, temb=temb) if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, scale=scale) + hidden_states = upsampler(hidden_states) return hidden_states @@ -3044,18 +3080,22 @@ class AttnSkipUpBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, skip_sample=None, - scale: float = 1.0, + *args, + **kwargs, ) -> Tuple[torch.FloatTensor, torch.FloatTensor]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for resnet in 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, scale=scale) + hidden_states = resnet(hidden_states, temb) - cross_attention_kwargs = {"scale": scale} - hidden_states = self.attentions[0](hidden_states, **cross_attention_kwargs) + hidden_states = self.attentions[0](hidden_states) if skip_sample is not None: skip_sample = self.upsampler(skip_sample) @@ -3069,7 +3109,7 @@ class AttnSkipUpBlock2D(nn.Module): skip_sample = skip_sample + skip_sample_states - hidden_states = self.resnet_up(hidden_states, temb, scale=scale) + hidden_states = self.resnet_up(hidden_states, temb) return hidden_states, skip_sample @@ -3152,15 +3192,20 @@ class SkipUpBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, skip_sample=None, - scale: float = 1.0, + *args, + **kwargs, ) -> Tuple[torch.FloatTensor, torch.FloatTensor]: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for resnet in 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, scale=scale) + hidden_states = resnet(hidden_states, temb) if skip_sample is not None: skip_sample = self.upsampler(skip_sample) @@ -3174,7 +3219,7 @@ class SkipUpBlock2D(nn.Module): skip_sample = skip_sample + skip_sample_states - hidden_states = self.resnet_up(hidden_states, temb, scale=scale) + hidden_states = self.resnet_up(hidden_states, temb) return hidden_states, skip_sample @@ -3254,8 +3299,13 @@ class ResnetUpsampleBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, upsample_size: Optional[int] = None, - scale: float = 1.0, + *args, + **kwargs, ) -> torch.FloatTensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + for resnet in self.resnets: # pop res hidden states res_hidden_states = res_hidden_states_tuple[-1] @@ -3279,11 +3329,11 @@ class ResnetUpsampleBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale=scale) + hidden_states = resnet(hidden_states, temb) if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, temb, scale=scale) + hidden_states = upsampler(hidden_states, temb) return hidden_states @@ -3399,8 +3449,9 @@ class SimpleCrossAttnUpBlock2D(nn.Module): encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") - lora_scale = cross_attention_kwargs.get("scale", 1.0) if attention_mask is None: # if encoder_hidden_states is defined: we are doing cross-attn, so we should use cross-attn mask. mask = None if encoder_hidden_states is None else encoder_attention_mask @@ -3438,7 +3489,7 @@ class SimpleCrossAttnUpBlock2D(nn.Module): **cross_attention_kwargs, ) else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, @@ -3449,7 +3500,7 @@ class SimpleCrossAttnUpBlock2D(nn.Module): if self.upsamplers is not None: for upsampler in self.upsamplers: - hidden_states = upsampler(hidden_states, temb, scale=lora_scale) + hidden_states = upsampler(hidden_states, temb) return hidden_states @@ -3510,8 +3561,13 @@ class KUpBlock2D(nn.Module): res_hidden_states_tuple: Tuple[torch.FloatTensor, ...], temb: Optional[torch.FloatTensor] = None, upsample_size: Optional[int] = None, - scale: float = 1.0, + *args, + **kwargs, ) -> torch.FloatTensor: + if len(args) > 0 or kwargs.get("scale", None) is not None: + deprecation_message = "The `scale` argument is deprecated and will be ignored. Please remove it, as passing it will raise an error in the future. `scale` should directly be passed while calling the underlying pipeline component i.e., via `cross_attention_kwargs`." + deprecate("scale", "1.0.0", deprecation_message) + res_hidden_states_tuple = res_hidden_states_tuple[-1] if res_hidden_states_tuple is not None: hidden_states = torch.cat([hidden_states, res_hidden_states_tuple], dim=1) @@ -3534,7 +3590,7 @@ class KUpBlock2D(nn.Module): create_custom_forward(resnet), hidden_states, temb ) else: - hidden_states = resnet(hidden_states, temb, scale=scale) + hidden_states = resnet(hidden_states, temb) if self.upsamplers is not None: for upsampler in self.upsamplers: @@ -3644,7 +3700,6 @@ class KCrossAttnUpBlock2D(nn.Module): if res_hidden_states_tuple is not None: hidden_states = torch.cat([hidden_states, res_hidden_states_tuple], dim=1) - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 for resnet, attn in zip(self.resnets, self.attentions): if self.training and self.gradient_checkpointing: @@ -3673,7 +3728,7 @@ class KCrossAttnUpBlock2D(nn.Module): encoder_attention_mask=encoder_attention_mask, ) else: - hidden_states = resnet(hidden_states, temb, scale=lora_scale) + hidden_states = resnet(hidden_states, temb) hidden_states = attn( hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -3776,6 +3831,8 @@ class KAttentionBlock(nn.Module): encoder_attention_mask: Optional[torch.FloatTensor] = None, ) -> torch.FloatTensor: cross_attention_kwargs = cross_attention_kwargs if cross_attention_kwargs is not None else {} + if cross_attention_kwargs.get("scale", None) is not None: + logger.warning("Passing `scale` to `cross_attention_kwargs` is deprecated. `scale` will be ignored.") # 1. Self-Attention if self.add_self_attention: diff --git a/brushnet/unet_2d_condition.py b/brushnet/unet_2d_condition.py index 4c0eb0b..783a199 100644 --- a/brushnet/unet_2d_condition.py +++ b/brushnet/unet_2d_condition.py @@ -1188,7 +1188,14 @@ class UNet2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin, cross_attention_kwargs["gligen"] = {"objs": self.position_net(**gligen_args)} # 3. down - lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 + # we're popping the `scale` instead of getting it because otherwise `scale` will be propagated + # to the internal blocks and will raise deprecation warnings. this will be confusing for our users. + if cross_attention_kwargs is not None: + cross_attention_kwargs = cross_attention_kwargs.copy() + lora_scale = cross_attention_kwargs.pop("scale", 1.0) + else: + lora_scale = 1.0 + if USE_PEFT_BACKEND: # weight the lora layers by setting `lora_scale` for each PEFT layer scale_lora_layers(self, lora_scale) @@ -1243,7 +1250,7 @@ class UNet2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin, additional_residuals["down_block_add_samples"] = [down_block_add_samples.pop(0) for _ in range(len(downsample_block.resnets)+(downsample_block.downsamplers !=None))] - sample, res_samples = downsample_block(hidden_states=sample, temb=emb, scale=lora_scale, **additional_residuals) + sample, res_samples = downsample_block(hidden_states=sample, temb=emb, **additional_residuals) if is_adapter and len(down_intrablock_additional_residuals) > 0: sample += down_intrablock_additional_residuals.pop(0) @@ -1328,7 +1335,6 @@ class UNet2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin, temb=emb, res_hidden_states_tuple=res_samples, upsample_size=upsample_size, - scale=lora_scale, **additional_residuals, ) diff --git a/ella/ella_unet.py b/ella/ella_unet.py new file mode 100644 index 0000000..feedecf --- /dev/null +++ b/ella/ella_unet.py @@ -0,0 +1,67 @@ +from typing import Any, Optional, Union, Tuple +import torch +class ELLAProxyUNet(torch.nn.Module): + def __init__(self, ella, unet): + super().__init__() + # In order to still use the diffusers pipeline, including various workaround + + self.ella = ella + self.unet = unet + self.config = unet.config + self.dtype = unet.dtype + self.device = unet.device + + self.flexible_max_length_workaround = None + + def forward( + self, + sample: torch.FloatTensor, + timestep: Union[torch.Tensor, float, int], + encoder_hidden_states: torch.Tensor, + class_labels: Optional[torch.Tensor] = None, + timestep_cond: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + cross_attention_kwargs: Optional[dict[str, Any]] = None, + added_cond_kwargs: Optional[dict[str, torch.Tensor]] = None, + down_block_additional_residuals: Optional[tuple[torch.Tensor]] = None, + mid_block_additional_residual: Optional[torch.Tensor] = None, + down_intrablock_additional_residuals: Optional[tuple[torch.Tensor]] = None, + encoder_attention_mask: Optional[torch.Tensor] = None, + return_dict: bool = True, + down_block_add_samples: Optional[Tuple[torch.Tensor]] = None, + mid_block_add_sample: Optional[Tuple[torch.Tensor]] = None, + up_block_add_samples: Optional[Tuple[torch.Tensor]] = None, + ): + if self.flexible_max_length_workaround is not None: + time_aware_encoder_hidden_state_list = [] + for i, max_length in enumerate(self.flexible_max_length_workaround): + time_aware_encoder_hidden_state_list.append( + self.ella(encoder_hidden_states[i : i + 1, :max_length], timestep) + ) + # No matter how many tokens are text features, the ella output must be 64 tokens. + time_aware_encoder_hidden_states = torch.cat( + time_aware_encoder_hidden_state_list, dim=0 + ) + else: + time_aware_encoder_hidden_states = self.ella( + encoder_hidden_states, timestep + ) + + return self.unet( + sample=sample, + timestep=timestep, + encoder_hidden_states=time_aware_encoder_hidden_states, + class_labels=class_labels, + timestep_cond=timestep_cond, + attention_mask=attention_mask, + cross_attention_kwargs=cross_attention_kwargs, + added_cond_kwargs=added_cond_kwargs, + down_block_additional_residuals=down_block_additional_residuals, + mid_block_additional_residual=mid_block_additional_residual, + down_intrablock_additional_residuals=down_intrablock_additional_residuals, + encoder_attention_mask=encoder_attention_mask, + return_dict=return_dict, + down_block_add_samples=down_block_add_samples, + mid_block_add_sample=mid_block_add_sample, + up_block_add_samples=up_block_add_samples, + ) \ No newline at end of file diff --git a/ella/model.py b/ella/model.py new file mode 100644 index 0000000..f8235da --- /dev/null +++ b/ella/model.py @@ -0,0 +1,218 @@ +from collections import OrderedDict +from typing import Optional + +import torch +import torch.nn as nn +from diffusers.models.embeddings import TimestepEmbedding, Timesteps +from transformers import T5EncoderModel, T5Tokenizer + + +class AdaLayerNorm(nn.Module): + def __init__(self, embedding_dim: int, time_embedding_dim: Optional[int] = None): + super().__init__() + + if time_embedding_dim is None: + time_embedding_dim = embedding_dim + + self.silu = nn.SiLU() + self.linear = nn.Linear(time_embedding_dim, 2 * embedding_dim, bias=True) + nn.init.zeros_(self.linear.weight) + nn.init.zeros_(self.linear.bias) + + self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) + + def forward( + self, x: torch.Tensor, timestep_embedding: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + emb = self.linear(self.silu(timestep_embedding)) + shift, scale = emb.view(len(x), 1, -1).chunk(2, dim=-1) + x = self.norm(x) * (1 + scale) + shift + return x + + +class SquaredReLU(nn.Module): + def forward(self, x: torch.Tensor): + return torch.square(torch.relu(x)) + + +class PerceiverAttentionBlock(nn.Module): + def __init__( + self, d_model: int, n_heads: int, time_embedding_dim: Optional[int] = None + ): + super().__init__() + self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True) + + self.mlp = nn.Sequential( + OrderedDict( + [ + ("c_fc", nn.Linear(d_model, d_model * 4)), + ("sq_relu", SquaredReLU()), + ("c_proj", nn.Linear(d_model * 4, d_model)), + ] + ) + ) + + self.ln_1 = AdaLayerNorm(d_model, time_embedding_dim) + self.ln_2 = AdaLayerNorm(d_model, time_embedding_dim) + self.ln_ff = AdaLayerNorm(d_model, time_embedding_dim) + + def attention(self, q: torch.Tensor, kv: torch.Tensor): + attn_output, attn_output_weights = self.attn(q, kv, kv, need_weights=False) + return attn_output + + def forward( + self, + x: torch.Tensor, + latents: torch.Tensor, + timestep_embedding: torch.Tensor = None, + ): + normed_latents = self.ln_1(latents, timestep_embedding) + latents = latents + self.attention( + q=normed_latents, + kv=torch.cat([normed_latents, self.ln_2(x, timestep_embedding)], dim=1), + ) + latents = latents + self.mlp(self.ln_ff(latents, timestep_embedding)) + return latents + + +class PerceiverResampler(nn.Module): + def __init__( + self, + width: int = 768, + layers: int = 6, + heads: int = 8, + num_latents: int = 64, + output_dim=None, + input_dim=None, + time_embedding_dim: Optional[int] = None, + ): + super().__init__() + self.output_dim = output_dim + self.input_dim = input_dim + self.latents = nn.Parameter(width**-0.5 * torch.randn(num_latents, width)) + self.time_aware_linear = nn.Linear( + time_embedding_dim or width, width, bias=True + ) + + if self.input_dim is not None: + self.proj_in = nn.Linear(input_dim, width) + + self.perceiver_blocks = nn.Sequential( + *[ + PerceiverAttentionBlock( + width, heads, time_embedding_dim=time_embedding_dim + ) + for _ in range(layers) + ] + ) + + if self.output_dim is not None: + self.proj_out = nn.Sequential( + nn.Linear(width, output_dim), nn.LayerNorm(output_dim) + ) + + def forward(self, x: torch.Tensor, timestep_embedding: torch.Tensor = None): + learnable_latents = self.latents.unsqueeze(dim=0).repeat(len(x), 1, 1) + latents = learnable_latents + self.time_aware_linear( + torch.nn.functional.silu(timestep_embedding) + ) + if self.input_dim is not None: + x = self.proj_in(x) + for p_block in self.perceiver_blocks: + latents = p_block(x, latents, timestep_embedding=timestep_embedding) + + if self.output_dim is not None: + latents = self.proj_out(latents) + + return latents + + +class T5TextEmbedder(nn.Module): + def __init__(self, pretrained_path="ybelkada/flan-t5-xl-sharded-bf16", max_length=None): + super().__init__() + self.model = T5EncoderModel.from_pretrained(pretrained_path) + self.tokenizer = T5Tokenizer.from_pretrained(pretrained_path) + self.max_length = max_length + + def forward( + self, caption, text_input_ids=None, attention_mask=None, max_length=None + ): + if max_length is None: + max_length = self.max_length + + if text_input_ids is None or attention_mask is None: + if max_length is not None: + text_inputs = self.tokenizer( + caption, + return_tensors="pt", + add_special_tokens=True, + max_length=max_length, + padding="max_length", + truncation=True, + ) + else: + text_inputs = self.tokenizer( + caption, return_tensors="pt", add_special_tokens=True + ) + text_input_ids = text_inputs.input_ids + attention_mask = text_inputs.attention_mask + text_input_ids = text_input_ids.to(self.model.device) + attention_mask = attention_mask.to(self.model.device) + outputs = self.model(text_input_ids, attention_mask=attention_mask) + + embeddings = outputs.last_hidden_state + return embeddings + + +class ELLA(nn.Module): + def __init__( + self, + time_channel=320, + time_embed_dim=768, + act_fn: str = "silu", + out_dim: Optional[int] = None, + width=768, + layers=6, + heads=8, + num_latents=64, + input_dim=2048, + ): + super().__init__() + + self.position = Timesteps( + time_channel, flip_sin_to_cos=True, downscale_freq_shift=0 + ) + self.time_embedding = TimestepEmbedding( + in_channels=time_channel, + time_embed_dim=time_embed_dim, + act_fn=act_fn, + out_dim=out_dim, + ) + + self.connector = PerceiverResampler( + width=width, + layers=layers, + heads=heads, + num_latents=num_latents, + input_dim=input_dim, + time_embedding_dim=time_embed_dim, + ) + + def forward(self, text_encode_features, timesteps): + device = text_encode_features.device + dtype = text_encode_features.dtype + + ori_time_feature = self.position(timesteps.view(-1)).to(device, dtype=dtype) + ori_time_feature = ( + ori_time_feature.unsqueeze(dim=1) + if ori_time_feature.ndim == 2 + else ori_time_feature + ) + ori_time_feature = ori_time_feature.expand(len(text_encode_features), -1, -1) + time_embedding = self.time_embedding(ori_time_feature) + + encoder_hidden_states = self.connector( + text_encode_features, timestep_embedding=time_embedding + ) + + return encoder_hidden_states diff --git a/examples/brushnet_example_incremental_outpainting.json b/examples/brushnet_example_incremental_outpainting.json index bf1aff4..88dfa4f 100644 --- a/examples/brushnet_example_incremental_outpainting.json +++ b/examples/brushnet_example_incremental_outpainting.json @@ -1,6 +1,6 @@ { - "last_node_id": 74, - "last_link_id": 155, + "last_node_id": 80, + "last_link_id": 175, "nodes": [ { "id": 24, @@ -14,7 +14,7 @@ "1": 218 }, "flags": {}, - "order": 4, + "order": 3, "mode": 0, "inputs": [ { @@ -59,31 +59,6 @@ 2 ] }, - { - "id": 54, - "type": "MaskPreview+", - "pos": [ - 1295, - 764 - ], - "size": { - "0": 315.012451171875, - "1": 337.9891662597656 - }, - "flags": {}, - "order": 8, - "mode": 0, - "inputs": [ - { - "name": "mask", - "type": "MASK", - "link": 135 - } - ], - "properties": { - "Node name for S&R": "MaskPreview+" - } - }, { "id": 55, "type": "ImageRemoveBackground+", @@ -133,6 +108,465 @@ "Node name for S&R": "ImageRemoveBackground+" } }, + { + "id": 56, + "type": "RemBGSession+", + "pos": [ + 436, + 970 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "REMBG_SESSION", + "type": "REMBG_SESSION", + "links": [ + 110 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "RemBGSession+" + }, + "widgets_values": [ + "u2net_human_seg: human segmentation", + "CPU" + ] + }, + { + "id": 73, + "type": "ImagePadForOutpaintMasked", + "pos": [ + 973, + 1563 + ], + "size": { + "0": 315, + "1": 174 + }, + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 165 + }, + { + "name": "mask", + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 167 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 168 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "ImagePadForOutpaintMasked" + }, + "widgets_values": [ + 0, + 0, + 0, + 256, + 0 + ] + }, + { + "id": 66, + "type": "ImagePadForOutpaintMasked", + "pos": [ + 806, + 584 + ], + "size": { + "0": 315, + "1": 174 + }, + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 130 + }, + { + "name": "mask", + "type": "MASK", + "link": 143 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 157 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 135, + 158 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "ImagePadForOutpaintMasked" + }, + "widgets_values": [ + 128, + 128, + 128, + 0, + 0 + ] + }, + { + "id": 7, + "type": "LoadImage", + "pos": [ + 1631, + 173 + ], + "size": { + "0": 316, + "1": 405 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 39 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "oldman.jpg", + "image" + ] + }, + { + "id": 54, + "type": "MaskPreview+", + "pos": [ + 1293, + 729 + ], + "size": { + "0": 315.012451171875, + "1": 337.9891662597656 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 135 + } + ], + "properties": { + "Node name for S&R": "MaskPreview+" + } + }, + { + "id": 70, + "type": "ImagePadForOutpaintMasked", + "pos": [ + 857, + 1122 + ], + "size": { + "0": 315, + "1": 174 + }, + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 160 + }, + { + "name": "mask", + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 162 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 163 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "ImagePadForOutpaintMasked" + }, + "widgets_values": [ + 0, + 0, + 0, + 256, + 0 + ] + }, + { + "id": 47, + "type": "PreviewImage", + "pos": [ + 1972, + 197 + ], + "size": { + "0": 598.1146240234375, + "1": 551.4878540039062 + }, + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 159, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 71, + "type": "PreviewImage", + "pos": [ + 2591, + 187 + ], + "size": { + "0": 704.5072021484375, + "1": 839.8850708007812 + }, + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 164, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 74, + "type": "PreviewImage", + "pos": [ + 3320, + 187 + ], + "size": { + "0": 786.09912109375, + "1": 1164.818603515625 + }, + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 169, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 75, + "type": "brushnet_sampler", + "pos": [ + 1208, + 145 + ], + "size": [ + 393.54825627441437, + 477.48640561035177 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "link": 156 + }, + { + "name": "image", + "type": "IMAGE", + "link": 157 + }, + { + "name": "mask", + "type": "MASK", + "link": 158 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 159, + 160 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "brushnet_sampler" + }, + "widgets_values": [ + 25, + 7.5, + 1, + 0, + 1, + false, + 0, + 127, + "fixed", + "UniPCMultistepScheduler", + "seaside, old man, scenery", + "bad quality" + ] + }, + { + "id": 1, + "type": "brushnet_model_loader", + "pos": [ + 688, + 402 + ], + "size": [ + 378.24762954139715, + 98.14945294189465 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 1, + "slot_index": 0 + }, + { + "name": "clip", + "type": "CLIP", + "link": 2 + }, + { + "name": "vae", + "type": "VAE", + "link": 3 + } + ], + "outputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "links": [ + 156, + 161, + 166 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "brushnet_model_loader" + }, + "widgets_values": [ + "brushnet_segmentation_mask" + ] + }, { "id": 3, "type": "CheckpointLoaderSimple", @@ -140,12 +574,12 @@ 202, 402 ], - "size": { - "0": 351.8843078613281, - "1": 98 - }, + "size": [ + 401.15305795288054, + 98 + ], "flags": {}, - "order": 0, + "order": 2, "mode": 0, "outputs": [ { @@ -179,167 +613,20 @@ "Node name for S&R": "CheckpointLoaderSimple" }, "widgets_values": [ - "1_5/epicphotogasm_x.safetensors" + "1_5\\realisticVisionV60B1_v51VAE.safetensors" ] }, { - "id": 56, - "type": "RemBGSession+", - "pos": [ - 436, - 970 - ], - "size": { - "0": 315, - "1": 82 - }, - "flags": {}, - "order": 1, - "mode": 0, - "outputs": [ - { - "name": "REMBG_SESSION", - "type": "REMBG_SESSION", - "links": [ - 110 - ], - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "RemBGSession+" - }, - "widgets_values": [ - "u2net_human_seg: human segmentation", - "CPU" - ] - }, - { - "id": 66, - "type": "ImagePadForOutpaintMasked", - "pos": [ - 845, - 612 - ], - "size": { - "0": 315, - "1": 174 - }, - "flags": {}, - "order": 6, - "mode": 0, - "inputs": [ - { - "name": "image", - "type": "IMAGE", - "link": 130 - }, - { - "name": "mask", - "type": "MASK", - "link": 143 - } - ], - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 131 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [ - 132, - 135 - ], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "ImagePadForOutpaintMasked" - }, - "widgets_values": [ - 128, - 128, - 128, - 0, - 0 - ] - }, - { - "id": 70, - "type": "ImagePadForOutpaintMasked", - "pos": [ - 875, - 1181 - ], - "size": { - "0": 315, - "1": 174 - }, - "flags": {}, - "order": 10, - "mode": 0, - "inputs": [ - { - "name": "image", - "type": "IMAGE", - "link": 147 - }, - { - "name": "mask", - "type": "MASK", - "link": null - } - ], - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 148 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [ - 149 - ], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "ImagePadForOutpaintMasked" - }, - "widgets_values": [ - 0, - 0, - 0, - 256, - 0 - ] - }, - { - "id": 69, + "id": 76, "type": "brushnet_sampler", "pos": [ - 1299, - 1175 + 1371, + 1159 + ], + "size": [ + 393.5482482910156, + 477.4864196777344 ], - "size": { - "0": 399, - "1": 299 - }, "flags": {}, "order": 11, "mode": 0, @@ -347,19 +634,17 @@ { "name": "brushnet", "type": "BRUSHNET", - "link": 144, - "slot_index": 0 + "link": 161 }, { "name": "image", "type": "IMAGE", - "link": 148, - "slot_index": 1 + "link": 162 }, { "name": "mask", "type": "MASK", - "link": 149 + "link": 163 } ], "outputs": [ @@ -367,171 +652,41 @@ "name": "images", "type": "IMAGE", "links": [ - 150, - 151 + 164, + 165 ], - "shape": 3, - "slot_index": 0 + "shape": 3 } ], "properties": { "Node name for S&R": "brushnet_sampler" }, "widgets_values": [ - 32, + 25, + 7.5, 1, - 36, + 0, + 1, + false, + 0, + 137, "fixed", - "DPMSolverMultistepScheduler_SDE_karras", - "old man, jacket" + "UniPCMultistepScheduler", + "closeup, jacket, lower body, arms crossed", + "bad quality" ] }, { - "id": 73, - "type": "ImagePadForOutpaintMasked", - "pos": [ - 973, - 1563 - ], - "size": { - "0": 315, - "1": 174 - }, - "flags": {}, - "order": 13, - "mode": 0, - "inputs": [ - { - "name": "image", - "type": "IMAGE", - "link": 151 - }, - { - "name": "mask", - "type": "MASK", - "link": null - } - ], - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 152 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [ - 155 - ], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "ImagePadForOutpaintMasked" - }, - "widgets_values": [ - 0, - 0, - 0, - 256, - 0 - ] - }, - { - "id": 1, - "type": "brushnet_model_loader", - "pos": [ - 688, - 402 - ], - "size": { - "0": 337, - "1": 98 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "MODEL", - "link": 1, - "slot_index": 0 - }, - { - "name": "clip", - "type": "CLIP", - "link": 2 - }, - { - "name": "vae", - "type": "VAE", - "link": 3 - } - ], - "outputs": [ - { - "name": "brushnet", - "type": "BRUSHNET", - "links": [ - 4, - 144, - 154 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "brushnet_model_loader" - }, - "widgets_values": [ - "brushnet_segmentation_mask" - ] - }, - { - "id": 71, - "type": "PreviewImage", - "pos": [ - 2600, - 230 - ], - "size": [ - 704.5071846093756, - 839.885078935547 - ], - "flags": {}, - "order": 12, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 150, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, - { - "id": 72, + "id": 77, "type": "brushnet_sampler", "pos": [ - 1322, - 1531 + 1379, + 1690 + ], + "size": [ + 393.5482482910156, + 477.4864196777344 ], - "size": { - "0": 399, - "1": 299 - }, "flags": {}, "order": 14, "mode": 0, @@ -539,20 +694,17 @@ { "name": "brushnet", "type": "BRUSHNET", - "link": 154, - "slot_index": 0 + "link": 166 }, { "name": "image", "type": "IMAGE", - "link": 152, - "slot_index": 1 + "link": 167 }, { "name": "mask", "type": "MASK", - "link": 155, - "slot_index": 2 + "link": 168 } ], "outputs": [ @@ -560,171 +712,27 @@ "name": "images", "type": "IMAGE", "links": [ - 153 + 169 ], - "shape": 3, - "slot_index": 0 + "shape": 3 } ], "properties": { "Node name for S&R": "brushnet_sampler" }, "widgets_values": [ - 32, + 25, + 7.5, 1, - 37, - "fixed", - "DPMSolverMultistepScheduler_SDE_karras", - "old man, pants" - ] - }, - { - "id": 5, - "type": "brushnet_sampler", - "pos": [ - 1153, - 386 - ], - "size": { - "0": 399, - "1": 299 - }, - "flags": {}, - "order": 7, - "mode": 0, - "inputs": [ - { - "name": "brushnet", - "type": "BRUSHNET", - "link": 4, - "slot_index": 0 - }, - { - "name": "image", - "type": "IMAGE", - "link": 131, - "slot_index": 1 - }, - { - "name": "mask", - "type": "MASK", - "link": 132 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 86, - 147 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "brushnet_sampler" - }, - "widgets_values": [ - 32, + 0, 1, - 35, + false, + 0, + 128, "fixed", - "DPMSolverMultistepScheduler_SDE_karras", - "seaside, landscape, old man, jacket" - ] - }, - { - "id": 74, - "type": "PreviewImage", - "pos": [ - 3320, - 230 - ], - "size": [ - 786.0991246093758, - 1164.8186089355472 - ], - "flags": {}, - "order": 15, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 153, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, - { - "id": 47, - "type": "PreviewImage", - "pos": [ - 1970, - 240 - ], - "size": [ - 598.1146253564457, - 551.4878638635256 - ], - "flags": {}, - "order": 9, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 86, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, - { - "id": 7, - "type": "LoadImage", - "pos": [ - 1630, - 190 - ], - "size": { - "0": 316, - "1": 405 - }, - "flags": {}, - "order": 2, - "mode": 0, - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 39 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "LoadImage" - }, - "widgets_values": [ - "ComfyUI_temp_cytrk_00031_ (2).png", - "image" + "UniPCMultistepScheduler", + "old man, chair", + "bad quality" ] } ], @@ -753,14 +761,6 @@ 2, "VAE" ], - [ - 4, - 1, - 0, - 5, - 0, - "BRUSHNET" - ], [ 39, 7, @@ -769,14 +769,6 @@ 0, "IMAGE" ], - [ - 86, - 5, - 0, - 47, - 0, - "IMAGE" - ], [ 110, 56, @@ -801,22 +793,6 @@ 0, "IMAGE" ], - [ - 131, - 66, - 0, - 5, - 1, - "IMAGE" - ], - [ - 132, - 66, - 1, - 5, - 2, - "MASK" - ], [ 135, 66, @@ -834,84 +810,116 @@ "MASK" ], [ - 144, + 156, 1, 0, - 69, + 75, 0, "BRUSHNET" ], [ - 147, - 5, + 157, + 66, 0, - 70, - 0, - "IMAGE" - ], - [ - 148, - 70, - 0, - 69, + 75, 1, "IMAGE" ], [ - 149, - 70, + 158, + 66, 1, - 69, + 75, 2, "MASK" ], [ - 150, - 69, + 159, + 75, + 0, + 47, + 0, + "IMAGE" + ], + [ + 160, + 75, + 0, + 70, + 0, + "IMAGE" + ], + [ + 161, + 1, + 0, + 76, + 0, + "BRUSHNET" + ], + [ + 162, + 70, + 0, + 76, + 1, + "IMAGE" + ], + [ + 163, + 70, + 1, + 76, + 2, + "MASK" + ], + [ + 164, + 76, 0, 71, 0, "IMAGE" ], [ - 151, - 69, + 165, + 76, 0, 73, 0, "IMAGE" ], [ - 152, - 73, - 0, - 72, - 1, - "IMAGE" - ], - [ - 153, - 72, - 0, - 74, - 0, - "IMAGE" - ], - [ - 154, + 166, 1, 0, - 72, + 77, 0, "BRUSHNET" ], [ - 155, + 167, + 73, + 0, + 77, + 1, + "IMAGE" + ], + [ + 168, 73, 1, - 72, + 77, 2, "MASK" + ], + [ + 169, + 77, + 0, + 74, + 0, + "IMAGE" ] ], "groups": [], diff --git a/examples/brushnet_example_outpaint_rembg.json b/examples/brushnet_example_inpaint_ella.json similarity index 51% rename from examples/brushnet_example_outpaint_rembg.json rename to examples/brushnet_example_inpaint_ella.json index 6ef920d..cdbe2eb 100644 --- a/examples/brushnet_example_outpaint_rembg.json +++ b/examples/brushnet_example_inpaint_ella.json @@ -1,7 +1,377 @@ { - "last_node_id": 68, - "last_link_id": 143, + "last_node_id": 60, + "last_link_id": 129, "nodes": [ + { + "id": 50, + "type": "MaskToImage", + "pos": [ + 1251, + 1020 + ], + "size": { + "0": 210, + "1": 26 + }, + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 96 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 92 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "MaskToImage" + } + }, + { + "id": 49, + "type": "ImageCompositeMasked", + "pos": [ + 1251, + 1093 + ], + "size": { + "0": 315, + "1": 146 + }, + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "destination", + "type": "IMAGE", + "link": 90 + }, + { + "name": "source", + "type": "IMAGE", + "link": 92 + }, + { + "name": "mask", + "type": "MASK", + "link": 89, + "slot_index": 2 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 93 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ImageCompositeMasked" + }, + "widgets_values": [ + 0, + 0, + true + ] + }, + { + "id": 46, + "type": "PreviewImage", + "pos": [ + 2124, + 794 + ], + "size": { + "0": 491.6383056640625, + "1": 480.1783142089844 + }, + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 85, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 6, + "type": "PreviewImage", + "pos": [ + 1593, + 795 + ], + "size": { + "0": 491.6383056640625, + "1": 480.1783142089844 + }, + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 93, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 44, + "type": "ImageCompositeMasked", + "pos": [ + 1722, + 568 + ], + "size": { + "0": 315, + "1": 146 + }, + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "destination", + "type": "IMAGE", + "link": 81 + }, + { + "name": "source", + "type": "IMAGE", + "link": 108 + }, + { + "name": "mask", + "type": "MASK", + "link": 95, + "slot_index": 2 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 85 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ImageCompositeMasked" + }, + "widgets_values": [ + 0, + 0, + true + ] + }, + { + "id": 48, + "type": "MaskPreview+", + "pos": [ + 987, + 991 + ], + "size": { + "0": 210, + "1": 246 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 88 + } + ], + "properties": { + "Node name for S&R": "MaskPreview+" + } + }, + { + "id": 47, + "type": "PreviewImage", + "pos": [ + 1711, + 147 + ], + "size": { + "0": 361.5766296386719, + "1": 361.18719482421875 + }, + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 109, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 51, + "type": "RemapMaskRange", + "pos": [ + 1240, + 885 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 94 + } + ], + "outputs": [ + { + "name": "mask", + "type": "MASK", + "links": [ + 95, + 96 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "RemapMaskRange" + }, + "widgets_values": [ + 0, + 0.9500000000000001 + ] + }, + { + "id": 1, + "type": "brushnet_model_loader", + "pos": [ + 688, + 402 + ], + "size": { + "0": 337, + "1": 98 + }, + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MODEL", + "link": 1, + "slot_index": 0 + }, + { + "name": "clip", + "type": "CLIP", + "link": 2 + }, + { + "name": "vae", + "type": "VAE", + "link": 3 + } + ], + "outputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "links": [ + 104 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "brushnet_model_loader" + }, + "widgets_values": [ + "brushnet_segmentation_mask" + ] + }, + { + "id": 53, + "type": "brushnet_ella_loader", + "pos": [ + 941, + 312 + ], + "size": { + "0": 216.59999084472656, + "1": 26 + }, + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "link": 104 + } + ], + "outputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "links": [ + 110 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "brushnet_ella_loader" + } + }, { "id": 3, "type": "CheckpointLoaderSimple", @@ -48,131 +418,15 @@ "Node name for S&R": "CheckpointLoaderSimple" }, "widgets_values": [ - "1_5/darkSushi25D25D_v40.safetensors" - ] - }, - { - "id": 1, - "type": "brushnet_model_loader", - "pos": [ - 688, - 402 - ], - "size": { - "0": 337, - "1": 98 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "MODEL", - "link": 1, - "slot_index": 0 - }, - { - "name": "clip", - "type": "CLIP", - "link": 2 - }, - { - "name": "vae", - "type": "VAE", - "link": 3 - } - ], - "outputs": [ - { - "name": "brushnet", - "type": "BRUSHNET", - "links": [ - 4 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "brushnet_model_loader" - }, - "widgets_values": [ - "brushnet_segmentation_mask" - ] - }, - { - "id": 47, - "type": "PreviewImage", - "pos": [ - 1728, - 406 - ], - "size": { - "0": 555.6796875, - "1": 582.3743896484375 - }, - "flags": {}, - "order": 9, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 86, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, - { - "id": 7, - "type": "LoadImage", - "pos": [ - 28, - 570 - ], - "size": { - "0": 316, - "1": 405 - }, - "flags": {}, - "order": 1, - "mode": 0, - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 39 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "LoadImage" - }, - "widgets_values": [ - "clipspace/clipspace-mask-1261422.png [input]", - "image" + "1_5\\realisticVisionV60B1_v51VAE.safetensors" ] }, { "id": 24, "type": "ImageResize+", "pos": [ - 453, - 569 + 578, + 591 ], "size": { "0": 315, @@ -193,8 +447,9 @@ "name": "IMAGE", "type": "IMAGE", "links": [ - 111, - 130 + 81, + 90, + 116 ], "shape": 3, "slot_index": 0 @@ -225,37 +480,79 @@ ] }, { - "id": 66, - "type": "ImagePadForOutpaintMasked", + "id": 45, + "type": "GrowMaskWithBlur", "pos": [ - 845, - 612 + 644, + 948 ], "size": { "0": 315, - "1": 174 + "1": 246 }, "flags": {}, - "order": 6, + "order": 5, "mode": 0, "inputs": [ { - "name": "image", - "type": "IMAGE", - "link": 130 - }, + "name": "mask", + "type": "MASK", + "link": 129 + } + ], + "outputs": [ { "name": "mask", "type": "MASK", - "link": 143 + "links": [ + 88, + 89, + 94 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "mask_inverted", + "type": "MASK", + "links": null, + "shape": 3 } ], + "properties": { + "Node name for S&R": "GrowMaskWithBlur" + }, + "widgets_values": [ + 2, + 0, + true, + false, + 25.5, + 1, + 1, + false + ] + }, + { + "id": 7, + "type": "LoadImage", + "pos": [ + 215, + 582 + ], + "size": { + "0": 316, + "1": 405 + }, + "flags": {}, + "order": 1, + "mode": 0, "outputs": [ { "name": "IMAGE", "type": "IMAGE", "links": [ - 131 + 39 ], "shape": 3, "slot_index": 0 @@ -264,55 +561,56 @@ "name": "MASK", "type": "MASK", "links": [ - 132, - 135 + 128, + 129 ], "shape": 3, "slot_index": 1 } ], "properties": { - "Node name for S&R": "ImagePadForOutpaintMasked" + "Node name for S&R": "LoadImage" }, "widgets_values": [ - 128, - 128, - 128, - 128, - 0 + "clipspace/clipspace-mask-6229541.png [input]", + "image" ] }, { - "id": 5, - "type": "brushnet_sampler", + "id": 54, + "type": "brushnet_sampler_ella", "pos": [ - 1283, - 404 + 1250, + 278 ], "size": { - "0": 399, - "1": 299 + "0": 348.9030456542969, + "1": 395.87213134765625 }, "flags": {}, - "order": 7, + "order": 9, "mode": 0, "inputs": [ { "name": "brushnet", "type": "BRUSHNET", - "link": 4, - "slot_index": 0 + "link": 110 + }, + { + "name": "ella_embeds", + "type": "ELLAEMBEDS", + "link": 111, + "slot_index": 1 }, { "name": "image", "type": "IMAGE", - "link": 131, - "slot_index": 1 + "link": 116 }, { "name": "mask", "type": "MASK", - "link": 132 + "link": 128 } ], "outputs": [ @@ -320,128 +618,61 @@ "name": "images", "type": "IMAGE", "links": [ - 86 + 108, + 109 ], - "shape": 3, - "slot_index": 0 + "shape": 3 } ], "properties": { - "Node name for S&R": "brushnet_sampler" + "Node name for S&R": "brushnet_sampler_ella" }, "widgets_values": [ - 30, + 25, + 7.5, 1, - 31, + 0, + 1, + false, + 0, + 30, "fixed", - "UniPCMultistepScheduler", - "1girl, blue dress, forest, best quality, masterpiece" + "UniPCMultistepScheduler" ] }, - { - "id": 54, - "type": "MaskPreview+", - "pos": [ - 1295, - 764 - ], - "size": { - "0": 315.012451171875, - "1": 337.9891662597656 - }, - "flags": {}, - "order": 8, - "mode": 0, - "inputs": [ - { - "name": "mask", - "type": "MASK", - "link": 135 - } - ], - "properties": { - "Node name for S&R": "MaskPreview+" - } - }, { "id": 55, - "type": "ImageRemoveBackground+", + "type": "ella_t5_embeds", "pos": [ - 505, - 849 + 404, + 50 ], "size": { - "0": 241.79998779296875, - "1": 46 - }, - "flags": {}, - "order": 5, - "mode": 0, - "inputs": [ - { - "name": "rembg_session", - "type": "REMBG_SESSION", - "link": 110, - "slot_index": 0 - }, - { - "name": "image", - "type": "IMAGE", - "link": 111 - } - ], - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": null, - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [ - 143 - ], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "ImageRemoveBackground+" - } - }, - { - "id": 56, - "type": "RemBGSession+", - "pos": [ - 436, - 970 - ], - "size": { - "0": 315, - "1": 82 + "0": 485.3834228515625, + "1": 248.846435546875 }, "flags": {}, "order": 2, "mode": 0, "outputs": [ { - "name": "REMBG_SESSION", - "type": "REMBG_SESSION", + "name": "ella_embeds", + "type": "ELLAEMBEDS", "links": [ - 110 + 111 ], "shape": 3 } ], "properties": { - "Node name for S&R": "RemBGSession+" + "Node name for S&R": "ella_t5_embeds" }, "widgets_values": [ - "isnet-anime: anime illustrations", - "CPU" + "painting of a woman waving with her left hand raised up", + 1, + 128, + false, + true ] } ], @@ -470,14 +701,6 @@ 2, "VAE" ], - [ - 4, - 1, - 0, - 5, - 0, - "BRUSHNET" - ], [ 39, 7, @@ -487,8 +710,104 @@ "IMAGE" ], [ - 86, - 5, + 81, + 24, + 0, + 44, + 0, + "IMAGE" + ], + [ + 85, + 44, + 0, + 46, + 0, + "IMAGE" + ], + [ + 88, + 45, + 0, + 48, + 0, + "MASK" + ], + [ + 89, + 45, + 0, + 49, + 2, + "MASK" + ], + [ + 90, + 24, + 0, + 49, + 0, + "IMAGE" + ], + [ + 92, + 50, + 0, + 49, + 1, + "IMAGE" + ], + [ + 93, + 49, + 0, + 6, + 0, + "IMAGE" + ], + [ + 94, + 45, + 0, + 51, + 0, + "MASK" + ], + [ + 95, + 51, + 0, + 44, + 2, + "MASK" + ], + [ + 96, + 51, + 0, + 50, + 0, + "MASK" + ], + [ + 104, + 1, + 0, + 53, + 0, + "BRUSHNET" + ], + [ + 108, + 54, + 0, + 44, + 1, + "IMAGE" + ], + [ + 109, + 54, 0, 47, 0, @@ -496,58 +815,42 @@ ], [ 110, - 56, + 53, 0, - 55, + 54, 0, - "REMBG_SESSION" + "BRUSHNET" ], [ 111, - 24, - 0, 55, + 0, + 54, 1, - "IMAGE" + "ELLAEMBEDS" ], [ - 130, + 116, 24, 0, - 66, - 0, - "IMAGE" - ], - [ - 131, - 66, - 0, - 5, - 1, - "IMAGE" - ], - [ - 132, - 66, - 1, - 5, + 54, 2, - "MASK" + "IMAGE" ], [ - 135, - 66, + 128, + 7, 1, 54, - 0, + 3, "MASK" ], [ - 143, - 55, - 1, - 66, + 129, + 7, 1, + 45, + 0, "MASK" ] ], diff --git a/examples/brushnet_example_outpaint.json b/examples/brushnet_example_outpaint.json index e8c762d..3e97d8f 100644 --- a/examples/brushnet_example_outpaint.json +++ b/examples/brushnet_example_outpaint.json @@ -1,7 +1,188 @@ { - "last_node_id": 54, - "last_link_id": 109, + "last_node_id": 69, + "last_link_id": 147, "nodes": [ + { + "id": 47, + "type": "PreviewImage", + "pos": [ + 1728, + 406 + ], + "size": { + "0": 555.6796875, + "1": 582.3743896484375 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 147, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PreviewImage" + } + }, + { + "id": 24, + "type": "ImageResize+", + "pos": [ + 453, + 569 + ], + "size": { + "0": 315, + "1": 218 + }, + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 39 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 130 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResize+" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + "always", + 2 + ] + }, + { + "id": 7, + "type": "LoadImage", + "pos": [ + 28, + 570 + ], + "size": { + "0": 316, + "1": 405 + }, + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 39 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "lighthouse (5).png", + "image" + ] + }, + { + "id": 69, + "type": "brushnet_sampler", + "pos": [ + 1281, + 400 + ], + "size": [ + 391.887911987302, + 422.8043090820303 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "brushnet", + "type": "BRUSHNET", + "link": 144 + }, + { + "name": "image", + "type": "IMAGE", + "link": 145 + }, + { + "name": "mask", + "type": "MASK", + "link": 146 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 147 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "brushnet_sampler" + }, + "widgets_values": [ + 25, + 7.5, + 1, + 0, + 1, + false, + 0, + 178746569802073, + "randomize", + "UniPCMultistepScheduler", + "miniature, lighthouse", + "bad quality" + ] + }, { "id": 3, "type": "CheckpointLoaderSimple", @@ -14,7 +195,7 @@ "1": 98 }, "flags": {}, - "order": 0, + "order": 1, "mode": 0, "outputs": [ { @@ -48,47 +229,7 @@ "Node name for S&R": "CheckpointLoaderSimple" }, "widgets_values": [ - "1_5/darkSushi25D25D_v40.safetensors" - ] - }, - { - "id": 7, - "type": "LoadImage", - "pos": [ - 215, - 582 - ], - "size": { - "0": 316, - "1": 405 - }, - "flags": {}, - "order": 1, - "mode": 0, - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 39 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "LoadImage" - }, - "widgets_values": [ - "clipspace/clipspace-mask-1261422.png [input]", - "image" + "1_5\\photon_v1.safetensors" ] }, { @@ -98,12 +239,12 @@ 688, 402 ], - "size": { - "0": 337, - "1": 98 - }, + "size": [ + 390.7779070281963, + 98 + ], "flags": {}, - "order": 2, + "order": 3, "mode": 0, "inputs": [ { @@ -128,7 +269,7 @@ "name": "brushnet", "type": "BRUSHNET", "links": [ - 4 + 144 ], "shape": 3, "slot_index": 0 @@ -142,67 +283,11 @@ ] }, { - "id": 24, - "type": "ImageResize+", + "id": 66, + "type": "ImagePadForOutpaintMasked", "pos": [ - 578, - 591 - ], - "size": { - "0": 315, - "1": 218 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [ - { - "name": "image", - "type": "IMAGE", - "link": 39 - } - ], - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 99 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "width", - "type": "INT", - "links": null, - "shape": 3 - }, - { - "name": "height", - "type": "INT", - "links": null, - "shape": 3 - } - ], - "properties": { - "Node name for S&R": "ImageResize+" - }, - "widgets_values": [ - 512, - 512, - "lanczos", - true, - "always", - 2 - ] - }, - { - "id": 52, - "type": "ImagePadForOutpaint", - "pos": [ - 918, - 598 + 845, + 612 ], "size": { "0": 315, @@ -215,7 +300,12 @@ { "name": "image", "type": "IMAGE", - "link": 99 + "link": 130 + }, + { + "name": "mask", + "type": "MASK", + "link": null } ], "outputs": [ @@ -223,8 +313,7 @@ "name": "IMAGE", "type": "IMAGE", "links": [ - 100, - 108 + 145 ], "shape": 3, "slot_index": 0 @@ -233,155 +322,48 @@ "name": "MASK", "type": "MASK", "links": [ - 101, - 109 + 135, + 146 ], "shape": 3, "slot_index": 1 } ], "properties": { - "Node name for S&R": "ImagePadForOutpaint" + "Node name for S&R": "ImagePadForOutpaintMasked" }, "widgets_values": [ - 256, - 256, - 256, - 256, - 1 + 128, + 128, + 128, + 0, + 2 ] }, { - "id": 53, - "type": "PreviewImage", + "id": 54, + "type": "MaskPreview+", "pos": [ - 928, - 824 + 871, + 845 ], "size": [ - 210, - 246 + 334.2839078979473, + 295.2053316650374 ], "flags": {}, "order": 5, "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 108 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } - }, - { - "id": 54, - "type": "MaskPreview+", - "pos": [ - 1148, - 826 - ], - "size": [ - 210, - 246 - ], - "flags": {}, - "order": 7, - "mode": 0, "inputs": [ { "name": "mask", "type": "MASK", - "link": 109 + "link": 135 } ], "properties": { "Node name for S&R": "MaskPreview+" } - }, - { - "id": 5, - "type": "brushnet_sampler", - "pos": [ - 1283, - 404 - ], - "size": { - "0": 399, - "1": 299 - }, - "flags": {}, - "order": 6, - "mode": 0, - "inputs": [ - { - "name": "brushnet", - "type": "BRUSHNET", - "link": 4, - "slot_index": 0 - }, - { - "name": "image", - "type": "IMAGE", - "link": 100, - "slot_index": 1 - }, - { - "name": "mask", - "type": "MASK", - "link": 101 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 86 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "brushnet_sampler" - }, - "widgets_values": [ - 30, - 1, - 29, - "fixed", - "UniPCMultistepScheduler", - "1girl, blue dress, forest, best quality, masterpiece" - ] - }, - { - "id": 47, - "type": "PreviewImage", - "pos": [ - 1728, - 406 - ], - "size": [ - 555.6796875, - 582.3743743896484 - ], - "flags": {}, - "order": 8, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 86, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PreviewImage" - } } ], "links": [ @@ -409,14 +391,6 @@ 2, "VAE" ], - [ - 4, - 1, - 0, - 5, - 0, - "BRUSHNET" - ], [ 39, 7, @@ -426,52 +400,52 @@ "IMAGE" ], [ - 86, - 5, - 0, - 47, - 0, - "IMAGE" - ], - [ - 99, + 130, 24, 0, - 52, + 66, 0, "IMAGE" ], [ - 100, - 52, - 0, - 5, - 1, - "IMAGE" - ], - [ - 101, - 52, - 1, - 5, - 2, - "MASK" - ], - [ - 108, - 52, - 0, - 53, - 0, - "IMAGE" - ], - [ - 109, - 52, + 135, + 66, 1, 54, 0, "MASK" + ], + [ + 144, + 1, + 0, + 69, + 0, + "BRUSHNET" + ], + [ + 145, + 66, + 0, + 69, + 1, + "IMAGE" + ], + [ + 146, + 66, + 1, + 69, + 2, + "MASK" + ], + [ + 147, + 69, + 0, + 47, + 0, + "IMAGE" ] ], "groups": [], diff --git a/examples/brushnet_example_workflow_blend.json b/examples/brushnet_example_workflow_blend.json index 932f6a6..cf0108f 100644 --- a/examples/brushnet_example_workflow_blend.json +++ b/examples/brushnet_example_workflow_blend.json @@ -1,56 +1,7 @@ { - "last_node_id": 51, - "last_link_id": 98, + "last_node_id": 52, + "last_link_id": 103, "nodes": [ - { - "id": 3, - "type": "CheckpointLoaderSimple", - "pos": [ - 202, - 402 - ], - "size": { - "0": 351.8843078613281, - "1": 98 - }, - "flags": {}, - "order": 0, - "mode": 0, - "outputs": [ - { - "name": "MODEL", - "type": "MODEL", - "links": [ - 1 - ], - "shape": 3 - }, - { - "name": "CLIP", - "type": "CLIP", - "links": [ - 2 - ], - "shape": 3, - "slot_index": 1 - }, - { - "name": "VAE", - "type": "VAE", - "links": [ - 3 - ], - "shape": 3, - "slot_index": 2 - } - ], - "properties": { - "Node name for S&R": "CheckpointLoaderSimple" - }, - "widgets_values": [ - "1_5/darkSushi25D25D_v40.safetensors" - ] - }, { "id": 24, "type": "ImageResize+", @@ -63,7 +14,7 @@ "1": 218 }, "flags": {}, - "order": 3, + "order": 2, "mode": 0, "inputs": [ { @@ -78,8 +29,8 @@ "type": "IMAGE", "links": [ 81, - 87, - 90 + 90, + 100 ], "shape": 3, "slot_index": 0 @@ -109,47 +60,6 @@ 2 ] }, - { - "id": 51, - "type": "RemapMaskRange", - "pos": [ - 1240, - 885 - ], - "size": { - "0": 315, - "1": 82 - }, - "flags": {}, - "order": 7, - "mode": 0, - "inputs": [ - { - "name": "mask", - "type": "MASK", - "link": 94 - } - ], - "outputs": [ - { - "name": "mask", - "type": "MASK", - "links": [ - 95, - 96 - ], - "shape": 3, - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "RemapMaskRange" - }, - "widgets_values": [ - 0, - 0.8 - ] - }, { "id": 50, "type": "MaskToImage", @@ -162,7 +72,7 @@ "1": 26 }, "flags": {}, - "order": 10, + "order": 8, "mode": 0, "inputs": [ { @@ -198,7 +108,7 @@ "1": 146 }, "flags": {}, - "order": 12, + "order": 11, "mode": 0, "inputs": [ { @@ -250,7 +160,7 @@ "1": 480.1783142089844 }, "flags": {}, - "order": 11, + "order": 12, "mode": 0, "inputs": [ { @@ -313,7 +223,7 @@ { "name": "source", "type": "IMAGE", - "link": 82 + "link": 102 }, { "name": "mask", @@ -349,12 +259,12 @@ 987, 991 ], - "size": [ - 210, - 246 - ], + "size": { + "0": 210, + "1": 246 + }, "flags": {}, - "order": 6, + "order": 5, "mode": 0, "inputs": [ { @@ -379,7 +289,7 @@ "1": 246 }, "flags": {}, - "order": 4, + "order": 3, "mode": 0, "inputs": [ { @@ -421,49 +331,6 @@ false ] }, - { - "id": 7, - "type": "LoadImage", - "pos": [ - 215, - 582 - ], - "size": { - "0": 316, - "1": 405 - }, - "flags": {}, - "order": 1, - "mode": 0, - "outputs": [ - { - "name": "IMAGE", - "type": "IMAGE", - "links": [ - 39 - ], - "shape": 3, - "slot_index": 0 - }, - { - "name": "MASK", - "type": "MASK", - "links": [ - 84, - 98 - ], - "shape": 3, - "slot_index": 1 - } - ], - "properties": { - "Node name for S&R": "LoadImage" - }, - "widgets_values": [ - "clipspace/clipspace-mask-381446.png [input]", - "image" - ] - }, { "id": 47, "type": "PreviewImage", @@ -471,18 +338,18 @@ 1711, 147 ], - "size": [ - 361.57662353515616, - 361.1871826171875 - ], + "size": { + "0": 361.5766296386719, + "1": 361.18719482421875 + }, "flags": {}, - "order": 8, + "order": 10, "mode": 0, "inputs": [ { "name": "images", "type": "IMAGE", - "link": 86, + "link": 103, "slot_index": 0 } ], @@ -502,7 +369,7 @@ "1": 98 }, "flags": {}, - "order": 2, + "order": 4, "mode": 0, "inputs": [ { @@ -527,7 +394,7 @@ "name": "brushnet", "type": "BRUSHNET", "links": [ - 4 + 99 ], "shape": 3, "slot_index": 0 @@ -541,36 +408,167 @@ ] }, { - "id": 5, - "type": "brushnet_sampler", + "id": 7, + "type": "LoadImage", "pos": [ - 1250, - 401 + 215, + 582 ], "size": { - "0": 399, - "1": 299 + "0": 316, + "1": 405 }, "flags": {}, - "order": 5, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 39 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 84, + 101 + ], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "clipspace/clipspace-mask-5637958.png [input]", + "image" + ] + }, + { + "id": 3, + "type": "CheckpointLoaderSimple", + "pos": [ + 202, + 402 + ], + "size": { + "0": 351.8843078613281, + "1": 98 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 1 + ], + "shape": 3 + }, + { + "name": "CLIP", + "type": "CLIP", + "links": [ + 2 + ], + "shape": 3, + "slot_index": 1 + }, + { + "name": "VAE", + "type": "VAE", + "links": [ + 3 + ], + "shape": 3, + "slot_index": 2 + } + ], + "properties": { + "Node name for S&R": "CheckpointLoaderSimple" + }, + "widgets_values": [ + "1_5\\photon_v1.safetensors" + ] + }, + { + "id": 51, + "type": "RemapMaskRange", + "pos": [ + 1240, + 885 + ], + "size": { + "0": 315, + "1": 82 + }, + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "mask", + "type": "MASK", + "link": 94 + } + ], + "outputs": [ + { + "name": "mask", + "type": "MASK", + "links": [ + 95, + 96 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "RemapMaskRange" + }, + "widgets_values": [ + 0, + 0.9500000000000001 + ] + }, + { + "id": 52, + "type": "brushnet_sampler", + "pos": [ + 1206, + 272 + ], + "size": [ + 396.6837166259743, + 448.22052941894424 + ], + "flags": {}, + "order": 7, "mode": 0, "inputs": [ { "name": "brushnet", "type": "BRUSHNET", - "link": 4, - "slot_index": 0 + "link": 99 }, { "name": "image", "type": "IMAGE", - "link": 87, - "slot_index": 1 + "link": 100 }, { "name": "mask", "type": "MASK", - "link": 98 + "link": 101 } ], "outputs": [ @@ -578,23 +576,28 @@ "name": "images", "type": "IMAGE", "links": [ - 82, - 86 + 102, + 103 ], - "shape": 3, - "slot_index": 0 + "shape": 3 } ], "properties": { "Node name for S&R": "brushnet_sampler" }, "widgets_values": [ - 30, + 25, + 7.5, 1, - 27, - "fixed", + 0, + 1, + false, + 0, + 492040977673787, + "randomize", "UniPCMultistepScheduler", - "purple eye" + "sunglasses", + "bad quality" ] } ], @@ -623,14 +626,6 @@ 2, "VAE" ], - [ - 4, - 1, - 0, - 5, - 0, - "BRUSHNET" - ], [ 39, 7, @@ -647,14 +642,6 @@ 0, "IMAGE" ], - [ - 82, - 5, - 0, - 44, - 1, - "IMAGE" - ], [ 84, 7, @@ -671,22 +658,6 @@ 0, "IMAGE" ], - [ - 86, - 5, - 0, - 47, - 0, - "IMAGE" - ], - [ - 87, - 24, - 0, - 5, - 1, - "IMAGE" - ], [ 88, 45, @@ -752,12 +723,44 @@ "MASK" ], [ - 98, + 99, + 1, + 0, + 52, + 0, + "BRUSHNET" + ], + [ + 100, + 24, + 0, + 52, + 1, + "IMAGE" + ], + [ + 101, 7, 1, - 5, + 52, 2, "MASK" + ], + [ + 102, + 52, + 0, + 44, + 1, + "IMAGE" + ], + [ + 103, + 52, + 0, + 47, + 0, + "IMAGE" ] ], "groups": [], diff --git a/inference.py b/inference.py deleted file mode 100644 index 79ef24b..0000000 --- a/inference.py +++ /dev/null @@ -1,419 +0,0 @@ -from pathlib import Path -from typing import Any, Optional, Union - -import fire -import gradio as gr -import safetensors.torch -import torch -from diffusers import DPMSolverMultistepScheduler, StableDiffusionPipeline -from torchvision.utils import save_image - -from model import ELLA, T5TextEmbedder - - -class ELLAProxyUNet(torch.nn.Module): - def __init__(self, ella, unet): - super().__init__() - # In order to still use the diffusers pipeline, including various workaround - - self.ella = ella - self.unet = unet - self.config = unet.config - self.dtype = unet.dtype - self.device = unet.device - - self.flexible_max_length_workaround = None - - def forward( - self, - sample: torch.FloatTensor, - timestep: Union[torch.Tensor, float, int], - encoder_hidden_states: torch.Tensor, - class_labels: Optional[torch.Tensor] = None, - timestep_cond: Optional[torch.Tensor] = None, - attention_mask: Optional[torch.Tensor] = None, - cross_attention_kwargs: Optional[dict[str, Any]] = None, - added_cond_kwargs: Optional[dict[str, torch.Tensor]] = None, - down_block_additional_residuals: Optional[tuple[torch.Tensor]] = None, - mid_block_additional_residual: Optional[torch.Tensor] = None, - down_intrablock_additional_residuals: Optional[tuple[torch.Tensor]] = None, - encoder_attention_mask: Optional[torch.Tensor] = None, - return_dict: bool = True, - ): - if self.flexible_max_length_workaround is not None: - time_aware_encoder_hidden_state_list = [] - for i, max_length in enumerate(self.flexible_max_length_workaround): - time_aware_encoder_hidden_state_list.append( - self.ella(encoder_hidden_states[i : i + 1, :max_length], timestep) - ) - # No matter how many tokens are text features, the ella output must be 64 tokens. - time_aware_encoder_hidden_states = torch.cat( - time_aware_encoder_hidden_state_list, dim=0 - ) - else: - time_aware_encoder_hidden_states = self.ella( - encoder_hidden_states, timestep - ) - - return self.unet( - sample=sample, - timestep=timestep, - encoder_hidden_states=time_aware_encoder_hidden_states, - class_labels=class_labels, - timestep_cond=timestep_cond, - attention_mask=attention_mask, - cross_attention_kwargs=cross_attention_kwargs, - added_cond_kwargs=added_cond_kwargs, - down_block_additional_residuals=down_block_additional_residuals, - mid_block_additional_residual=mid_block_additional_residual, - down_intrablock_additional_residuals=down_intrablock_additional_residuals, - encoder_attention_mask=encoder_attention_mask, - return_dict=return_dict, - ) - - -def generate_image_with_flexible_max_length( - pipe, t5_encoder, prompt, fixed_negative=False, output_type="pt", **pipe_kwargs -): - device = pipe.device - dtype = pipe.dtype - prompt = [prompt] if isinstance(prompt, str) else prompt - batch_size = len(prompt) - - prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype) - negative_prompt_embeds = t5_encoder( - [""] * batch_size, max_length=128 if fixed_negative else None - ).to(device, dtype) - - # diffusers pipeline concatenate `prompt_embeds` too early... - # https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913 - pipe.unet.flexible_max_length_workaround = [ - negative_prompt_embeds.size(1) - ] * batch_size + [prompt_embeds.size(1)] * batch_size - - max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)]) - b, _, d = prompt_embeds.shape - prompt_embeds = torch.cat( - [ - prompt_embeds, - torch.zeros( - (b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype - ), - ], - dim=1, - ) - negative_prompt_embeds = torch.cat( - [ - negative_prompt_embeds, - torch.zeros( - (b, max_length - negative_prompt_embeds.size(1), d), - device=device, - dtype=dtype, - ), - ], - dim=1, - ) - - images = pipe( - prompt_embeds=prompt_embeds, - negative_prompt_embeds=negative_prompt_embeds, - **pipe_kwargs, - output_type=output_type, - ).images - pipe.unet.flexible_max_length_workaround = None - return images - - -def load_ella(filename, device, dtype): - ella = ELLA() - safetensors.torch.load_model(ella, filename, strict=True) - ella.to(device, dtype=dtype) - return ella - - -def load_ella_for_pipe(pipe, ella): - pipe.unet = ELLAProxyUNet(ella, pipe.unet) - - -def offload_ella_for_pipe(pipe): - pipe.unet = pipe.unet.unet - - -def generate_image_with_fixed_max_length( - pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs -): - prompt = [prompt] if isinstance(prompt, str) else prompt - - prompt_embeds = t5_encoder(prompt, max_length=128).to(pipe.device, pipe.dtype) - negative_prompt_embeds = t5_encoder([""] * len(prompt), max_length=128).to( - pipe.device, pipe.dtype - ) - - return pipe( - prompt_embeds=prompt_embeds, - negative_prompt_embeds=negative_prompt_embeds, - **pipe_kwargs, - output_type=output_type, - ).images - - -def build_demo(ella_path, sd_path="runwayml/stable-diffusion-v1-5"): - pipe = StableDiffusionPipeline.from_pretrained( - sd_path, - torch_dtype=torch.float16, - safety_checker=None, - feature_extractor=None, - requires_safety_checker=False, - ) - pipe = pipe.to("cuda") - pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config) - - ella = load_ella(ella_path, pipe.device, pipe.dtype) - t5_encoder = T5TextEmbedder().to(pipe.device, dtype=torch.float16) - - def generate_images( - prompt, guidance_scale, seed, num_inference_steps, size=512, _batch_size=2 - ): - print("#" * 50) - print(prompt) - load_ella_for_pipe(pipe, ella) - image_flexible = generate_image_with_flexible_max_length( - pipe, - t5_encoder, - [prompt] * _batch_size, - guidance_scale=guidance_scale, - num_inference_steps=num_inference_steps, - height=size, - width=size, - generator=[ - torch.Generator(device="cuda").manual_seed(seed + i) - for i in range(_batch_size) - ], - output_type="pil", - ) - offload_ella_for_pipe(pipe) - - image_ori = pipe( - [prompt] * _batch_size, - output_type="pil", - guidance_scale=guidance_scale, - num_inference_steps=num_inference_steps, - height=size, - width=size, - generator=[ - torch.Generator(device="cuda").manual_seed(seed + i) - for i in range(_batch_size) - ], - ).images - - return image_ori, image_flexible - - with gr.Blocks() as app: - gr.Markdown( - """ - # ELLA-SD1.5 vs SD1.5 - - [ELLA Project](https://ella-diffusion.github.io/) - - ## Notes - - ** short prompt also works, but the result is much better after the caption is refined. ** - - ### Caption Refining with In Context Learning(ICL) - - caption refining instruction example: - ``` - Please generate the long prompt version of the short one according to the given examples. Long prompt version should consist of 3 to 5 sentences. Long prompt version must sepcify the color, shape, texture or spatial relation of the included objects. DO NOT generate sentences that describe any atmosphere!!! - - Short: A calico cat with eyes closed is perched upon a Mercedes. - Long: a multicolored cat perched atop a shiny black car. the car is parked in front of a building with wooden walls and a green fence. the reflection of the car and the surrounding environment can be seen on the car's glossy surface. - - Short: A boys sitting on a chair holding a video game remote. - Long: a young boy sitting on a chair, wearing a blue shirt and a baseball cap with the letter 'm'. he has a red medal around his neck and is holding a white game controller. behind him, there are two other individuals, one of whom is wearing a backpack. to the right of the boy, there's a blue trash bin with a sign that reads 'automatic party'. - - Short: A man is on the bank of the water fishing. - Long: a serene waterscape where a person, dressed in a blue jacket and a red beanie, stands in shallow waters, fishing with a long rod. the calm waters are dotted with several sailboats anchored at a distance, and a mountain range can be seen in the background under a cloudy sky. - - Short: A kitchen with a cluttered counter and wooden cabinets. - Long: a well-lit kitchen with wooden cabinets, a black and white checkered floor, and a refrigerator adorned with a floral decal on its side. the kitchen countertop holds various items, including a coffee maker, jars, and fruits. - - Short: a racoon holding a shiny red apple over its head - ``` - - using: https://huggingface.co/spaces/Qwen/Qwen-72B-Chat-Demo - got: a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby. - """ - ) - with gr.Row(): - input_caption = gr.Textbox( - value="A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers." - ) - with gr.Column(): - guidance_scale = gr.Slider( - minimum=1.0, maximum=16.0, value=10, label="guidance_scale" - ) - seed = gr.Slider( - minimum=1000, maximum=2**20, value=1000, label="random seed" - ) - num_inference_steps = gr.Slider( - minimum=15, maximum=100, value=25, label="num_inference_steps" - ) - - with gr.Row(): - with gr.Column(): - gr.Markdown(f"### ORIGINAL Stable Diffusion Model") - sd_output_image_gallery = gr.Gallery(columns=2, label="ORIGINAL SD") - with gr.Column(): - gr.Markdown(f"### ELLA") - ella_output_image_gallery = gr.Gallery(columns=2, label="ELLA") - submit_button = gr.Button() - - submit_button.click( - fn=generate_images, - inputs=[input_caption, guidance_scale, seed, num_inference_steps], - outputs=[sd_output_image_gallery, ella_output_image_gallery], - ) - - app.queue(concurrency_count=1, api_open=False) - app.launch(share=False) - - -def main(save_folder, ella_path): - save_folder = Path(save_folder) - save_folder.mkdir(exist_ok=True) - - pipe = StableDiffusionPipeline.from_pretrained( - "runwayml/stable-diffusion-v1-5", - torch_dtype=torch.float16, - safety_checker=None, - feature_extractor=None, - requires_safety_checker=False, - ) - pipe = pipe.to("cuda") - pipe.scheduler = DPMSolverMultistepScheduler.from_config(pipe.scheduler.config) - - ella = load_ella(ella_path, pipe.device, pipe.dtype) - t5_encoder = T5TextEmbedder().to(pipe.device, dtype=torch.float16) - - # prompt from ViLG-300, PartiPrompts - # short prompt also works, but the result is much better after the caption is refined. - - # caption refining instruction example: - # ``` - # Please generate the long prompt version of the short one according to the given examples. Long prompt version should consist of 3 to 5 sentences. Long prompt version must sepcify the color, shape, texture or spatial relation of the included objects. DO NOT generate sentences that describe any atmosphere!!! - # - # Short: A calico cat with eyes closed is perched upon a Mercedes. - # Long: a multicolored cat perched atop a shiny black car. the car is parked in front of a building with wooden walls and a green fence. the reflection of the car and the surrounding environment can be seen on the car's glossy surface. - # - # Short: A boys sitting on a chair holding a video game remote. - # Long: a young boy sitting on a chair, wearing a blue shirt and a baseball cap with the letter 'm'. he has a red medal around his neck and is holding a white game controller. behind him, there are two other individuals, one of whom is wearing a backpack. to the right of the boy, there's a blue trash bin with a sign that reads 'automatic party'. - # - # Short: A man is on the bank of the water fishing. - # Long: a serene waterscape where a person, dressed in a blue jacket and a red beanie, stands in shallow waters, fishing with a long rod. the calm waters are dotted with several sailboats anchored at a distance, and a mountain range can be seen in the background under a cloudy sky. - # - # Short: A kitchen with a cluttered counter and wooden cabinets. - # Long: a well-lit kitchen with wooden cabinets, a black and white checkered floor, and a refrigerator adorned with a floral decal on its side. the kitchen countertop holds various items, including a coffee maker, jars, and fruits. - # - # Short: a racoon holding a shiny red apple over its head - # ``` - # - # using: https://huggingface.co/spaces/Qwen/Qwen-72B-Chat-Demo - # got: a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby. - - prompt_name_examples1 = [ - ("crocodile_sweater", "Crocodile in a sweater"), - ( - "crocodile_sweater-gpt4_refined_caption", - "a large, textured green crocodile lying comfortably on a patch of grass with a cute, knitted orange sweater enveloping its scaly body. Around its neck, the sweater features a whimsical pattern of blue and yellow stripes. In the background, a smooth, grey rock partially obscures the view of a small pond with lily pads floating on the surface.", - ), - ("red_book-yellow_vase", "A red book and a yellow vase."), - ( - "red_book-yellow_vase-gpt4_refined_caption", - "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.", - ), - ("racoon_apple", "a racoon holding a shiny red apple over its head"), - ( - "racoon_apple_Qwen-72B-Chat-refined", - "a mischievous raccoon standing on its hind legs, holding a bright red apple aloft in its furry paws. the apple shines brightly against the backdrop of a dense forest, with leaves rustling in the gentle breeze. a few scattered rocks can be seen on the ground beneath the raccoon's feet, while a gnarled tree trunk stands nearby.", - ), - ] - - # hard example prompt. - prompt_name_examples2 = [ - ( - "falcon_chinese", - "a chinese man wearing a white shirt and a checkered headscarf, holds a large falcon near his shoulder. the falcon has dark feathers with a distinctive beak. the background consists of a clear sky and a fence, suggesting an outdoor setting, possibly a desert or arid region", - ), - ( - "wombat", - "A close-up photo of a wombat wearing a red backpack and raising both arms in the air. Mount Rushmore is in the background", - ), - ( - "bakkot_AstralCodexTen_2", - "An oil painting of a man in a factory looking at a cat wearing a top hat", - ), - ] - - for name, prompt in prompt_name_examples1 + prompt_name_examples2: - print("#" * 80) - print(f'{name}: "{prompt}"') - _batch_size = 1 - size = 512 - seed = 1001 - prompt = [prompt] * _batch_size - - load_ella_for_pipe(pipe, ella) - image_flexible = generate_image_with_flexible_max_length( - pipe, - t5_encoder, - prompt, - guidance_scale=12, - num_inference_steps=50, - height=size, - width=size, - generator=[ - torch.Generator(device="cuda").manual_seed(seed + i) - for i in range(_batch_size) - ], - ) - image_fixed = generate_image_with_fixed_max_length( - pipe, - t5_encoder, - prompt, - guidance_scale=12, - num_inference_steps=50, - height=size, - width=size, - generator=[ - torch.Generator(device="cuda").manual_seed(seed + i) - for i in range(_batch_size) - ], - ) - offload_ella_for_pipe(pipe) - - image_ori = pipe( - prompt, - output_type="pt", - guidance_scale=12, - num_inference_steps=50, - height=size, - width=size, - generator=[ - torch.Generator(device="cuda").manual_seed(seed + i) - for i in range(_batch_size) - ], - ).images - - print(f'save image at {save_folder / f"{name}.png"}') - print( - "original SD1.5\t|\tELLA-SD1.5(fixed token length)\t|\tELLA-SD1.5(flexible token length)" - ) - save_image( - torch.cat([image_ori, image_fixed, image_flexible], dim=0), - save_folder / f"{name}.png", - nrow=3, - ) - - -if __name__ == "__main__": - fire.Fire(dict(test=main, demo=build_demo)) diff --git a/nodes.py b/nodes.py index b936499..c39c92b 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,4 @@ import os -from contextlib import nullcontext import torch import torch.nn.functional as F @@ -23,20 +22,23 @@ try: create_text_encoder_from_ldm_clip_checkpoint ) except: - raise ImportError("Diffusers version too old. Please update to 0.26.0 minimum.") + raise ImportError("Diffusers version too old. Please update to 0.27.2 minimum.") from .brushnet.pipeline_brushnet import StableDiffusionBrushNetPipeline from .brushnet.brushnet import BrushNetModel from .brushnet.unet_2d_condition import UNet2DConditionModel +import safetensors.torch from omegaconf import OmegaConf from transformers import CLIPTokenizer + import comfy.model_management as mm import comfy.utils import folder_paths script_directory = os.path.dirname(os.path.abspath(__file__)) - +IS_MODEL_CPU_OFFLOAD_ENABLED = False + class brushnet_model_loader: @classmethod def INPUT_TYPES(s): @@ -70,6 +72,8 @@ class brushnet_model_loader: "brushnet_model": brushnet_model } if not hasattr(self, "model") or self.model == None or custom_config != self.current_config: + global IS_MODEL_CPU_OFFLOAD_ENABLED + IS_MODEL_CPU_OFFLOAD_ENABLED = False pbar = comfy.utils.ProgressBar(5) self.current_config = custom_config @@ -144,14 +148,14 @@ class brushnet_model_loader: safety_checker=None, feature_extractor=None ) - self.pipe.enable_model_cpu_offload() + #self.pipe.enable_model_cpu_offload() pbar.update(1) brushnet = { "pipe": self.pipe, } return (brushnet,) - + class brushnet_sampler: @classmethod def INPUT_TYPES(s): @@ -160,7 +164,12 @@ class brushnet_sampler: "image": ("IMAGE",), "mask": ("MASK",), "steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}), - "guidance_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01}), + "cfg": ("FLOAT", {"default": 7.5, "min": 0.0, "max": 20.0, "step": 0.01}), + "cfg_brushnet": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01}), + "control_guidance_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "control_guidance_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "guess_mode": ("BOOLEAN", {"default": False}), + "clip_skip": ("INT", {"default": 0, "min": 0, "max": 20, "step": 1}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "scheduler": ( [ @@ -177,8 +186,8 @@ class brushnet_sampler: "default": "UniPCMultistepScheduler" }), "prompt": ("STRING", {"multiline": True, "default": "caption",}), - - }, + "n_prompt": ("STRING", {"multiline": True, "default": "caption",}), + }, } RETURN_TYPES = ("IMAGE",) @@ -186,12 +195,16 @@ class brushnet_sampler: FUNCTION = "process" CATEGORY = "BrushNetWrapper" - def process(self, brushnet, image, mask, prompt, steps, guidance_scale, seed, scheduler): + def process(self, brushnet, image, mask, prompt, n_prompt, steps, cfg, guess_mode, clip_skip, + cfg_brushnet, control_guidance_start, control_guidance_end, seed, scheduler): device = mm.get_torch_device() - mm.unload_all_models() mm.soft_empty_cache() pipe=brushnet["pipe"] - #pipe.to(device, dtype=dtype) + + global IS_MODEL_CPU_OFFLOAD_ENABLED + if not IS_MODEL_CPU_OFFLOAD_ENABLED: + pipe.enable_model_cpu_offload() + IS_MODEL_CPU_OFFLOAD_ENABLED = True scheduler_config = { "num_train_timesteps": 1000, @@ -244,16 +257,28 @@ class brushnet_sampler: if len(prompt_list) < B: prompt_list += [prompt_list[-1]] * (B - len(prompt_list)) + n_prompt_list = [] + n_prompt_list.append(n_prompt) + if len(n_prompt_list) < B: + n_prompt_list += [n_prompt_list[-1]] * (B - len(n_prompt_list)) + #sample generator = torch.Generator(device).manual_seed(seed) - + print(prompt_list) images = pipe( - prompt_list, - image=image, + prompt_list, + negative_prompt=n_prompt_list, + image=image, + ipadapter_image=None, mask=resized_mask, num_inference_steps=steps, generator=generator, - brushnet_conditioning_scale=guidance_scale, + guidance_scale=cfg, + guess_mode=guess_mode, + clip_skip=clip_skip if clip_skip > 0 else None, + brushnet_conditioning_scale=cfg_brushnet, + control_guidance_start=control_guidance_start, + control_guidance_end=control_guidance_end, output_type="pt" ).images @@ -261,11 +286,165 @@ class brushnet_sampler: return (image_out,) +class brushnet_ella_loader: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "brushnet": ("BRUSHNET",), + }, + } + + RETURN_TYPES = ("BRUSHNET",) + RETURN_NAMES = ("brushnet",) + FUNCTION = "loadmodel" + CATEGORY = "ELLA-Wrapper" + + def loadmodel(self, brushnet): + print("loading ELLA") + from .ella.model import ELLA + from .ella.ella_unet import ELLAProxyUNet + checkpoint_path = os.path.join(folder_paths.models_dir,'ella') + ella_path = os.path.join(checkpoint_path, 'ella-sd1.5-tsc-t5xl.safetensors') + if not os.path.exists(ella_path): + from huggingface_hub import snapshot_download + snapshot_download(repo_id="QQGYLab/ELLA", local_dir=checkpoint_path, local_dir_use_symlinks=False) + + ella = ELLA() + safetensors.torch.load_model(ella, ella_path, strict=True) + ella_unet = ELLAProxyUNet(ella, brushnet['pipe'].unet) + brushnet['pipe'].unet = ella_unet + + return (brushnet,) + +class brushnet_sampler_ella: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "brushnet": ("BRUSHNET",), + "ella_embeds": ("ELLAEMBEDS",), + "image": ("IMAGE",), + "mask": ("MASK",), + "steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}), + "cfg": ("FLOAT", {"default": 7.5, "min": 0.0, "max": 20.0, "step": 0.01}), + "cfg_brushnet": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01}), + "control_guidance_start": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "control_guidance_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "guess_mode": ("BOOLEAN", {"default": False}), + "clip_skip": ("INT", {"default": 0, "min": 0, "max": 20, "step": 1}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + "scheduler": ( + [ + "DPMSolverMultistepScheduler", + "DPMSolverMultistepScheduler_SDE_karras", + "DDPMScheduler", + "LCMScheduler", + "PNDMScheduler", + "DEISMultistepScheduler", + "EulerDiscreteScheduler", + "EulerAncestralDiscreteScheduler", + "UniPCMultistepScheduler" + ], { + "default": "UniPCMultistepScheduler" + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "process" + CATEGORY = "BrushNetWrapper" + + def process(self, brushnet, image, mask, steps, cfg, guess_mode, clip_skip, ella_embeds, + cfg_brushnet, control_guidance_start, control_guidance_end, seed, scheduler): + device = mm.get_torch_device() + dtype = mm.unet_dtype() + mm.soft_empty_cache() + pipe=brushnet["pipe"].to(dtype) + + global IS_MODEL_CPU_OFFLOAD_ENABLED + if not IS_MODEL_CPU_OFFLOAD_ENABLED: + pipe.enable_model_cpu_offload() + IS_MODEL_CPU_OFFLOAD_ENABLED = True + + scheduler_config = { + "num_train_timesteps": 1000, + "beta_start": 0.00085, + "beta_end": 0.012, + "beta_schedule": "scaled_linear", + "steps_offset": 1, + } + if scheduler == "DPMSolverMultistepScheduler": + noise_scheduler = DPMSolverMultistepScheduler(**scheduler_config) + elif scheduler == "DPMSolverMultistepScheduler_SDE_karras": + scheduler_config.update({"algorithm_type": "sde-dpmsolver++"}) + scheduler_config.update({"use_karras_sigmas": True}) + noise_scheduler = DPMSolverMultistepScheduler(**scheduler_config) + elif scheduler == "DDPMScheduler": + noise_scheduler = DDPMScheduler(**scheduler_config) + elif scheduler == "LCMScheduler": + noise_scheduler = LCMScheduler(**scheduler_config) + elif scheduler == "PNDMScheduler": + scheduler_config.update({"set_alpha_to_one": False}) + scheduler_config.update({"trained_betas": None}) + noise_scheduler = PNDMScheduler(**scheduler_config) + elif scheduler == "DEISMultistepScheduler": + noise_scheduler = DEISMultistepScheduler(**scheduler_config) + elif scheduler == "EulerDiscreteScheduler": + noise_scheduler = EulerDiscreteScheduler(**scheduler_config) + elif scheduler == "EulerAncestralDiscreteScheduler": + noise_scheduler = EulerAncestralDiscreteScheduler(**scheduler_config) + elif scheduler == "UniPCMultistepScheduler": + noise_scheduler = UniPCMultistepScheduler(**scheduler_config) + pipe.scheduler = noise_scheduler + + B, H, W, C = image.shape + image = image.permute(0, 3, 1, 2).to(device) + + #handle masks + if len(mask.shape) == 2: + mask = mask.unsqueeze(0) + mask = mask.to(device) + if mask.shape[0] < B: + repeat_times = B // mask.shape[0] + mask = mask.repeat(repeat_times, 1, 1, 1) + resized_mask = F.interpolate(mask.unsqueeze(1), size=[H, W], mode='nearest').squeeze(1) + + image = image * (1-resized_mask) + + #sample + generator = torch.Generator(device).manual_seed(seed) + + images = pipe( + prompt=None, + negative_prompt=None, + prompt_embeds=ella_embeds["prompt_embeds"], + negative_prompt_embeds=ella_embeds["negative_prompt_embeds"], + image=image, + ipadapter_image=None, + mask=resized_mask, + num_inference_steps=steps, + generator=generator, + guidance_scale=cfg, + guess_mode=guess_mode, + clip_skip=clip_skip if clip_skip > 0 else None, + brushnet_conditioning_scale=cfg_brushnet, + control_guidance_start=control_guidance_start, + control_guidance_end=control_guidance_end, + output_type="pt" + ).images + + image_out = images.permute(0, 2, 3, 1).cpu().float() + return (image_out,) + NODE_CLASS_MAPPINGS = { "brushnet_model_loader": brushnet_model_loader, "brushnet_sampler": brushnet_sampler, + "brushnet_sampler_ella": brushnet_sampler_ella, + "brushnet_ella_loader": brushnet_ella_loader } NODE_DISPLAY_NAME_MAPPINGS = { "brushnet_model_loader": "BrushNet Model Loader", "brushnet_sampler": "BrushNet Sampler", + "brushnet_sampler_ella": "BrushNet Sampler (ELLA)", + "brushnet_ella_loader": "BrushNet ELLA Loader" }