Added Flex.2 conditioning node.
This commit is contained in:
@@ -11,3 +11,4 @@ Clone this repo into your `custom_nodes` directory.
|
||||
- **Flex Guidance**: Allows you to set the guidance for the Flex.1 guidance embedder, or bypass it completly to use true CFG.
|
||||
- **Flex LoRA Loader**: Loads LoRAs and automatically prunes them to Flex.1 layers. It will not be perfect as Flex is heavily diverged from Flux dev and is not a direct ancenstor of it, but it should be good enough for most purposes.
|
||||
- **Flex LoRA Loader (Model Only)**: Same as Flex LoRA Loader, but only loads the model and not the text encoder. Most Flux LoRAs do not train the text encoder.
|
||||
- **Flex2 Conditioner** A conditionaing node for controlling all of the [Flex.2-preview](https://huggingface.co/ostris/Flex.2-preview) conditioning for inpaint and universal controls. This node is currently required for [Flex.2-preview](https://huggingface.co/ostris/Flex.2-preview) inference.
|
||||
+316
@@ -1,7 +1,13 @@
|
||||
from functools import partial
|
||||
import torch
|
||||
from comfy.model_base import Flux
|
||||
import folder_paths
|
||||
import node_helpers
|
||||
import comfy.sd
|
||||
import comfy.utils
|
||||
import comfy.patcher_extension
|
||||
import comfy.conds
|
||||
from comfy.patcher_extension import CallbacksMP, WrappersMP
|
||||
|
||||
|
||||
class FlexGuidance:
|
||||
@@ -150,14 +156,324 @@ class FlexLoraLoaderModelOnly(FlexLoraLoader):
|
||||
return (self.load_lora(model, None, lora_name, strength_model, 0)[0],)
|
||||
|
||||
|
||||
def flex2_concat_cond(self: Flux, **kwargs):
|
||||
# will break otherwise
|
||||
return None
|
||||
|
||||
|
||||
def flex2_extra_conds(self, **kwargs):
|
||||
out = self._flex2_orig_extra_conds(**kwargs)
|
||||
noise = kwargs.get("noise", None)
|
||||
device = kwargs["device"]
|
||||
flex2_concat_latent = kwargs.get("flex2_concat_latent", None)
|
||||
flex2_concat_latent_no_control = kwargs.get(
|
||||
"flex2_concat_latent_no_control", None)
|
||||
control_strength = kwargs.get("flex2_control_strength", 1.0)
|
||||
control_start_percent = kwargs.get("flex2_control_start_percent", 0.0)
|
||||
control_end_percent = kwargs.get("flex2_control_end_percent", 0.1)
|
||||
if flex2_concat_latent is not None:
|
||||
flex2_concat_latent = comfy.utils.resize_to_batch_size(
|
||||
flex2_concat_latent, noise.shape[0])
|
||||
flex2_concat_latent = self.process_latent_in(flex2_concat_latent)
|
||||
flex2_concat_latent = flex2_concat_latent.to(device)
|
||||
out['flex2_concat_latent'] = comfy.conds.CONDNoiseShape(
|
||||
flex2_concat_latent)
|
||||
if flex2_concat_latent_no_control is not None:
|
||||
flex2_concat_latent_no_control = comfy.utils.resize_to_batch_size(
|
||||
flex2_concat_latent_no_control, noise.shape[0])
|
||||
flex2_concat_latent_no_control = self.process_latent_in(
|
||||
flex2_concat_latent_no_control)
|
||||
flex2_concat_latent_no_control = flex2_concat_latent_no_control.to(
|
||||
device)
|
||||
out['flex2_concat_latent_no_control'] = comfy.conds.CONDNoiseShape(
|
||||
flex2_concat_latent_no_control)
|
||||
|
||||
out['flex2_control_start_percent'] = comfy.conds.CONDConstant(
|
||||
control_start_percent)
|
||||
out['flex2_control_end_percent'] = comfy.conds.CONDConstant(
|
||||
control_end_percent)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def flex_apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs):
|
||||
sigma = t
|
||||
xc = self.model_sampling.calculate_input(sigma, x)
|
||||
if c_concat is not None:
|
||||
xc = torch.cat([xc] + [c_concat], dim=1)
|
||||
|
||||
flex2_control_start_sigma = 1.0 - \
|
||||
kwargs.get("flex2_control_start_percent", 0.0)
|
||||
flex2_control_end_sigma = 1.0 - \
|
||||
kwargs.get("flex2_control_end_percent", 1.0)
|
||||
|
||||
flex2_concat_latent_active = kwargs.get("flex2_concat_latent", None)
|
||||
flex2_concat_latent_inactive = kwargs.get(
|
||||
"flex2_concat_latent_no_control", None)
|
||||
|
||||
sigma_float = sigma.mean().cpu().item()
|
||||
sigma_int = int(sigma_float * 1000)
|
||||
|
||||
# simple, but doesnt work right because of shift
|
||||
is_being_controlled = sigma_float <= flex2_control_start_sigma and sigma_float >= flex2_control_end_sigma
|
||||
|
||||
sigmas = transformer_options.get("sample_sigmas", None)
|
||||
|
||||
if sigmas is not None:
|
||||
# we have all the timesteps here, determine what percent we are through the
|
||||
# timesteps we are doing. This way is more intuitive to user.
|
||||
all_timesteps = [int(sigma.cpu().item() * 1000) for sigma in sigmas]
|
||||
current_idx = all_timesteps.index(sigma_int)
|
||||
current_percent = current_idx / len(all_timesteps)
|
||||
current_percent_sigma = 1.0 - current_percent
|
||||
is_being_controlled = current_percent_sigma <= flex2_control_start_sigma and current_percent_sigma >= flex2_control_end_sigma
|
||||
|
||||
if is_being_controlled:
|
||||
# it is active
|
||||
xc = torch.cat([xc] + [flex2_concat_latent_active], dim=1)
|
||||
else:
|
||||
# it is inactive
|
||||
xc = torch.cat([xc] + [flex2_concat_latent_inactive], dim=1)
|
||||
|
||||
context = c_crossattn
|
||||
dtype = self.get_dtype()
|
||||
|
||||
if self.manual_cast_dtype is not None:
|
||||
dtype = self.manual_cast_dtype
|
||||
|
||||
xc = xc.to(dtype)
|
||||
t = self.model_sampling.timestep(t).float()
|
||||
if context is not None:
|
||||
context = context.to(dtype)
|
||||
|
||||
extra_conds = {}
|
||||
for o in kwargs:
|
||||
extra = kwargs[o]
|
||||
if hasattr(extra, "dtype"):
|
||||
if extra.dtype != torch.int and extra.dtype != torch.long:
|
||||
extra = extra.to(dtype)
|
||||
extra_conds[o] = extra
|
||||
|
||||
t = self.process_timestep(t, x=x, **extra_conds)
|
||||
model_output = self.diffusion_model(
|
||||
xc, t, context=context, control=control, transformer_options=transformer_options, **extra_conds).float()
|
||||
return self.model_sampling.calculate_denoised(sigma, model_output, x)
|
||||
|
||||
|
||||
class Flex2Conditioner:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL", ),
|
||||
"vae": ("VAE", ),
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"bypass_guidance_embedder": (["yes", "no"], {"default": "no"}),
|
||||
"guidance": ("FLOAT", {"default": 3.5, "min": 0.0, "max": 100.0, "step": 0.1}),
|
||||
"control_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"control_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"control_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01})
|
||||
},
|
||||
"optional": {
|
||||
"latent": ("LATENT", ),
|
||||
"inpaint_image": ("IMAGE", ),
|
||||
"inpaint_mask": ("MASK", ),
|
||||
"control_image": ("IMAGE", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "LATENT")
|
||||
RETURN_NAMES = ("model", "positive", "negative", "latent")
|
||||
FUNCTION = "do_it"
|
||||
|
||||
CATEGORY = "advanced/conditioning/flex"
|
||||
|
||||
def do_it(
|
||||
self,
|
||||
model,
|
||||
vae,
|
||||
positive,
|
||||
negative,
|
||||
guidance,
|
||||
bypass_guidance_embedder,
|
||||
control_strength,
|
||||
control_start_percent,
|
||||
control_end_percent,
|
||||
latent=None,
|
||||
inpaint_image=None,
|
||||
inpaint_mask=None,
|
||||
control_image=None,
|
||||
):
|
||||
# replace concat_cond of the flux model as default one breaks flex2
|
||||
# todo, find a better way to do this
|
||||
if not hasattr(model.model, "_flex2_orig_concat_cond"):
|
||||
model.model._flex2_orig_concat_cond = model.model.concat_cond
|
||||
model.model.concat_cond = partial(flex2_concat_cond, model.model)
|
||||
# replace extra_conds
|
||||
if not hasattr(model.model, "_flex2_orig_extra_conds"):
|
||||
model.model._flex2_orig_extra_conds = model.model.extra_conds
|
||||
model.model.extra_conds = partial(flex2_extra_conds, model.model)
|
||||
|
||||
if not hasattr(model.model, "_flex2_orig_apply_model"):
|
||||
model.model._flex2_orig_apply_model = model.model._apply_model
|
||||
model.model._apply_model = partial(flex_apply_model, model.model)
|
||||
|
||||
# masks come in as (bs, h, w) 0 to 1
|
||||
# images come in as (bs, h, w, c) 0 to 1
|
||||
# latents come in as (bs, ch, h, w) -1 to 1
|
||||
batch_size = 1
|
||||
latent_height: int = None
|
||||
latent_width: int = None
|
||||
|
||||
# DETERIMINE SIZES
|
||||
# size order is latent size, inpaint size, control size
|
||||
if latent is not None:
|
||||
latent_height = latent['samples'].shape[2]
|
||||
latent_width = latent['samples'].shape[3]
|
||||
if latent['samples'].shape[1] == 4:
|
||||
# make it 16 channels
|
||||
latent['samples'] = torch.cat(
|
||||
[latent['samples'] for _ in range(4)], dim=1)
|
||||
batch_size = latent['samples'].shape[0]
|
||||
elif inpaint_image is not None:
|
||||
batch_size = inpaint_image.shape[0]
|
||||
latent_height = inpaint_image.shape[1] // 8
|
||||
latent_width = inpaint_image.shape[2] // 8
|
||||
elif control_image is not None:
|
||||
batch_size = control_image.shape[0]
|
||||
latent_height = control_image.shape[1] // 8
|
||||
latent_width = control_image.shape[2] // 8
|
||||
else:
|
||||
raise ValueError("No latent, inpaint or control image provided")
|
||||
|
||||
img_width = latent_width * 8
|
||||
img_height = latent_height * 8
|
||||
|
||||
# apply differential diffusion to model
|
||||
model = model.clone()
|
||||
# model.set_model_denoise_mask_function(self.denoise_mask_function)
|
||||
|
||||
# guidance embedder
|
||||
bypass_guidance_embedder = bypass_guidance_embedder == "yes"
|
||||
|
||||
guidance_value = guidance
|
||||
if bypass_guidance_embedder:
|
||||
guidance_value = None
|
||||
positive = node_helpers.conditioning_set_values(
|
||||
positive,
|
||||
{
|
||||
"guidance": guidance_value
|
||||
}
|
||||
)
|
||||
|
||||
# out input is our latent(16) + (inpaint_image(16) + mask(1) + control image(16))
|
||||
# We just need to build the non latent part
|
||||
concat_latent = torch.zeros(
|
||||
(batch_size, 33, latent_height, latent_width),
|
||||
device='cpu',
|
||||
dtype=torch.float32
|
||||
)
|
||||
|
||||
# when we are not using inpainting, the mask needs to be all 1s
|
||||
concat_latent[:, 16:17, :, :] = torch.ones(
|
||||
(batch_size, 1, latent_height, latent_width),
|
||||
device='cpu',
|
||||
dtype=torch.float32
|
||||
)
|
||||
|
||||
out_latent = {
|
||||
"samples": torch.zeros(
|
||||
(batch_size, 16, latent_height, latent_width),
|
||||
device='cpu',
|
||||
dtype=torch.float32
|
||||
)
|
||||
}
|
||||
|
||||
# inpaint
|
||||
if inpaint_image is not None:
|
||||
if inpaint_image.shape[1] != img_height or inpaint_image.shape[2] != img_width:
|
||||
inpaint_image = torch.nn.functional.interpolate(
|
||||
inpaint_image.permute(0, 3, 1, 2),
|
||||
size=(img_height, img_width),
|
||||
mode="bilinear"
|
||||
).permute(0, 2, 3, 1)
|
||||
|
||||
if inpaint_mask is not None:
|
||||
inpaint_mask_latent = torch.nn.functional.interpolate(
|
||||
inpaint_mask.reshape(
|
||||
(-1, 1, inpaint_mask.shape[-2], inpaint_mask.shape[-1])),
|
||||
size=(latent_height, latent_width),
|
||||
mode="bilinear"
|
||||
)
|
||||
else:
|
||||
# make it all 1s
|
||||
inpaint_mask_latent = torch.ones(
|
||||
(batch_size, 1, latent_height, latent_width),
|
||||
device='cpu',
|
||||
dtype=torch.float32
|
||||
)
|
||||
|
||||
inpaint_latent_orig = vae.encode(inpaint_image)
|
||||
|
||||
# set this so we can partially denoise with it if desired
|
||||
out_latent["samples"] = inpaint_latent_orig.clone()
|
||||
|
||||
# mask is currently 0 for keep and 1 for inpaint
|
||||
inpaint_latent_masked = inpaint_latent_orig * \
|
||||
(1 - inpaint_mask_latent)
|
||||
|
||||
# put it on our concat latent
|
||||
concat_latent[:, 0:16, :, :] = inpaint_latent_masked
|
||||
# put the mask in the last channel, 0 for keep, 1 for inpaint
|
||||
concat_latent[:, 16:17, :, :] = inpaint_mask_latent
|
||||
|
||||
concat_latent_no_control = concat_latent.clone()
|
||||
|
||||
# control
|
||||
if control_image is not None:
|
||||
if control_image.shape[1] != img_height or control_image.shape[2] != img_width:
|
||||
control_image = torch.nn.functional.interpolate(
|
||||
control_image.permute(0, 3, 1, 2),
|
||||
size=(img_height, img_width),
|
||||
mode="bilinear"
|
||||
).permute(0, 2, 3, 1)
|
||||
|
||||
control_latent = vae.encode(control_image)
|
||||
|
||||
# put the control image in the last 16 channels
|
||||
concat_latent[:, 17:33, :, :] = control_latent * \
|
||||
control_strength
|
||||
|
||||
out = []
|
||||
for conditioning in [positive, negative]:
|
||||
c = node_helpers.conditioning_set_values(
|
||||
conditioning,
|
||||
{
|
||||
"flex2_concat_latent": concat_latent,
|
||||
"flex2_concat_latent_no_control": concat_latent_no_control,
|
||||
"flex2_control_strength": control_strength,
|
||||
"flex2_control_start_percent": control_start_percent,
|
||||
"flex2_control_end_percent": control_end_percent,
|
||||
}
|
||||
)
|
||||
out.append(c)
|
||||
positive = out[0]
|
||||
negative = out[1]
|
||||
|
||||
return (model, positive, negative, out_latent)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"FlexGuidance": FlexGuidance,
|
||||
"FlexLoraLoader": FlexLoraLoader,
|
||||
"FlexLoraLoaderModelOnly": FlexLoraLoaderModelOnly,
|
||||
"Flex2Conditioner": Flex2Conditioner,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FlexGuidance": "Flex Guidance",
|
||||
"FlexLoraLoader": "Flex LoRA Loader",
|
||||
"FlexLoraLoaderModelOnly": "Flex LoRA Loader (Model Only)",
|
||||
"Flex2Conditioner": "Flex 2 Conditioner",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user