From cf7a9c41e81e8dd461ab9dfa3c05bb8e2cdf2a67 Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Sat, 15 Feb 2025 23:44:34 +0100 Subject: [PATCH] =?UTF-8?q?feat:=20=E2=9C=A8=20improve=20the=20debug=20nod?= =?UTF-8?q?e?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - preserve input order - new "as_detailed_type" option - support mask preview - improved styling a bit for readibility --- nodes/debug.py | 125 +++++++++++++++++++++++++++++++++-------- web/debug.js | 147 ++++++++++++++++++++++++++++++++++--------------- 2 files changed, 204 insertions(+), 68 deletions(-) diff --git a/nodes/debug.py b/nodes/debug.py index 5deac9f..89f8849 100644 --- a/nodes/debug.py +++ b/nodes/debug.py @@ -2,7 +2,6 @@ import base64 import io import json from pathlib import Path -from typing import Optional import folder_paths import torch @@ -11,13 +10,66 @@ from ..log import log from ..utils import tensor2pil +def get_detailed_type_info(obj): + type_info = [] + + 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)): + type_info.extend( + [ + f"Length: {len(obj)}", + f"Container type: {type_name}", + ] + ) + if obj: + type_info.append(f"Element type: {type(obj[0]).__name__}") + elif isinstance(obj, dict): + type_info.extend( + [ + f"Length: {len(obj)}", + f"Keys: {list(obj.keys())}", + ] + ) + elif hasattr(obj, "__dict__"): + attributes = [attr for attr in dir(obj) if not attr.startswith("_")] + type_info.append(f"Attributes: {attributes}") + + return type_info + + # region processors -def process_tensor(tensor): +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( @@ -28,11 +80,16 @@ def process_tensor(tensor): return {"b64_images": b64_imgs} -def process_list(anything): +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) @@ -54,25 +111,41 @@ def process_list(anything): return {"text": text} -def process_dict(anything): +def process_dict(anything, as_type=False): text = [] + if as_type: + return {"text": get_detailed_type_info(anything)} + if "samples" in anything: is_empty = ( "(empty)" if torch.count_nonzero(anything["samples"]) == 0 else "" ) text.append(f"Latent Samples: {anything['samples'].shape} {is_empty}") + elif "waveform" in anything: + is_empty = ( + "(empty) " if torch.count_nonzero(anything["samples"]) == 0 else "" + ) + + text.append( + f"Audio Samples: {anything['waveform'].shape}{is_empty} | sample rate {anything['sample_rate']}" + ) + else: + log.debug(f"Unhandled dict: {anything.keys()}") text.append(json.dumps(anything, indent=2)) return {"text": text} -def process_bool(anything): +def process_bool(anything, as_type=False): return {"text": ["True" if anything else "False"]} -def process_text(anything): +def process_text(anything, as_type=False): + if as_type: + return {"text": get_detailed_type_info(anything)} + return {"text": [str(anything)]} @@ -89,6 +162,7 @@ class MTB_Debug: def INPUT_TYPES(cls): return { "required": {"output_to_console": ("BOOLEAN", {"default": False})}, + "optional": {"as_detailed_types": ("BOOLEAN", {"default": False})}, } RETURN_TYPES = () @@ -96,29 +170,25 @@ class MTB_Debug: CATEGORY = "mtb/debug" OUTPUT_NODE = True - def do_debug(self, output_to_console: bool, **kwargs): - output = { - "ui": {"b64_images": [], "text": []}, - # "result": ("A"), - } + def do_debug( + self, output_to_console: bool, as_detailed_types: bool, **kwargs + ): + output = {"ui": {"items": []}} - processors = { - torch.Tensor: process_tensor, - list: process_list, - dict: process_dict, - bool: process_bool, - } if output_to_console: for k, v in kwargs.items(): log.info(f"{k}: {v}") - for anything in kwargs.values(): + for input_name, anything in kwargs.items(): processor = processors.get(type(anything), process_text) - processed_data = processor(anything) + processed = processor(anything, as_detailed_types) - for ui_key, ui_value in processed_data.items(): - output["ui"][ui_key].extend(ui_value) + item = { + "input": input_name, + **processed, + } + output["ui"]["items"].append(item) return output @@ -154,9 +224,9 @@ class MTB_SaveTensors: def save( self, filename_prefix, - image: Optional[torch.Tensor] = None, - mask: Optional[torch.Tensor] = None, - latent: Optional[torch.Tensor] = None, + image: torch.Tensor | None = None, + mask: torch.Tensor | None = None, + latent: torch.Tensor | None = None, ): ( full_output_folder, @@ -188,4 +258,11 @@ class MTB_SaveTensors: return f"{filename_prefix}_{counter:05}" +processors = { + torch.Tensor: process_tensor, + list: process_list, + dict: process_dict, + bool: process_bool, +} + __nodes__ = [MTB_Debug, MTB_SaveTensors] diff --git a/web/debug.js b/web/debug.js index 102b5b2..511de6f 100644 --- a/web/debug.js +++ b/web/debug.js @@ -11,13 +11,9 @@ /// import { app } from '../../scripts/app.js' - import * as shared from './comfy_shared.js' -import { MtbWidgets } from './mtb_widgets.js' import * as mtb_ui from './mtb_ui.js' -// TODO: respect inputs order... - function escapeHtml(unsafe) { return unsafe .replace(/&/g, '&') @@ -26,6 +22,54 @@ function escapeHtml(unsafe) { .replace(/"/g, '"') .replace(/'/g, ''') } + +function createDebugSection(title) { + const section = mtb_ui.makeElement('div', { + margin: '8px 0', + padding: '8px', + borderRadius: '4px', + backgroundColor: 'rgba(0,0,0,0.2)' + }) + + const header = mtb_ui.makeElement('h3', { + margin: '0 0 8px 0', + padding: '4px 0', + borderBottom: '1px solid rgba(255,255,255,0.1)', + fontSize: '14px', + fontWeight: 'bold', + color: '#9f9' + }) + header.textContent = title + section.appendChild(header) + + return section +} + +function createDebugContent(content, type) { + const wrapper = mtb_ui.makeElement('div', { + margin: '4px 0' + }) + + if (type === 'text') { + const text = mtb_ui.makeElement('p', { + margin: '2px 0', + fontFamily: 'monospace', + whiteSpace: 'pre-wrap' + }) + text.innerHTML = content + wrapper.appendChild(text) + } else if (type === 'image') { + const img = mtb_ui.makeElement('img', { + width: '100%', + borderRadius: '2px' + }) + img.src = content + wrapper.appendChild(img) + } + + return wrapper +} + app.registerExtension({ name: 'mtb.Debug', @@ -84,63 +128,78 @@ app.registerExtension({ onExecuted?.apply(this, args) const [data, ..._rest] = args - const prefix = 'anything_' - if (this.widgets) { + let tgt_len = this.widgets.length for (let i = 0; i < this.widgets.length; i++) { - if (this.widgets[i].name !== 'output_to_console') { + if ( + this.widgets[i].name !== 'output_to_console' && + this.widgets[i].name !== 'as_detailed_types' + ) { this.widgets[i].onRemove?.() this.widgets[i].onRemoved?.() + tgt_len -= 1 } } - this.widgets.length = 1 + this.widgets.length = tgt_len } + + const inputData = {} + + const uiData = data.ui || data + + if (uiData.items) { + uiData.items.forEach(item => { + const inputName = item.input + if (!inputData[inputName]) { + inputData[inputName] = { text: [], b64_images: [] } + } + if (item.text) { + inputData[inputName].text.push(...item.text) + } + if (item.b64_images) { + inputData[inputName].b64_images.push(...item.b64_images) + } + }) + } + let widgetI = 1 - // console.log(message) - if (data.text) { - for (const txt of data.text) { - const textDom = mtb_ui.makeElement('p', { fontFamily: 'monospace' }) - textDom.innerHTML = txt - - this.addDOMWidget( - `${prefix}_${widgetI}`, - 'CUSTOM_TEXT', - textDom, - {}, - ) - widgetI++ + for (const [inputName, content] of Object.entries(inputData)) { + if (content.text.length === 0 && content.b64_images.length === 0) { + continue } - } - if (data.b64_images) { - for (const img of data.b64_images) { - const imgDom = mtb_ui.makeElement('img', { width: '100%' }) - imgDom.src = img - this.addDOMWidget( - `${prefix}_${widgetI}`, - 'CUSTOM_IMG_B64', - mtb_ui.wrapElement(imgDom, { - overflow: 'hidden', - }), - {}, - ) + const section = createDebugSection(inputName) - widgetI++ + if (content.text.length > 0) { + content.text.forEach(text => { + section.appendChild(createDebugContent(text, 'text')) + }) } - } - // this.setSize(this.computeSize()) + if (content.b64_images.length > 0) { + content.b64_images.forEach(img => { + section.appendChild(createDebugContent(img, 'image')) + }) + } + + this.addDOMWidget( + `debug_section_${widgetI}`, + 'CUSTOM', + section, + {} + ) + widgetI++ + } this.onRemoved = function () { - // When removing this node we need to remove the input from the DOM - for (const y in this.widgets) { - if (this.widgets[y].canvas) { - this.widgets[y].canvas.remove() + for (const widget of this.widgets) { + if (widget.canvas) { + widget.canvas.remove() } - shared.cleanupNode(this) - this.widgets[y].onRemoved?.() - this.widgets[y].onRemove?.() + widget.onRemoved?.() + widget.onRemove?.() } + shared.cleanupNode(this) } } }