From 80a240efee55aca299002f3d7eb76b871a285668 Mon Sep 17 00:00:00 2001 From: Daniel Martinek Date: Wed, 1 Apr 2026 22:57:09 +0200 Subject: [PATCH] AI: add unet hooks --- more_math/SelectiveGuiderMathNode.py | 68 +++++++++++++++++++++++++++- 1 file changed, 67 insertions(+), 1 deletion(-) diff --git a/more_math/SelectiveGuiderMathNode.py b/more_math/SelectiveGuiderMathNode.py index 0b06767..d7471bb 100644 --- a/more_math/SelectiveGuiderMathNode.py +++ b/more_math/SelectiveGuiderMathNode.py @@ -72,7 +72,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): tree = parse_expr(Expression) if isinstance(Expression, str) else Expression MAX_HOOK_INDEX = 999 - hit_flags = {"attn1": False, "attn2": False, "dit": False, "attn_unknown": False} + hit_flags = {"attn1": False, "attn2": False, "dit": False, "attn_unknown": False, "unet": False} def resolve_attn_kind(transformer_options: dict, q: torch.Tensor | None = None, k: torch.Tensor | None = None, v: torch.Tensor | None = None) -> str: raw = transformer_options.get( @@ -147,6 +147,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): "hook_domain": "unknown", "attn_kind": "none", "is_dit": 0.0, + "is_unet_block": 0.0, "is_attn1": 0.0, "is_attn2": 0.0, "block_name": "", @@ -298,11 +299,74 @@ class SelectiveGuiderMathNode(io.ComfyNode): patched_dict["transformer_options"] = topts return patched_dict + def register_unet_blocks(patched_dict, forced_side=None): + topts = patched_dict["transformer_options"].copy() + patches = topts.get("patches", {}).copy() + + def _parse_block(transformer_options: dict, fallback_stage: str): + blk = transformer_options.get("block", None) + if isinstance(blk, (tuple, list)) and len(blk) >= 2: + return str(blk[0]), int(blk[1]) + return fallback_stage, int(transformer_options.get("block_index", -1)) + + def _apply_block_expr(t: torch.Tensor, transformer_options: dict, stage: str, idx: int): + if not match_index(idx, None, stage, transformer_options, "unet_block"): + return t + + co = transformer_options.get("cond_or_uncond", None) + auto_side, auto_idx = side_from_cond_or_uncond(co) + side = forced_side if forced_side is not None else auto_side + idx_side = 0 if side == "positive" else (1 if side == "negative" else auto_idx) + + if not hit_flags["unet"]: + logging.warning(f"[SelectiveGuiderMathNode] UNET BLOCK HIT stage={stage} layer_id={idx} side={side}") + hit_flags["unet"] = True + + meta = { + "hook_kind": "unet_block", + "hook_domain": "diffusion", + "attn_kind": "none", + "is_dit": 0.0, + "is_unet_block": 1.0, + "is_attn1": 0.0, + "is_attn2": 0.0, + "block_name": stage, + "layer_id": float(idx), + "layer": float(idx), + "i": float(idx), + "total_blocks": -1.0, + "has_qkv": 0.0, + "layer_key": f"unet.{stage}.{idx}", + "cond_side": side, + "cond_index": float(idx_side), + "is_positive": 1.0 if side == "positive" else 0.0, + "is_negative": 1.0 if side == "negative" else 0.0, + } + return run_expr(t, meta) + + def input_block_patch(h, transformer_options): + stage, idx = _parse_block(transformer_options, "input") + return _apply_block_expr(h, transformer_options, stage, idx) + + def output_block_patch(h, hsp, transformer_options): + stage, idx = _parse_block(transformer_options, "output") + h2 = _apply_block_expr(h, transformer_options, stage, idx) + return h2, hsp + + patches["input_block_patch"] = list(patches.get("input_block_patch", [])) + [input_block_patch] + patches["output_block_patch"] = list(patches.get("output_block_patch", [])) + [output_block_patch] + + topts["patches"] = patches + patched_dict["transformer_options"] = topts + return patched_dict + def build_transformers_for_side(side_label=None): p = comfy.model_patcher.create_model_options_clone(V.model_options) p.setdefault("transformer_options", {}) if hook_target in ("all", "dit_block"): register_dit(p, forced_side=side_label) + if hook_target in ("all", "unet_block"): + p = register_unet_blocks(p, forced_side=side_label) if hook_target in ("all", "attn1", "attn2"): p = register_attn(p, forced_side=side_label) return p["transformer_options"] @@ -339,6 +403,8 @@ class SelectiveGuiderMathNode(io.ComfyNode): patched.setdefault("transformer_options", {}) if hook_target in ("all", "dit_block"): register_dit(patched, forced_side=None) + if hook_target in ("all", "unet_block"): + patched = register_unet_blocks(patched, forced_side=None) if hook_target in ("all", "attn1", "attn2"): patched = register_attn(patched, forced_side=None) V.model_options = patched