Files
laksjdjf-cgem156-ComfyUI/scripts/reference/reference.py
T

218 lines
7.9 KiB
Python

import torch
from ... import ROOT_NAME
CATEGORY_NAME = ROOT_NAME + "reference"
class ReferenceApply:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"index": ("INT", {"default": 0, "min": 0, "max": 256}),
"mode": (["concat", "replace"], {"default": "concat"}),
"depth": ("INT", {"default": 12, "min": -1, "max": 12}),
"start_step": ("FLOAT", {"default": 0,"min": 0, "max": 1, "step": 0.01}),
"end_step": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.01}),
"apply_input": ("BOOLEAN", {"default": True}),
"apply_middle": ("BOOLEAN", {"default": True}),
"apply_output": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("MODEL", )
FUNCTION = "reference_only"
CATEGORY = CATEGORY_NAME
def reference_only(self, model, index, mode, depth, start_step, end_step, apply_input, apply_middle, apply_output):
model_reference = model.clone()
start_sigma = model_reference.model.model_sampling.percent_to_sigma(start_step)
end_sigma = model_reference.model.model_sampling.percent_to_sigma(end_step)
self.depth = depth
self.sdxl = hasattr(model_reference.model.diffusion_model, "label_emb")
self.num_blocks = 8 if self.sdxl else 11
def reference_apply(q, k, v, extra_options):
block_name, block_id = extra_options["block"]
chunks = len(extra_options["cond_or_uncond"])
batch_size = q.shape[0] // chunks
if block_name == "input" and not apply_input:
return q, k, v
if block_name == "middle" and not apply_middle:
return q, k, v
if block_name == "output" and not apply_output:
return q, k, v
if block_name == "output":
block_number = self.num_blocks - block_id
else:
block_number = block_id
q_out = q.clone()
k_out = k.clone()
v_out = v.clone()
sigma = extra_options["sigmas"][0].item()
if end_sigma <= sigma <= start_sigma and block_number <= self.depth:
k_ref = k_out[index::batch_size].repeat_interleave(batch_size, dim=0).clone()
v_ref = v_out[index::batch_size].repeat_interleave(batch_size, dim=0).clone()
k_out = torch.cat([k_out, k_ref], dim=1) if mode == "concat" else k_ref
v_out = torch.cat([v_out, v_ref], dim=1) if mode == "concat" else v_ref
return q_out, k_out, v_out
model_reference.set_model_attn1_patch(reference_apply)
return (model_reference, )
class ReferenceLatent:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"latent": ("LATENT",),
"index": ("INT", {"default": 0, "min": 0, "max": 256}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256}),
}
}
RETURN_TYPES = ("LATENT", )
FUNCTION = "reference_latent"
CATEGORY = CATEGORY_NAME
def reference_latent(self, latent, index, batch_size):
latent_new = latent.copy()
sample = latent_new["samples"]
height, width = sample.shape[2], sample.shape[3]
empty_latent = torch.zeros_like(latent["samples"]).repeat(batch_size , 1, 1, 1)
empty_latent[index] = sample[0]
noise_mask = torch.ones(batch_size, 1, height * 8, width * 8).to(sample)
noise_mask[index] = 0.0
latent_new["samples"] = empty_latent
latent_new["noise_mask"] = noise_mask
return (latent_new, )
class MultipleReferenceApply:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"indices": ("STRING", {"default": "0"}),
"depth": ("INT", {"default": 12, "min": -1, "max": 12}),
"start_step": ("FLOAT", {"default": 0,"min": 0, "max": 1, "step": 0.01}),
"end_step": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.01}),
"apply_input": ("BOOLEAN", {"default": True}),
"apply_middle": ("BOOLEAN", {"default": True}),
"apply_output": ("BOOLEAN", {"default": True}),
"weights": ("STRING", {"default": ""}),
}
}
RETURN_TYPES = ("MODEL", )
FUNCTION = "multiple_reference_only"
CATEGORY = CATEGORY_NAME
def multiple_reference_only(self, model, indices, depth, start_step, end_step, apply_input, apply_middle, apply_output, weights):
model_reference = model.clone()
start_sigma = model_reference.model.model_sampling.percent_to_sigma(start_step)
end_sigma = model_reference.model.model_sampling.percent_to_sigma(end_step)
self.depth = depth
self.sdxl = hasattr(model_reference.model.diffusion_model, "label_emb")
self.num_blocks = 8 if self.sdxl else 11
indices = [int(i) for i in indices.split(",") if i.strip().isdigit()]
weights = [float(i) for i in weights.split(",") if i.strip()] if weights else [1.0] * len(indices)
def reference_apply(q, k, v, extra_options):
block_name, block_id = extra_options["block"]
if block_name == "input" and not apply_input:
return q, k, v
if block_name == "middle" and not apply_middle:
return q, k, v
if block_name == "output" and not apply_output:
return q, k, v
if block_name == "output":
block_number = self.num_blocks - block_id
else:
block_number = block_id
q_out = q.clone()
k_out = k.clone()
v_out = v.clone()
sigma = extra_options["sigmas"][0].item()
if end_sigma <= sigma <= start_sigma and block_number <= self.depth:
chunks = len(extra_options["cond_or_uncond"])
batch_size = q.shape[0] // chunks
num_tokens = q.shape[1]
k_refs = torch.cat([k_out[i::batch_size] for i in indices], dim=1)
v_refs = torch.cat([v_out[i::batch_size] * weight for i, weight in zip(indices, weights)], dim=1)
k_out = k_out.repeat(1, len(indices)+1, 1).clone()
v_out = v_out.repeat(1, len(indices)+1, 1).clone()
for i in range(batch_size):
if i not in indices:
k_out[i::batch_size, num_tokens:] = k_refs.clone()
v_out[i::batch_size, num_tokens:] = v_refs.clone()
return q_out, k_out, v_out
model_reference.set_model_attn1_patch(reference_apply)
return (model_reference,)
class MultipleReferenceLatent:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"latent": ("LATENT",),
"indices": ("STRING", {"default": "0"}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256}),
}
}
RETURN_TYPES = ("LATENT", )
FUNCTION = "reference_latent"
CATEGORY = CATEGORY_NAME
def reference_latent(self, latent, indices, batch_size):
latent_new = latent.copy()
indices = [int(i) for i in indices.split(",") if i.strip().isdigit()]
sample = latent_new["samples"]
b, _, height, width = sample.shape
assert len(indices) == b
empty_latent = torch.zeros_like(latent["samples"][:1]).repeat(batch_size , 1, 1, 1)
empty_latent[torch.tensor(indices)] = sample
noise_mask = torch.ones(batch_size, 1, height * 8, width * 8).to(sample)
noise_mask[torch.tensor(indices)] = 0.0
latent_new["samples"] = empty_latent
latent_new["noise_mask"] = noise_mask
return (latent_new, )