From d202da0e928fa8a662a969bf8362ae74baeda39e Mon Sep 17 00:00:00 2001 From: melMass Date: Fri, 1 Mar 2024 00:28:11 +0100 Subject: [PATCH] =?UTF-8?q?fix=20=F0=9F=90=9B:=20stack=20images?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit also update pillow min version --- nodes/image_utils.py | 59 ++++++++++++++++++++++++------- requirements.txt | 3 +- web/mtb_widgets.js | 83 +++++++++++++++++++++++--------------------- 3 files changed, 91 insertions(+), 54 deletions(-) diff --git a/nodes/image_utils.py b/nodes/image_utils.py index 575831e..74675ab 100644 --- a/nodes/image_utils.py +++ b/nodes/image_utils.py @@ -1,10 +1,9 @@ import torch - from ..log import log class StackImages: - """Stack the input images horizontally or vertically""" + """Stack the input images horizontally or vertically.""" @classmethod def INPUT_TYPES(cls): @@ -20,26 +19,59 @@ class StackImages: tensors = list(kwargs.values()) log.debug( - f"Stacking {len(tensors)} tensors {'vertically' if vertical else 'horizontally'}" + f"Stacking {len(tensors)} tensors " + f"{'vertically' if vertical else 'horizontally'}" ) - log.debug(list(kwargs.keys())) - ref_shape = tensors[0].shape - for tensor in tensors[1:]: - if tensor.shape[1:] != ref_shape[1:]: + normalized_tensors = [ + self.normalize_to_rgba(tensor) for tensor in tensors + ] + + if vertical: + width = normalized_tensors[0].shape[2] + if any(tensor.shape[2] != width for tensor in normalized_tensors): raise ValueError( - "All tensors must have the same dimensions except for the stacking dimension." + "All tensors must have the same width " + "for vertical stacking." ) + dim = 1 + else: + height = normalized_tensors[0].shape[1] + if any(tensor.shape[1] != height for tensor in normalized_tensors): + raise ValueError( + "All tensors must have the same height " + "for horizontal stacking." + ) + dim = 2 - dim = 1 if vertical else 2 - - stacked_tensor = torch.cat(tensors, dim=dim) + stacked_tensor = torch.cat(normalized_tensors, dim=dim) return (stacked_tensor,) + def normalize_to_rgba(self, tensor): + """Normalize tensor to have 4 channels (RGBA).""" + _, _, _, channels = tensor.shape + # already RGBA + if channels == 4: + return tensor + # RGB to RGBA + elif channels == 3: + alpha_channel = torch.ones( + tensor.shape[:-1] + (1,), device=tensor.device + ) # Add an alpha channel + return torch.cat((tensor, alpha_channel), dim=-1) + else: + raise ValueError( + "Tensor has an unsupported number of channels: " + "expected 3 (RGB) or 4 (RGBA)." + ) + class PickFromBatch: - """Pick a specific number of images from a batch, either from the start or end.""" + """Pick a specific number of images from a batch. + + either from the start or end. + """ @classmethod def INPUT_TYPES(cls): @@ -62,7 +94,8 @@ class PickFromBatch: count = min(count, batch_size) if count < batch_size: log.warning( - f"Requested {count} images, but only {batch_size} are available." + f"Requested {count} images, " + f"but only {batch_size} are available." ) if from_direction == "end": diff --git a/requirements.txt b/requirements.txt index 96d14b6..a5a35e8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,4 +6,5 @@ rembg imageio_ffmpeg rich rich_argparse -matplotlib \ No newline at end of file +matplotlib +pillow>=10 diff --git a/web/mtb_widgets.js b/web/mtb_widgets.js index 5f88a20..9f26b87 100644 --- a/web/mtb_widgets.js +++ b/web/mtb_widgets.js @@ -43,7 +43,7 @@ const calculateTextDimensions = (ctx, value, width, fontSize = 16) => { const textHeight = (lines.length + 1) * fontSize const maxLineWidth = lines.reduce( (maxWidth, line) => Math.max(maxWidth, ctx.measureText(line).width), - 0 + 0, ) return { textHeight, maxLineWidth } } @@ -106,7 +106,7 @@ export const MtbWidgets = { ctx.fillText( this.label || this.name, margin * 2 + 5, - currentY + H * 0.7 + currentY + H * 0.7, ) ctx.fillStyle = text_color ctx.textAlign = 'right' @@ -115,10 +115,10 @@ export const MtbWidgets = { Number(this.value).toFixed( this.options?.precision !== undefined ? this.options.precision - : 3 + : 3, ), widget_width - margin * 2 - 20, - currentY + H * 0.7 + currentY + H * 0.7, ) } } @@ -195,12 +195,12 @@ export const MtbWidgets = { try { //solve the equation if possible v = eval(v) - } catch (e) { } + } catch (e) {} } this.value = Number(v) shared.inner_value_change(this, this.value, event) }.bind(w), - event + event, ) } } @@ -210,7 +210,7 @@ export const MtbWidgets = { function () { shared.inner_value_change(this, this.value, event) }.bind(this), - 20 + 20, ) app.canvas.setDirty(true) @@ -259,7 +259,7 @@ export const MtbWidgets = { border, widgetY + border, widgetWidth - border * 2, - height - border * 2 + height - border * 2, ) const color = parseCss(this.value.default || this.value) if (!color) { @@ -362,7 +362,7 @@ export const MtbWidgets = { }) const widgetWidth = Math.max( width || this.width || 32, - dimensions.maxLineWidth + dimensions.maxLineWidth, ) const widgetHeight = dimensions.textHeight * 1.5 return [widgetWidth, widgetHeight] @@ -386,7 +386,7 @@ export const MtbWidgets = { text-align: center; font-size: ${fontSize}px; color: var(--input-text); - line-height: 0; + line-height: 1em; font-family: monospace; ` w.value = val @@ -445,7 +445,7 @@ const mtb_widgets = { enabled: value, }), }) - .then((response) => { }) + .then((response) => {}) .catch((error) => { console.error('Error:', error) }) @@ -460,7 +460,7 @@ const mtb_widgets = { return { widget: node.addCustomWidget( - MtbWidgets.BOOL(inputName, inputData[1]?.default || false) + MtbWidgets.BOOL(inputName, inputData[1]?.default || false), ), minWidth: 150, minHeight: 30, @@ -471,7 +471,7 @@ const mtb_widgets = { console.debug('Registering color') return { widget: node.addCustomWidget( - MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000') + MtbWidgets.COLOR(inputName, inputData[1]?.default || '#ff0000'), ), minWidth: 150, minHeight: 30, @@ -575,7 +575,7 @@ const mtb_widgets = { type, index, connected, - link_info + link_info, ) { const r = onConnectionsChange ? onConnectionsChange.apply(this, arguments) @@ -593,7 +593,7 @@ const mtb_widgets = { ? onNodeCreated.apply(this, arguments) : undefined const internal_count = this.widgets.find( - (w) => w.name === 'internal_count' + (w) => w.name === 'internal_count', ) shared.hideWidgetForGood(this, internal_count) internal_count.afterQueued = function () { @@ -633,24 +633,24 @@ const mtb_widgets = { imgURLs = imgURLs.concat( message.gif.map((params) => { return api.apiURL( - '/view?' + new URLSearchParams(params).toString() + '/view?' + new URLSearchParams(params).toString(), ) - }) + }), ) } if (message.apng) { imgURLs = imgURLs.concat( message.apng.map((params) => { return api.apiURL( - '/view?' + new URLSearchParams(params).toString() + '/view?' + new URLSearchParams(params).toString(), ) - }) + }), ) } let i = 0 for (const img of imgURLs) { const w = this.addCustomWidget( - MtbWidgets.DEBUG_IMG(`${prefix}_${i}`, img) + MtbWidgets.DEBUG_IMG(`${prefix}_${i}`, img), ) w.parent = this i++ @@ -678,12 +678,12 @@ const mtb_widgets = { this.changeMode(LiteGraph.ALWAYS) const raw_iteration = this.widgets.find( - (w) => w.name === 'raw_iteration' + (w) => w.name === 'raw_iteration', ) const raw_loop = this.widgets.find((w) => w.name === 'raw_loop') const total_frames = this.widgets.find( - (w) => w.name === 'total_frames' + (w) => w.name === 'total_frames', ) const loop_count = this.widgets.find((w) => w.name === 'loop_count') @@ -693,12 +693,12 @@ const mtb_widgets = { raw_iteration._value = 0 const value_preview = this.addCustomWidget( - MtbWidgets['DEBUG_STRING']('value_preview', 'Idle') + MtbWidgets['DEBUG_STRING']('value_preview', 'Idle'), ) value_preview.parent = this const loop_preview = this.addCustomWidget( - MtbWidgets['DEBUG_STRING']('loop_preview', 'Iteration: Idle') + MtbWidgets['DEBUG_STRING']('loop_preview', 'Iteration: Idle'), ) loop_preview.parent = this @@ -716,16 +716,17 @@ const mtb_widgets = { 'button', `Reset`, 'reset', - onReset + onReset, ) const run_button = this.addWidget('button', `Queue`, 'queue', () => { onReset() // this could maybe be a setting or checkbox app.queuePrompt(0, total_frames.value * loop_count.value) window.MTB?.notify?.( - `Started a queue of ${total_frames.value} frames (for ${loop_count.value + `Started a queue of ${total_frames.value} frames (for ${ + loop_count.value } loop, so ${total_frames.value * loop_count.value})`, - 5000 + 5000, ) }) @@ -738,14 +739,16 @@ const mtb_widgets = { this.value++ raw_loop.value = Math.floor(this.value / total_frames.value) - value_preview.value = `frame: ${raw_iteration.value % total_frames.value - } / ${total_frames.value - 1}` + value_preview.value = `frame: ${ + raw_iteration.value % total_frames.value + } / ${total_frames.value - 1}` if (raw_loop.value + 1 > loop_count.value) { loop_preview.value = 'Done 😎!' } else { - loop_preview.value = `current loop: ${raw_loop.value + 1}/${loop_count.value - }` + loop_preview.value = `current loop: ${raw_loop.value + 1}/${ + loop_count.value + }` } } @@ -760,7 +763,7 @@ const mtb_widgets = { type, index, connected, - link_info + link_info, ) { const r = onConnectionsChange ? onConnectionsChange.apply(this, arguments) @@ -781,7 +784,7 @@ const mtb_widgets = { const input = this.addInput( `replacement_${this.widgets.length}`, 'STRING', - '' + '', ) console.log(input) this.addWidget('STRING', `replacement_${this.widgets.length}`, '') @@ -798,7 +801,7 @@ const mtb_widgets = { 'remove', function (value, widget, node) { console.log(`Button clicked: ${value}`, widget, node) - } + }, ) return r @@ -839,7 +842,7 @@ const mtb_widgets = { if (style && style.length >= 1) { if (style[0]) { window.MTB?.notify?.( - `Extracted positive from ${this.widgets[0].value}` + `Extracted positive from ${this.widgets[0].value}`, ) const tn = LiteGraph.createNode('Text box') app.graph.add(tn) @@ -847,7 +850,7 @@ const mtb_widgets = { tn.widgets[0].value = style[0] } else { window.MTB?.notify?.( - `No positive to extract for ${this.widgets[0].value}` + `No positive to extract for ${this.widgets[0].value}`, ) } } @@ -860,7 +863,7 @@ const mtb_widgets = { if (style && style.length >= 2) { if (style[1]) { window.MTB?.notify?.( - `Extracted negative from ${this.widgets[0].value}` + `Extracted negative from ${this.widgets[0].value}`, ) const tn = LiteGraph.createNode('Text box') app.graph.add(tn) @@ -868,7 +871,7 @@ const mtb_widgets = { tn.widgets[0].value = style[1] } else { window.MTB.notify( - `No negative to extract for ${this.widgets[0].value}` + `No negative to extract for ${this.widgets[0].value}`, ) } } @@ -916,7 +919,7 @@ const mtb_widgets = { type, index, connected, - link_info + link_info, ) { const r = onConnectionsChange ? onConnectionsChange.apply(this, arguments) @@ -930,7 +933,7 @@ const mtb_widgets = { //- infer type if (link_info) { const fromNode = this.graph._nodes.find( - (otherNode) => otherNode.id == link_info.origin_id + (otherNode) => otherNode.id == link_info.origin_id, ) const type = fromNode.outputs[link_info.origin_slot].type this.inputs[index].type = type