Eexpose more options, add ELLA support, cleanup diffusers
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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
+593
-290
File diff suppressed because it is too large
Load Diff
@@ -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": [],
|
||||
|
||||
@@ -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
@@ -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))
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user