feat: ✨ batch apply text template if inputs are lists
same as previous commit
This commit is contained in:
+51
-5
@@ -19,6 +19,7 @@ from ..utils import (
|
||||
LazyProxyTensor,
|
||||
apply_easing,
|
||||
get_server_info,
|
||||
get_torch_tensor_info,
|
||||
numpy_NFOV,
|
||||
pil2tensor,
|
||||
tensor2np,
|
||||
@@ -135,12 +136,57 @@ class MTB_ApplyTextTemplate:
|
||||
CATEGORY = "mtb/utils"
|
||||
FUNCTION = "execute"
|
||||
|
||||
def execute(self, *, template: str, **kwargs):
|
||||
res = f"{template}"
|
||||
for k, v in kwargs.items():
|
||||
res = res.replace(f"{{{k}}}", f"{v}")
|
||||
def execute(self, *, template: str, **kwargs) -> tuple[str | list[str]]:
|
||||
keys = list(kwargs.keys())
|
||||
values = list(kwargs.values())
|
||||
|
||||
return (res,)
|
||||
has_list = any(isinstance(v, list) for v in values)
|
||||
target_length = -1
|
||||
|
||||
if has_list:
|
||||
first_list = next(x for x in values if isinstance(x, list))
|
||||
|
||||
target_length = len(first_list)
|
||||
same_length = all(
|
||||
len(v) == target_length for v in values if isinstance(v, list)
|
||||
)
|
||||
if not same_length:
|
||||
raise ValueError(
|
||||
"Text template received multiple list[str] but their size is varying, they should match..."
|
||||
)
|
||||
|
||||
if has_list:
|
||||
results = []
|
||||
|
||||
for it in range(target_length):
|
||||
res = f"{template}"
|
||||
for k, v in kwargs.items():
|
||||
if isinstance(v, list):
|
||||
res = self.apply_res(res, k, v[it])
|
||||
else:
|
||||
res = self.apply_res(res, k, v)
|
||||
results.append(res)
|
||||
|
||||
return (results,)
|
||||
|
||||
else:
|
||||
res = f"{template}"
|
||||
for k, v in kwargs.items():
|
||||
res = self.apply_res(res, k, v)
|
||||
|
||||
return (res,)
|
||||
|
||||
def apply_res(self, res, key, value):
|
||||
if isinstance(value, float):
|
||||
value = f"{value:.3f}"
|
||||
elif isinstance(value, torch.Tensor):
|
||||
value = get_torch_tensor_info(value)
|
||||
else:
|
||||
log.debug(
|
||||
f"Falling back to default string conversion for {key} of type {type(value).__name__}"
|
||||
)
|
||||
|
||||
return res.replace(f"{{{key}}}", f"{value}")
|
||||
|
||||
|
||||
class MTB_MatchDimensions:
|
||||
|
||||
+34
-11
@@ -431,7 +431,7 @@ const update_dynamic_properties = (node) => {
|
||||
* @param {NodeType} nodeType The nodetype to attach the documentation to
|
||||
* @param {str} prefix A prefix added to each dynamic inputs
|
||||
* @param {str | [str]} inputType The datatype(s) of those dynamic inputs
|
||||
* @param {{separator?:string, start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||
* @param {{separator?:string,rename_menu?:'label'|'name', start_index?:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} [opts] Extra options
|
||||
* @returns
|
||||
*/
|
||||
export const setupDynamicConnections = (
|
||||
@@ -445,31 +445,55 @@ export const setupDynamicConnections = (
|
||||
Object.getOwnPropertyDescriptors(nodeType).title.value,
|
||||
)
|
||||
|
||||
/** @type {{separator:string, start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||
/** @type {{separator:string,rename_menu?:"label"|"name" start_index:number, link?:LLink, ioSlot?:INodeInputSlot | INodeOutputSlot}?} */
|
||||
const options = Object.assign(
|
||||
{
|
||||
separator: '_',
|
||||
start_index: 1,
|
||||
rename_menu: 'label',
|
||||
},
|
||||
opts || {},
|
||||
)
|
||||
nodeType.prototype.getSlotMenuOptions = function (slot) {
|
||||
const is_valid_name = (node, val) => {
|
||||
return true
|
||||
}
|
||||
nodeType.prototype.getSlotMenuOptions = (slot) => {
|
||||
if (!slot.input) {
|
||||
return
|
||||
}
|
||||
infoLogger('Slot Menu', { slot })
|
||||
return [
|
||||
{
|
||||
content: 'Rename Input',
|
||||
content: `Rename Input (${options.rename_menu})`,
|
||||
callback: () => {
|
||||
let dialog = app.canvas.createDialog(
|
||||
const dialog = app.canvas.createDialog(
|
||||
"<span class='name'>Name</span><input autofocus type='text'/><button>OK</button>",
|
||||
{},
|
||||
)
|
||||
let dialogInput = dialog.querySelector('input')
|
||||
const dialogInput = dialog.querySelector('input')
|
||||
if (dialogInput) {
|
||||
dialogInput.value = slot.input.label || slot.input.name || ''
|
||||
if (options.rename_menu === 'label') {
|
||||
dialogInput.value = slot.input.label || slot.input.name || ''
|
||||
} else if (options.rename_menu === 'name') {
|
||||
dialogInput.value = slot.input.name || ''
|
||||
}
|
||||
}
|
||||
let inner = () => {
|
||||
app.graph.beforeChange()
|
||||
const inner = () => {
|
||||
// TODO: check if name exists or other guards
|
||||
slot.input.label = dialogInput.value
|
||||
const val = dialogInput.value
|
||||
if (!is_valid_name(slot.node, val)) {
|
||||
dialog.close()
|
||||
return
|
||||
}
|
||||
|
||||
app.graph.beforeChange()
|
||||
if (options.rename_menu === 'label') {
|
||||
slot.input.label = val
|
||||
} else if (options.rename_menu === 'name') {
|
||||
slot.input.name = val
|
||||
slot.input.label = val
|
||||
}
|
||||
|
||||
app.graph.afterChange()
|
||||
|
||||
dialog.close()
|
||||
@@ -491,7 +515,6 @@ export const setupDynamicConnections = (
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
}
|
||||
|
||||
const onConfigure = nodeType.prototype.onConfigure
|
||||
|
||||
+3
-1
@@ -1196,7 +1196,9 @@ const mtb_widgets = {
|
||||
|
||||
//NOTE: dynamic nodes
|
||||
case 'Apply Text Template (mtb)': {
|
||||
shared.setupDynamicConnections(nodeType, 'var', '*')
|
||||
shared.setupDynamicConnections(nodeType, 'var', '*', {
|
||||
rename_menu: 'name',
|
||||
})
|
||||
break
|
||||
}
|
||||
case 'Save Data Bundle (mtb)': {
|
||||
|
||||
Reference in New Issue
Block a user