diff --git a/more_math/SelectiveGuiderMathNode.py b/more_math/SelectiveGuiderMathNode.py index f132e36..3ba9a1a 100644 --- a/more_math/SelectiveGuiderMathNode.py +++ b/more_math/SelectiveGuiderMathNode.py @@ -10,7 +10,6 @@ from .Parser.UnifiedMathVisitor import UnifiedMathVisitor from .helper_functions import checkLazyNew, generate_dim_variables, parse_expr, as_tensor from .Stack import MrmthStack from .ParseTree import MrmthParseTree -from comfy.ldm.modules.attention import optimized_attention class SelectiveGuiderMathNode(io.ComfyNode): @@ -42,7 +41,8 @@ class SelectiveGuiderMathNode(io.ComfyNode): default="all" ), MrmthStack.Input(id="stack", tooltip="Access stack between nodes", optional=True), - io.Boolean.Input(id="remember_stack", default=False, display_name="Remember stack across batch") + io.Boolean.Input(id="remember_stack", default=False, display_name="Remember stack across batch"), + io.Boolean.Input(id="debug_hooks", default=False, display_name="Debug hook logs"), ], outputs=[io.Guider.Output(), MrmthStack.Output()], ) @@ -56,6 +56,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): layer_x=0, hook_target="all", remember_stack=False, + debug_hooks=False, stack={}, ): return checkLazyNew(Expression, {}, F) @@ -69,6 +70,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): layer_x=0, hook_target="all", remember_stack=False, + debug_hooks=False, stack=None, ): if V is None or not hasattr(V, "model_options"): @@ -77,7 +79,6 @@ class SelectiveGuiderMathNode(io.ComfyNode): stack = stack if remember_stack else (copy.deepcopy(stack) if stack is not None else {}) tree = parse_expr(Expression) if isinstance(Expression, str) else Expression - MAX_HOOK_INDEX = 999 hit_flags = { "attn1": False, "attn2": False, @@ -88,6 +89,26 @@ class SelectiveGuiderMathNode(io.ComfyNode): "unet": False, } + debug_hook_events = [] + + def dbg(msg: str): + if debug_hooks: + logging.warning(f"[SelectiveGuiderMathNode] {msg}") + + def record_hook(kind: str, side: str, layer_id=None, stage=None, key=None): + if not debug_hooks: + return + parts = [f"hook={kind}", f"side={side}"] + if stage is not None: + parts.append(f"stage={stage}") + if layer_id is not None: + parts.append(f"layer_id={layer_id}") + if key is not None: + parts.append(f"key={key}") + line = " ".join(parts) + debug_hook_events.append(line) + dbg(f"HOOK {line}") + def resolve_attn_kind(transformer_options: dict, q: torch.Tensor | None = None, k: torch.Tensor | None = None, v: torch.Tensor | None = None) -> str: # Flux / MM-DiT btype = transformer_options.get("block_type", None) @@ -280,7 +301,8 @@ class SelectiveGuiderMathNode(io.ComfyNode): visitor = UnifiedMathVisitor(variables, inp.shape, inp.device, state_storage=stack) out = visitor.visit(tree) - return as_tensor(out, inp.shape).to(inp.device) + # preserve dtype/device (float8-safe vs float16 promotion side effects) + return as_tensor(out, inp.shape).to(device=inp.device, dtype=inp.dtype) def register_dit(patched_dict, forced_side=None): topts = patched_dict["transformer_options"] @@ -292,8 +314,8 @@ class SelectiveGuiderMathNode(io.ComfyNode): if not match_index(idx, total_blocks, "dit", args.get("transformer_options", {}), "dit_block"): return extra_args["original_block"](args) if not hit_flags["dit"]: - logging.warning(f"[SelectiveGuiderMathNode] DIT HIT idx={idx} total={total_blocks}") - hit_flags["dit"] = True + dbg(f"DIT HIT idx={idx} total={total_blocks}") + hit_flags["dit"] = True btype = args.get("transformer_options", {}).get("block_type", "dit") co = args.get("transformer_options", {}).get("cond_or_uncond", None) auto_side, auto_idx = side_from_cond_or_uncond(co) @@ -374,12 +396,11 @@ class SelectiveGuiderMathNode(io.ComfyNode): idx_side = 0 if side == "positive" else (1 if side == "negative" else auto_idx) if not hit_flags.get(attn_kind, False): - logging.warning(f"[SelectiveGuiderMathNode] {attn_kind.upper()} HIT stage={stage} layer_id={layer_id} side={side}") + dbg(f"{attn_kind.upper()} HIT stage={stage} layer_id={layer_id} side={side}") hit_flags[attn_kind] = True act_shape = transformer_options.get("activations_shape", []) - if not isinstance(act_shape, (list, tuple)): - act_shape = [] + meta = { "hook_kind": attn_kind, @@ -399,7 +420,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): "k": k, "v": v, "heads": float(heads), - "dim_head": float(q.shape[-1] // heads) if heads > 0 else 0.0, + "dim_head": float(q.shape[-1] // heads) if heads > 0 : 0.0, "layer_key": f"{stage}.{layer_id}.{attn_kind}.{t_index}", "cond_side": side, "cond_index": float(idx_side), @@ -440,7 +461,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): 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}") + dbg(f"UNET BLOCK HIT stage={stage} layer_id={idx} side={side}") hit_flags["unet"] = True meta = { @@ -574,6 +595,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): V.original_conds[side_name] = new_conds # po parse_expr(...) + if hasattr(V, "original_conds"): def make_shallow_cond_copy(conds_dict): new_dict = {} @@ -581,7 +603,7 @@ class SelectiveGuiderMathNode(io.ComfyNode): # Kondice obvykle obsahují fragmentová pole (slovníky nebo listy) new_dict[k] = [item.copy() if hasattr(item, "copy") else copy.copy(item) for item in v_list] return new_dict - + if not hasattr(V, "_mrmth_base_original_conds"): V._mrmth_base_original_conds = make_shallow_cond_copy(V.original_conds) else: @@ -590,11 +612,11 @@ class SelectiveGuiderMathNode(io.ComfyNode): if hasattr(V, "original_conds"): attach_side_hook("positive") attach_side_hook("negative") - logging.warning(f"[SelectiveGuiderMathNode] has original_conds: {list(V.original_conds.keys())}") + dbg(f"has original_conds: {list(V.original_conds.keys())}") for side in ("positive", "negative"): if side in V.original_conds: cnt = sum(1 for c in V.original_conds[side] if c.get("hooks", None) is not None) - logging.warning(f"[SelectiveGuiderMathNode] side={side} conds_with_hooks={cnt}/{len(V.original_conds[side])}") + dbg(f"side={side} conds_with_hooks={cnt}/{len(V.original_conds[side])}") else: patched = comfy.model_patcher.create_model_options_clone(V.model_options) patched.setdefault("transformer_options", {}) @@ -614,17 +636,17 @@ class SelectiveGuiderMathNode(io.ComfyNode): reg_len = len(registered) if registered is not None else 0 in_len = len(hooks) if hooks is not None else 0 tr_count = len(registered.get_type(comfy.hooks.EnumHookType.TransformerOptions)) if registered is not None else 0 - logging.warning(f"[SelectiveGuiderMathNode] REGISTER hooks_in={in_len} registered={reg_len} transformer_registered={tr_count}") + dbg(f"REGISTER hooks_in={in_len} registered={reg_len} transformer_registered={tr_count}") if model_options is not None: rh = model_options.get("registered_hooks", None) - logging.warning(f"[SelectiveGuiderMathNode] model_options.registered_hooks={'set' if rh is not None else 'None'}") + dbg(f"model_options.registered_hooks={'set' if rh is not None else 'None'}") except Exception as e: logging.warning(f"[SelectiveGuiderMathNode] REGISTER debug error: {e}") def _on_apply_hooks(model_patcher, hooks): try: hl = len(hooks) if hooks is not None else 0 - logging.warning(f"[SelectiveGuiderMathNode] APPLY hooks_len={hl}") + dbg(f"APPLY hooks_len={hl}") except Exception as e: logging.warning(f"[SelectiveGuiderMathNode] APPLY debug error: {e}") @@ -641,4 +663,14 @@ class SelectiveGuiderMathNode(io.ComfyNode): V._mrmth_debug_callbacks_added = True stack = stack if remember_stack else copy.deepcopy(stack) - return (V, stack) \ No newline at end of file + result = (V, stack) + + # 6) Before return add compact summary: + + if debug_hooks: + uniq = sorted(set(debug_hook_events)) + dbg(f"HOOK SUMMARY total={len(debug_hook_events)} unique={len(uniq)}") + for item in uniq: + dbg(f"HOOK SUMMARY ITEM {item}") + + return result \ No newline at end of file