Typing fixes etc

This commit is contained in:
asagi4
2026-02-04 16:39:50 +02:00
parent 2c8727a75a
commit f98bf25a83
9 changed files with 29 additions and 17 deletions
+4 -2
View File
@@ -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]:
+3 -3
View File
@@ -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)
+1 -1
View File
@@ -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
+2 -1
View File
@@ -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}
+4 -5
View File
@@ -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, {}),
+3 -1
View File
@@ -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
View File
@@ -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
+11 -3
View File
@@ -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)
-1
View File
@@ -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