diff --git a/nodes/debug.py b/nodes/debug.py index 89f8849..cdae88d 100644 --- a/nodes/debug.py +++ b/nodes/debug.py @@ -1,33 +1,70 @@ import base64 import io -import json -from pathlib import Path +import textwrap +from collections.abc import Callable +from functools import wraps +from typing import Any, Literal, Protocol, TypedDict, runtime_checkable -import folder_paths import torch +from rich import inspect +from rich.console import Console from ..log import log -from ..utils import tensor2pil +from ..utils import LazyProxyTensor, get_torch_tensor_info, tensor2pil + +try: + import matplotlib.pyplot as plt + import numpy as np + + plt.style.use("dark_background") + MATPLOTLIB_AVAILABLE = True +except ImportError: + MATPLOTLIB_AVAILABLE = False -def get_detailed_type_info(obj): - type_info = [] +# region Decorator +def metadata(**meta_kwargs: Any) -> Callable[[Any], Any]: + """Add metadata to method (`__meta__` dict).""" + + def decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]: + @wraps(func) + def wrapper(*args, **kwargs): + return func(*args, **kwargs) + + wrapper.__meta__ = meta_kwargs + return wrapper + + return decorator + + +# endregion +class UIResult(TypedDict): + kind: Literal["text", "b64_images"] + data: str + + +def indent_results(results: list[UIResult], by: str = " "): + for res in results: + if res["kind"] == "text": + log.debug(f"Indenting: {res['data']}") + res["data"] = textwrap.indent(res["data"], by) + + return results + + +ProcessorResult = list[UIResult] + + +def _get_detailed_type_info(obj) -> str: + type_info: list[str] = [] type_name = type(obj).__name__ type_info.append(f"Type: {type_name}") if isinstance(obj, torch.Tensor): - type_info.extend( - [ - f"Shape: {obj.shape}", - f"Dtype: {obj.dtype}", - f"Device: {obj.device}", - f"Requires grad: {obj.requires_grad}", - f"Stride: {obj.stride()}", - f"Contiguous: {obj.is_contiguous()}", - ] - ) - elif isinstance(obj, (list, tuple)): + return get_torch_tensor_info(obj) + + elif isinstance(obj, list | tuple): type_info.extend( [ f"Length: {len(obj)}", @@ -47,122 +84,184 @@ def get_detailed_type_info(obj): attributes = [attr for attr in dir(obj) if not attr.startswith("_")] type_info.append(f"Attributes: {attributes}") - return type_info + return "\n".join(type_info) + + +def _apply_rich_results(processed, mode="none", title=""): + processing_text = False + acc = "" + reshaped: list[UIResult] = [] + for i in range(len(processed)): + if processed[i]["kind"] == "text": + if not processing_text: + processing_text = True + acc += processed[i]["data"] + "\n" + if len(processed) == (i + 1): + reshaped.append( + UIResult( + kind="text", data=_apply_rich(acc, mode, title=title) + ) + ) + else: + if processing_text: + processing_text = False + reshaped.append( + UIResult( + kind="text", data=_apply_rich(acc, mode, title=title) + ) + ) + acc = "" + reshaped.append(processed[i]) + + return reshaped + # for item in processed: # region processors -def process_tensor(tensor: torch.Tensor, as_type=False): - log.debug(f"Tensor: {tensor.shape}") - - if as_type: - return { - "text": [f"Tensor of shape {tensor.shape} of type {tensor.dtype}"] - } - - is_mask = len(tensor.shape) == 3 - - if is_mask: - tensor = tensor.unsqueeze(-1).repeat(1, 1, 1, 3) - - image = tensor2pil(tensor) - b64_imgs = [] - for im in image: - if is_mask: - im = im.convert("L") - - buffered = io.BytesIO() - im.save(buffered, format="PNG") - b64_imgs.append( - "data:image/png;base64," - + base64.b64encode(buffered.getvalue()).decode("utf-8") +def _apply_rich( + formatted: str | list[str], rich_mode: str | None = None, *, title="" +) -> str: + if rich_mode is None: + return ( + formatted if isinstance(formatted, str) else "\n".join(formatted) ) - return {"b64_images": b64_imgs} + from rich.console import Console + console = Console(record=True) -def process_list(anything, as_type=False): - text = [] - if not anything: - return {"text": []} - - if as_type: - type_info = get_detailed_type_info(anything) - type_info.extend(get_detailed_type_info(anything[0])) - return {"text": type_info} - - first_element = anything[0] - if ( - isinstance(first_element, list) - and first_element - and isinstance(first_element[0], torch.Tensor) - ): - text.append( - "List of List of Tensors: " - f"{first_element[0].shape} (x{len(anything)})" - ) - - elif isinstance(first_element, torch.Tensor): - text.append( - f"List of Tensors: {first_element.shape} (x{len(anything)})" - ) + if isinstance(formatted, list): + for line in formatted: + console.print(line) else: - text.append(f"Array ({len(anything)}): {anything}") + console.print(formatted) - return {"text": text} + CSV_CODE_FORMAT = """ + + + + + + + + + {lines} + + + {chrome} + + {backgrounds} + + {matrix} + + + +""" + + if rich_mode == "svg-window": + return console.export_svg(title=title, code_format=CSV_CODE_FORMAT) + elif rich_mode == "svg": + return console.export_svg( + title=title, + code_format=CSV_CODE_FORMAT.replace("{chrome}", ""), ) - text.append( - f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}" + elif rich_mode == "html": + CONSOLE_HTML_FORMAT = textwrap.dedent(""" +
+ {code} +
+ """).strip() + + import rich.terminal_theme + + return console.export_html( + inline_styles=True, + code_format=CONSOLE_HTML_FORMAT, + theme=rich.terminal_theme.MONOKAI, ) - else: - log.debug(f"Unhandled dict: {anything.keys()}") - text.append(json.dumps(anything, indent=2)) - - return {"text": text} - - -def process_bool(anything, as_type=False): - return {"text": ["True" if anything else "False"]} - - -def process_text(anything, as_type=False): - if as_type: - return {"text": get_detailed_type_info(anything)} - - return {"text": [str(anything)]} + log.error(f"Unknown rich mode: {rich_mode}") + return formatted if isinstance(formatted, str) else "\n".join(formatted) # endregion -class MTB_Debug: - """Experimental node to debug any Comfy values. +# region conditions - support for more types and widgets is planned. - """ + +# those are pretty dumb there is now probably a better way.. +def is_condition(item): + return ( + isinstance(item, list) + and all(isinstance(i, list) for i in item) + and isinstance(item[0][0], torch.Tensor) + ) + + +# endregion + +RICH_MODE = Literal["none", "html", "svg", "svg-window"] + + +@runtime_checkable +class Processor(Protocol): + """Generic protocol for processor functions.""" + + def __call__( + self, item: Any, *, as_type: bool = False, deep: bool = False + ) -> ProcessorResult: ... + + +class MTB_Debug: + """A debug node.""" @classmethod def INPUT_TYPES(cls): return { "required": {"output_to_console": ("BOOLEAN", {"default": False})}, - "optional": {"as_detailed_types": ("BOOLEAN", {"default": False})}, + "optional": { + "as_detailed_types": ("BOOLEAN", {"default": False}), + "deep_inspect": ("BOOLEAN", {"default": False}), + "rich_mode": ( + ("none", "html", "svg", "svg-window"), + {"default": "none"}, + ), + }, } RETURN_TYPES = () @@ -170,99 +269,426 @@ class MTB_Debug: CATEGORY = "mtb/debug" OUTPUT_NODE = True + _processors: dict[type, Processor] + + def __init__(self): + self._condition_processors = {is_condition: self._process_condition} + self._class_name_processors = { + "CLIP": self._process_clip, + "VAE": self._process_vae, + } + self._processors = { + torch.nn.Module: self._process_module, + torch.Tensor: self._process_tensor, + LazyProxyTensor: self._process_repr, + list: self._process_container, + tuple: self._process_container, + dict: self._process_dict, + bool: self._process_bool, + str: self._process_primitive, + int: self._process_primitive, + float: self._process_primitive, + type(None): self._process_primitive, + } + + # - Dispatchers ------------------------------------------------------------ + def _dispatch_processor( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + """Find and calls the appropriate processor for the given item.""" + # first conditions + for c, process in self._condition_processors.items(): + if c(item): + return process(item, as_type=as_type, deep=deep) + + # named class + class_name = type(item).__name__ + if class_name in self._class_name_processors: + return self._class_name_processors[class_name]( + item, as_type=as_type, deep=deep + ) + + # type based or unknown + processor = self._processors.get(type(item), self._process_unknown) + res = processor(item, as_type=as_type, deep=deep) + + return res + def do_debug( - self, output_to_console: bool, as_detailed_types: bool, **kwargs + self, + **kwargs, ): output = {"ui": {"items": []}} - if output_to_console: - for k, v in kwargs.items(): - log.info(f"{k}: {v}") + settings = {k: kwargs.pop(k) for k in self.INPUT_TYPES()["optional"]} + output_to_console = kwargs.pop("output_to_console") + as_type = settings.get("as_detailed_types", False) + deep = settings.get("deep_inspect", False) + rich_mode = settings.get("rich_mode", "none") - for input_name, anything in kwargs.items(): - processor = processors.get(type(anything), process_text) + for input_name, item in kwargs.items(): + processed = self._dispatch_processor( + item, as_type=as_type, deep=deep + ) + if processed is None: + continue - processed = processor(anything, as_detailed_types) + if rich_mode != "none": + title = f"{input_name} ({type(item).__name__})" + processed = _apply_rich_results(processed, rich_mode, title) - item = { - "input": input_name, - **processed, - } - output["ui"]["items"].append(item) + if output_to_console: + log.info(f"- Input '{input_name}':") + for p in processed: + if p["kind"] == "text": + log.info(f" {p['data']}") + if p["kind"] == "b64_image": + log.info(f" (contains {len(p['data'])} images)") + output["ui"]["items"].append( + {"input": input_name, "items": processed} + ) return output + def _process_unknown( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + console = Console( + record=True, + width=120, + ) -class MTB_SaveTensors: - """Save torch tensors (image, mask or latent) to disk. + console.print(f"Generic {type(item).__name__}", emoji=True) + if as_type: + inspect(item, console=console, all=deep, methods=deep, docs=deep) + else: + console.print(item, emoji=True) - useful to debug things outside comfy. - """ + text_output = console.export_text(clear=True) - def __init__(self): - self.output_dir = folder_paths.get_output_directory() - self.type = "mtb/debug" + return [UIResult(kind="text", data=text_output.strip())] - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "filename_prefix": ("STRING", {"default": "ComfyPickle"}), - }, - "optional": { - "image": ("IMAGE",), - "mask": ("MASK",), - "latent": ("LATENT",), - }, - } + def _process_repr( + self, item: Any, as_type=False, deep=False + ) -> ProcessorResult: + return [{"kind": "text", "data": item.__repr__()}] - FUNCTION = "save" - OUTPUT_NODE = True - RETURN_TYPES = () - CATEGORY = "mtb/debug" + def _process_primitive( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + if as_type: + return self._process_unknown(item, as_type=as_type, deep=deep) - def save( - self, - filename_prefix, - image: torch.Tensor | None = None, - mask: torch.Tensor | None = None, - latent: torch.Tensor | None = None, - ): - ( - full_output_folder, - filename, - counter, - subfolder, - filename_prefix, - ) = folder_paths.get_save_image_path(filename_prefix, self.output_dir) - full_output_folder = Path(full_output_folder) - if image is not None: - image_file = f"{filename}_image_{counter:05}.pt" - torch.save(image, full_output_folder / image_file) - # np.save(full_output_folder/ image_file, image.cpu().numpy()) + return [UIResult(kind="text", data=str(item))] - if mask is not None: - mask_file = f"{filename}_mask_{counter:05}.pt" - torch.save(mask, full_output_folder / mask_file) - # np.save(full_output_folder/ mask_file, mask.cpu().numpy()) + def _process_bool( + self, item: bool, *, as_type=False, deep=False + ) -> ProcessorResult: # noqa: FBT001 + return [{"kind": "text", "data": "True" if item else "False"}] - if latent is not None: - # for latent we must use pickle - latent_file = f"{filename}_latent_{counter:05}.pt" - torch.save(latent, full_output_folder / latent_file) - # pickle.dump(latent, open(full_output_folder/ latent_file, "wb")) + def _process_clip( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + try: + clip_model = getattr(item, "cond_stage_model", None) + tokenizer = getattr(item, "tokenizer", None) - # np.save(full_output_folder / latent_file, - # latent[""].cpu().numpy()) + text = [UIResult(kind="text", data="CLIP")] + if clip_model: + text.append(UIResult(kind="text", data="CLIP Model:")) + model_summary = self._process_module( + clip_model, as_type=as_type + ) + if model_summary: + text.extend(indent_results(model_summary, " ")) + else: + text.append( + UIResult( + kind="text", + data="[error] failed to get informations about clip model", + ) + ) - return f"{filename_prefix}_{counter:05}" + if tokenizer: + text.append(UIResult(kind="text", data="Tokenizer:")) + vocab_size = getattr(tokenizer, "vocab_size", "N/A") + text.append( + UIResult( + kind="text", + data=f" Class: {type(tokenizer).__name__}\n Vocab Size: {vocab_size}", + ) + ) + + return text + + except Exception as e: + log.error(f"Failed to process CLIP object: {e}") + return self._process_unknown(item, as_type=as_type, deep=deep) + + def _process_condition( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + count = len(item) + result = [UIResult(kind="text", data=f"Conditions: {count}")] + + for cond in item: + result.extend(self._preview_conditioning_tensor(cond[0])) + + return result + + def _process_vae( + self, item: Any, *, as_type=False, deep=False + ) -> ProcessorResult: + try: + vae_model = getattr( + item, "first_stage_model", getattr(item, "vae", item) + ) + text = [ + UIResult(kind="text", data="VAE"), + UIResult(kind="text", data="Internal Model:"), + ] + + model_summary = self._process_module( + vae_model, as_type=as_type, deep=deep + ) + text.extend(indent_results(model_summary, " ")) + + return text + except Exception as e: + log.error(f"Failed to process VAE object: {e}") + return self._process_unknown(item, as_type=as_type, deep=deep) + + def _process_module( + self, item: torch.nn.Module, *, as_type=False, deep=False + ) -> ProcessorResult: + if as_type and deep: + return self._process_unknown(item, as_type=as_type, deep=deep) + + total_params = sum(p.numel() for p in item.parameters()) + trainable_params = sum( + p.numel() for p in item.parameters() if p.requires_grad + ) + try: + device = next(item.parameters()).device + except StopIteration: + device = "cpu (no parameters)" + + train_percent = ( + f"{trainable_params / total_params:.2%}" + if total_params > 0 + else "0.00%" + ) + + text = [ + f"Model: {type(item).__name__} on {device}", + textwrap.dedent(f""" + - Parameters: {total_params:,} + - Trainable: {trainable_params:,} ({train_percent}) + """).strip(), + ] + return [{"kind": "text", "data": d} for d in text] + + def _process_tensor( + self, item: torch.Tensor, *, as_type=False, deep=False + ) -> ProcessorResult: + is_latent = item.ndim == 4 and item.shape[1] == 4 + is_image = ( + not is_latent and item.ndim == 4 and item.shape[3] in [1, 3, 4] + ) + is_conditioning = item.ndim == 3 and item.shape[2] in [ + 768, + 1024, + 1152, + 1280, + 2048, + 4096, + ] + is_mask = (item.ndim == 2) or (item.ndim == 3 and not is_conditioning) + + if as_type: + type_name = "Unknown Tensor" + if is_latent: + type_name = "Latent Tensor" + elif is_image: + type_name = "Image Tensor" + elif is_conditioning: + type_name = "CLIP Conditioning Tensor" + elif is_mask: + type_name = "Mask Tensor" + return [ + { + "kind": "text", + "data": get_torch_tensor_info(item, name=type_name), + } + ] + + if is_image or is_mask: + return self._render_image_tensor(item) + if is_latent: + return self._preview_latent_tensor(item) + if is_conditioning: + return self._preview_conditioning_tensor(item) + return self._process_unknown(item, as_type=as_type, deep=deep) + + def _visualize_tensor_heatmap( + self, tensor_2d: torch.Tensor, title: str + ) -> str | None: + if not MATPLOTLIB_AVAILABLE: + log.warning("Matplotlib not found. Skipping tensor visualization.") + return None + if tensor_2d.ndim != 2: + log.warning( + f"Cannot visualize tensor with {tensor_2d.ndim} dimensions. Requires 2." + ) + return None + + fig, ax = plt.subplots(figsize=(6, 4), dpi=100) + im = ax.imshow(tensor_2d.cpu().numpy(), cmap="viridis", aspect="auto") + fig.colorbar(im, ax=ax) + ax.set_title(title) + fig.tight_layout() + + buf = io.BytesIO() + fig.savefig(buf, format="png", bbox_inches="tight", pad_inches=0.1) + plt.close(fig) + buf.seek(0) + return "data:image/png;base64," + base64.b64encode(buf.read()).decode( + "utf-8" + ) + + def _render_image_tensor(self, item: torch.Tensor) -> ProcessorResult: + is_mask = (item.ndim == 2) or (item.ndim == 3 and item.shape[-1] != 3) + img_tensor = ( + item.unsqueeze(0) if item.ndim == 3 and not is_mask else item + ) + img_tensor = item.unsqueeze(0) if item.ndim == 2 else img_tensor + + images = tensor2pil(img_tensor) + b64_imgs = [] + for im in images: + if is_mask: + im = im.convert("L") + buffered = io.BytesIO() + im.save(buffered, format="PNG") + b64_imgs.append( + "data:image/png;base64," + + base64.b64encode(buffered.getvalue()).decode("utf-8") + ) + return [UIResult(kind="b64_images", data=b64_imgs)] + + def _preview_latent_tensor(self, item: torch.Tensor) -> ProcessorResult: + is_empty = "(empty)" if torch.count_nonzero(item) == 0 else "" + stats = [ + f"Min: {item.min():.4f}", + f"Max: {item.max():.4f}", + f"Mean: {item.mean():.4f}", + ] + text = [ + get_torch_tensor_info(item, name="Latent Tensor"), + is_empty, + ] + stats + + result = [UIResult(kind="text", data=t) for t in text] + vis_tensor = item[0].mean(dim=0) + heatmap_b64 = self._visualize_tensor_heatmap( + vis_tensor, "Latent Energy (Channel Mean)" + ) + if heatmap_b64: + result.append(UIResult(kind="b64_images", data=[heatmap_b64])) + return result + + def _preview_conditioning_tensor( + self, item: torch.Tensor + ) -> ProcessorResult: + _batch, tokens, embed_dim = item.shape + text = [ + get_torch_tensor_info(item, name="CLIP Conditioning Tensor"), + f"Token Count: {tokens}", + f"Embedding Dim: {embed_dim}", + ] + + result = [UIResult(kind="text", data=d) for d in text] + heatmap_b64 = self._visualize_tensor_heatmap( + item[0], "Token Embeddings (approx)" + ) + if heatmap_b64: + result.append(UIResult(kind="b64_images", data=[heatmap_b64])) + return result + + def _process_container( + self, item: list | tuple, *, as_type=False, deep=False + ) -> ProcessorResult: + if not item: + return [UIResult(kind="text", data=f"Empty {type(item).__name__}")] + + container_type = type(item).__name__ + element_type = type(item[0]).__name__ + + all_match = all(type(i) is type(item[0]) for i in item) + + result = [ + UIResult( + kind="text", + data=f"{container_type} of {len(item)} x {element_type}", + ), + UIResult(kind="text", data=f"(mixed types: {not all_match})"), + ] + + if not as_type or (as_type and deep): + for i, sub_item in enumerate(item): + res = self._dispatch_processor( + sub_item, as_type=as_type, deep=deep + ) + if res: + text = res[0].get("data", "Unknown") + res[0]["data"] = f"[{i}]: {text}" + + result.extend(res) + + return result + + first_item_result = self._dispatch_processor( + item[0], as_type=as_type, deep=deep + ) + if not first_item_result: + return result + + return ( + result + + [UIResult(kind="text", data="Preview of first element:")] + + indent_results(first_item_result, " - ") + ) + + def _process_dict( + self, item: dict, *, as_type=False, deep=False + ) -> ProcessorResult: + if "pooled_output" in item and isinstance( + item["pooled_output"], torch.Tensor + ): + return self._dispatch_processor( + item["pooled_output"], as_type=as_type, deep=deep + ) + + if "samples" in item and isinstance(item.get("samples"), torch.Tensor): + return self._dispatch_processor( + item["samples"], as_type=as_type, deep=deep + ) + + if "waveform" in item and isinstance( + item.get("waveform"), torch.Tensor + ): + waveform = item["waveform"] + is_empty = "(empty) " if torch.count_nonzero(waveform) == 0 else "" + text = textwrap.dedent(f""" + Audio Waveform: {waveform.shape}{is_empty} + Sample Rate: {item.get("sample_rate", "N/A")} + """).strip() + return [{"kind": "text", "data": text}] + + log.debug( + f"Processing generic dict with rich inspector: {item.keys()}" + ) + return self._process_unknown(item, as_type=as_type, deep=deep) -processors = { - torch.Tensor: process_tensor, - list: process_list, - dict: process_dict, - bool: process_bool, -} - -__nodes__ = [MTB_Debug, MTB_SaveTensors] +__nodes__ = [MTB_Debug] diff --git a/web/debug.js b/web/debug.js index daa1b66..2f6a355 100644 --- a/web/debug.js +++ b/web/debug.js @@ -132,8 +132,9 @@ app.registerExtension({ let tgt_len = this.widgets.length for (let i = 0; i < this.widgets.length; i++) { if ( - this.widgets[i].name !== 'output_to_console' && - this.widgets[i].name !== 'as_detailed_types' + !['output_to_console', 'as_detailed_types', 'rich_mode'].includes( + this.widgets[i].name, + ) ) { this.widgets[i].onRemove?.() this.widgets[i].onRemoved?.()