Make it possible to choose how to deal with long prompts

This commit is contained in:
asagi4
2024-12-10 19:19:12 +02:00
parent f0eb8206b1
commit 529d4e84c2
+35 -13
View File
@@ -204,6 +204,7 @@ class NPNetGoldenNoise:
seed = None
method = "nearest-exact"
strategy = "resize"
olp = "truncate"
@classmethod
def INPUT_TYPES(s):
@@ -215,14 +216,23 @@ class NPNetGoldenNoise:
return {
"required": {
"noise": ("NOISE",),
"prompt": ("CONDITIONING",),
"model": (folder_paths.get_filename_list("npnet"),),
"noise": ("NOISE", {"tooltip": "Connect the output of eg. RandomNoise to this node"}),
"prompt": ("CONDITIONING", {"tooltip": "This is the prompt you want the golden noise for"}),
"model": (
folder_paths.get_filename_list("npnet"),
{"tooltip": "Put your models under models/npnet in your ComfyUI directory"},
),
"device": (["cuda", "cpu"],),
},
"optional": {
"reshape": (["resize", "crop"],),
"method": (["nearest-exact", "bilinear", "area", "bicubic", "bislerp"],),
"reshape": (["resize", "crop"], {"tooltip": "What to do with latents that NPNet can't handle"}),
"reshape_method": (["nearest-exact", "bilinear", "area", "bicubic", "bislerp"],),
"on_long_prompt": (
["truncate", "average", "recurse"],
{
"tooltip": "What to do when the prompt is >77 tokens. 'truncate' will simply cut the prompt, average will apply the model to all prompts and average the result, and 'recurse' will apply npnet recursively for each 77-token chunk"
},
),
},
}
@@ -245,24 +255,36 @@ class NPNetGoldenNoise:
init_noise = self.noise.generate_noise(input_latent).to(self.npnet.device)
cond = self.cond[0].clone().to(self.npnet.device)
if cond.shape[1] != 77:
print("NPNet can't handle conds >77 tokens, running the model individually on each piece")
r = init_noise
for i, cond in enumerate(torch.split(cond, 77, 1)):
print("Applying NPNet to chunk", i + 1)
r = self.npnet(r, cond)
print(f"Prompt has {cond.shape[1]} tokens. NPNet can't handle prompts >77, a workaround will be applied")
if self.olp == "truncate":
print("Truncating prompt to 77 tokens")
cond = cond[:, :77, :]
r = self.npnet(init_noise, cond)
elif self.olp == "recurse":
print("Applying NPNet recursively to all prompt chunks")
r = init_noise
for i, cond in enumerate(torch.split(cond, 77, 1)):
r = self.npnet(r, cond)
else:
print("Averaging NPNet output for each chunk")
r = torch.stack([self.npnet(init_noise, c) for c in torch.split(cond, 77, 1)]).mean(dim=0)
else:
r = self.npnet(init_noise, cond)
return self.reshape(r.to("cpu"), orig_shape)
def doit(self, noise, prompt, model, device, reshape="resize", method="nearest-exact"):
def doit(
self, noise, prompt, model, device, reshape="resize", reshape_method="nearest-exact", on_long_prompt="truncate"
):
model_path = folder_paths.get_full_path("npnet", model)
if self.npnet is None or self.npnet.pretrained_path != model_path:
print("Loading NPNet from", model_path)
self.npnet = NPNet(model_path, device=device)
self.npnet.to(device)
self.method = method
self.method = reshape_method
self.strategy = reshape
self.noise = noise
self.olp = on_long_prompt
self.cond = prompt[0]
return (self,)