Typing fixes etc
This commit is contained in:
@@ -197,6 +197,7 @@ class AdvancedEncoder:
|
||||
def add_encoder(cls, name, fn):
|
||||
cls.STYLES[name] = fn
|
||||
|
||||
@classmethod
|
||||
def add_normalization_op(cls, name, fn):
|
||||
cls.NORMALIZATION_OPS[name] = fn
|
||||
|
||||
@@ -296,7 +297,7 @@ class AdvancedEncoder:
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
|
||||
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
|
||||
pooled = pooled_base
|
||||
if pooled is not None and self.max_length:
|
||||
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
@@ -328,12 +329,13 @@ class AdvancedEncoder:
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = None
|
||||
if pooled_base is not None and self.max_length:
|
||||
pooled = embs[0, self.max_length - 1 : self.max_length, :]
|
||||
pooled_start = pooled_base.expand(len(ws), -1)
|
||||
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
|
||||
pooled = (pooled - pooled_start) * (ws - 1)
|
||||
pooled = pooled.mean(axis=0, keepdim=True)
|
||||
pooled = pooled.mean(dim=0, keepdim=True)
|
||||
pooled = pooled_base + pooled
|
||||
|
||||
if embs.shape[0] != masks.shape[0]:
|
||||
|
||||
@@ -71,14 +71,13 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
}
|
||||
}
|
||||
self.has_negpip = False
|
||||
|
||||
# calculate later. All clones must refer to the same kv dict
|
||||
self.kv = {"k": None, "v": None}
|
||||
self.kv = {}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self.num_conds = len(conds) + 1
|
||||
self.base_strength = base_cond[1].get("strength", 1.0)
|
||||
self.strengths = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
|
||||
base_mask = base_cond[1].get("mask", None)
|
||||
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
|
||||
@@ -215,6 +214,7 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
assert self.num_conds is not None, "this is a bug"
|
||||
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
|
||||
|
||||
q = torch.cat(qs, dim=0)
|
||||
|
||||
@@ -215,7 +215,7 @@ def encode_regions(clip_regions, encode, tokenizer):
|
||||
region_emb *= region_masking
|
||||
|
||||
region_embeddings.append(region_emb)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
|
||||
|
||||
embeddings_final_mask = torch.tensor(
|
||||
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .prompts import encode_prompt
|
||||
|
||||
@@ -40,7 +41,7 @@ class PCTextEncode:
|
||||
DESCRIPTION = "Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling"
|
||||
|
||||
def apply(self, clip, text):
|
||||
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
|
||||
return PCTextEncodeWithRange().apply(clip, text, 0.0, 1.0)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
|
||||
import comfy.hooks
|
||||
@@ -38,7 +39,6 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
all_hooks = []
|
||||
|
||||
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
|
||||
nonlocal lora_cache
|
||||
hooks = []
|
||||
hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||
for path, info in loras.items():
|
||||
@@ -53,7 +53,8 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
|
||||
)
|
||||
# Set hook_ref so that identical hooks compare equal
|
||||
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
new_hook.hooks[0].hook_ref = ref # pyright: ignore[reportAttributeAccessIssue]
|
||||
hooks.append(new_hook)
|
||||
if start_pct > 0.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
||||
@@ -74,8 +75,6 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
all_hooks.append(hook)
|
||||
start_pct = end_pct
|
||||
|
||||
del lora_cache
|
||||
|
||||
all_hooks = [x for x in all_hooks if x]
|
||||
|
||||
if all_hooks:
|
||||
@@ -85,7 +84,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
|
||||
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> InputTypeDict:
|
||||
def INPUT_TYPES(s) -> InputTypeDict:
|
||||
return {
|
||||
"required": {
|
||||
"positive": (IO.CONDITIONING, {}),
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
from __future__ import annotations
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
@@ -167,7 +169,7 @@ class PCLazyLoraLoaderAdvanced:
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
|
||||
RETURN_TYPES: tuple[str, ...] = ("MODEL", "CLIP", "HOOKS")
|
||||
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules, expand_macros
|
||||
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
import lark
|
||||
import logging
|
||||
from math import ceil
|
||||
@@ -100,7 +101,15 @@ class CutTransform(lark.Transformer):
|
||||
def cut(self, args):
|
||||
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
|
||||
|
||||
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
|
||||
# prompts and cutouts are always sequences of str
|
||||
return (
|
||||
"".join(flatten(prompt)), # pyright: ignore
|
||||
"".join(flatten(cutout)), # pyright: ignore
|
||||
weight,
|
||||
strict_mask,
|
||||
start_from_masked,
|
||||
mask_token,
|
||||
) # pyright: ignore
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
@@ -301,8 +310,7 @@ def at_step(step, filters, tree):
|
||||
return name, params, lbw
|
||||
|
||||
def __default__(self, data, children, meta):
|
||||
for child in children:
|
||||
yield child
|
||||
return children
|
||||
|
||||
return AtStep().transform(tree)
|
||||
|
||||
|
||||
@@ -30,7 +30,6 @@ def compare_hookgroup_mask(h1, h2):
|
||||
class TestEncode(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
global clips
|
||||
print("Loading ComfyUI")
|
||||
from comfy.sd import load_clip
|
||||
from pathlib import Path
|
||||
|
||||
Reference in New Issue
Block a user