fix 🐛: stack images
also update pillow min version
This commit is contained in:
+46
-13
@@ -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":
|
||||
|
||||
+2
-1
@@ -6,4 +6,5 @@ rembg
|
||||
imageio_ffmpeg
|
||||
rich
|
||||
rich_argparse
|
||||
matplotlib
|
||||
matplotlib
|
||||
pillow>=10
|
||||
|
||||
+43
-40
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user