Eexpose more options, add ELLA support, cleanup diffusers

This commit is contained in:
kijai
2024-04-13 01:16:12 +03:00
parent c101fa2a9e
commit 870996362d
12 changed files with 2285 additions and 1887 deletions
+8 -6
View File
@@ -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
+3 -3
View File
@@ -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)
+143 -86
View File
@@ -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:
+9 -3
View File
@@ -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,
)
+67
View File
@@ -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,
)
+218
View File
@@ -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
File diff suppressed because it is too large Load Diff
+253 -279
View File
@@ -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": [],
+213 -210
View File
@@ -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": [],
-419
View File
@@ -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))
+194 -15
View File
@@ -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"
}