Compare commits

...
Author SHA1 Message Date
robertvoy 1d4dcb985e Feat: Audio Batch Divider 2026-01-25 17:32:24 +11:00
2 changed files with 145 additions and 72 deletions
+71
View File
@@ -2640,6 +2640,75 @@ class ImageBatchDivider:
return tuple(outputs)
class AudioBatchDivider:
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"audio": ("AUDIO",),
"divide_by": ("INT", {
"default": 2,
"min": 1,
"max": 10,
"step": 1,
"display": "number",
"tooltip": "Number of parts to divide the audio into"
}),
}
}
RETURN_TYPES = ByPassTypeTuple(("AUDIO",)) # Flexible for variable outputs
RETURN_NAMES = ByPassTypeTuple(tuple([f"audio_{i+1}" for i in range(10)]))
FUNCTION = "divide_audio"
OUTPUT_NODE = True
CATEGORY = "audio"
def divide_audio(self, audio, divide_by):
import torch
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate", 44100)
if waveform is None or waveform.numel() == 0:
# Return empty audio for all outputs
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": sample_rate}
return tuple([empty_audio] * 10)
total_splits = min(divide_by, 10) # Cap to max 10
# Waveform shape: [batch, channels, samples]
total_samples = waveform.shape[-1]
samples_per_split = total_samples // total_splits
remainder = total_samples % total_splits
outputs = []
start_idx = 0
for i in range(total_splits):
current_samples = samples_per_split + (1 if i < remainder else 0)
end_idx = start_idx + current_samples
split_waveform = waveform[..., start_idx:end_idx]
outputs.append({
"waveform": split_waveform,
"sample_rate": sample_rate
})
start_idx = end_idx
# Pad with empty audio up to max (10) to match RETURN_TYPES length
empty_audio = {
"waveform": torch.zeros(waveform.shape[0], waveform.shape[1], 1,
dtype=waveform.dtype, device=waveform.device),
"sample_rate": sample_rate
}
while len(outputs) < 10:
outputs.append(empty_audio)
return tuple(outputs)
class DistributedEmptyImage:
"""Produces an empty IMAGE batch used when the master delegates all work."""
@@ -2817,6 +2886,7 @@ NODE_CLASS_MAPPINGS = {
"DistributedSeed": DistributedSeed,
"DistributedModelName": DistributedModelName,
"ImageBatchDivider": ImageBatchDivider,
"AudioBatchDivider": AudioBatchDivider,
"DistributedEmptyImage": DistributedEmptyImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -2825,5 +2895,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DistributedSeed": "Distributed Seed",
"DistributedModelName": "Distributed Model Name",
"ImageBatchDivider": "Image Batch Divider",
"AudioBatchDivider": "Audio Batch Divider",
"DistributedEmptyImage": "Distributed Empty Image",
}
+74 -72
View File
@@ -1,86 +1,88 @@
import { app } from "/scripts/app.js";
// Configuration for each batch divider node type
const BATCH_DIVIDER_NODES = {
"ImageBatchDivider": { outputPrefix: "batch_", outputType: "IMAGE" },
"AudioBatchDivider": { outputPrefix: "audio_", outputType: "AUDIO" }
};
app.registerExtension({
name: "Distributed.ImageBatchDivider",
name: "Distributed.BatchDividers",
async nodeCreated(node) {
if (node.comfyClass === "ImageBatchDivider") {
try {
const updateOutputs = () => {
if (!node.widgets) return;
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (!divideByWidget) return;
const divideBy = parseInt(divideByWidget.value, 10) || 1;
const totalOutputs = divideBy; // Direct divide by value
// Ensure outputs array exists
if (!node.outputs) node.outputs = [];
// Remove excess outputs
while (node.outputs.length > totalOutputs) {
node.removeOutput(node.outputs.length - 1);
}
// Add missing outputs
while (node.outputs.length < totalOutputs) {
const outputIndex = node.outputs.length + 1;
node.addOutput(`batch_${outputIndex}`, "IMAGE");
}
if (node.setDirty) node.setDirty(true); // Refresh canvas
};
// Initial update with delay to allow workflow loading
setTimeout(updateOutputs, 200);
// Find the widget and set up responsive handlers
const config = BATCH_DIVIDER_NODES[node.comfyClass];
if (!config) return;
try {
const updateOutputs = () => {
if (!node.widgets) return;
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (divideByWidget) {
// Override callback for immediate trigger on value set
const originalCallback = divideByWidget.callback;
divideByWidget.callback = (value) => {
updateOutputs();
if (originalCallback) originalCallback.call(divideByWidget, value); // Preserve 'this' context
};
// Add event listener for real-time input changes (e.g., typing/dragging)
if (divideByWidget.inputEl) {
divideByWidget.inputEl.addEventListener('input', updateOutputs);
}
// Lightweight MutationObserver as fallback (observe attributes on widget element if available)
const observer = new MutationObserver(updateOutputs);
if (divideByWidget.element) {
observer.observe(divideByWidget.element, { attributes: true, childList: true, subtree: true });
}
// Store cleanup function
node._batchDividerCleanup = () => {
observer.disconnect();
if (divideByWidget.inputEl) {
divideByWidget.inputEl.removeEventListener('input', updateOutputs);
}
divideByWidget.callback = originalCallback; // Restore original
};
if (!divideByWidget) return;
const divideBy = parseInt(divideByWidget.value, 10) || 1;
const totalOutputs = divideBy;
// Ensure outputs array exists
if (!node.outputs) node.outputs = [];
// Remove excess outputs
while (node.outputs.length > totalOutputs) {
node.removeOutput(node.outputs.length - 1);
}
// Add post-configure hook for reliable workflow loading
const originalConfigure = node.configure;
node.configure = function(data) {
const result = originalConfigure ? originalConfigure.call(this, data) : undefined;
updateOutputs(); // Re-run after config load
return result;
// Add missing outputs
while (node.outputs.length < totalOutputs) {
const outputIndex = node.outputs.length + 1;
node.addOutput(`${config.outputPrefix}${outputIndex}`, config.outputType);
}
if (node.setDirty) node.setDirty(true);
};
// Initial update with delay to allow workflow loading
setTimeout(updateOutputs, 200);
// Find the widget and set up responsive handlers
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (divideByWidget) {
const originalCallback = divideByWidget.callback;
divideByWidget.callback = (value) => {
updateOutputs();
if (originalCallback) originalCallback.call(divideByWidget, value);
};
if (divideByWidget.inputEl) {
divideByWidget.inputEl.addEventListener('input', updateOutputs);
}
const observer = new MutationObserver(updateOutputs);
if (divideByWidget.element) {
observer.observe(divideByWidget.element, { attributes: true, childList: true, subtree: true });
}
node._batchDividerCleanup = () => {
observer.disconnect();
if (divideByWidget.inputEl) {
divideByWidget.inputEl.removeEventListener('input', updateOutputs);
}
divideByWidget.callback = originalCallback;
};
} catch (error) {
console.error("Error in ImageBatchDivider extension:", error);
}
const originalConfigure = node.configure;
node.configure = function(data) {
const result = originalConfigure ? originalConfigure.call(this, data) : undefined;
updateOutputs();
return result;
};
} catch (error) {
console.error(`Error in ${node.comfyClass} extension:`, error);
}
},
nodeBeforeRemove(node) {
if (node.comfyClass === "ImageBatchDivider" && node._batchDividerCleanup) {
if (BATCH_DIVIDER_NODES[node.comfyClass] && node._batchDividerCleanup) {
node._batchDividerCleanup();
}
}
});
});