Initial audio input and output handling

This commit is contained in:
Paul
2024-07-02 17:58:21 +00:00
parent 3200881e66
commit f3b74e2d3e
2 changed files with 218 additions and 34 deletions
+119 -8
View File
@@ -8,6 +8,8 @@ from torchvision import transforms
import torch
import base64
import time
import torchaudio
import soundfile as sf
from replicate.client import Client
from .schema_to_node import (
schema_to_comfyui_input_types,
@@ -32,7 +34,11 @@ def create_comfyui_node(schema):
def INPUT_TYPES(cls):
return schema_to_comfyui_input_types(schema)
RETURN_TYPES = (return_type,)
RETURN_TYPES = (
tuple(return_type.values())
if isinstance(return_type, dict)
else (return_type,)
)
FUNCTION = "run_replicate_model"
CATEGORY = "Replicate"
@@ -45,6 +51,8 @@ def create_comfyui_node(schema):
)
if input_type == "IMAGE":
kwargs[key] = self.image_to_base64(value)
elif input_type == "AUDIO":
kwargs[key] = self.audio_to_base64(value)
def image_to_base64(self, image):
if isinstance(image, torch.Tensor):
@@ -60,6 +68,31 @@ def create_comfyui_node(schema):
img_str = base64.b64encode(buffer.getvalue()).decode()
return f"data:image/png;base64,{img_str}"
def audio_to_base64(self, audio):
if (
isinstance(audio, dict)
and "waveform" in audio
and "sample_rate" in audio
):
waveform = audio["waveform"]
sample_rate = audio["sample_rate"]
else:
waveform, sample_rate = audio
# Ensure waveform is 2D
if waveform.dim() == 1:
waveform = waveform.unsqueeze(0)
elif waveform.dim() > 2:
waveform = waveform.squeeze()
if waveform.dim() > 2:
raise ValueError("Waveform must be 1D or 2D")
buffer = io.BytesIO()
sf.write(buffer, waveform.numpy().T, sample_rate, format="wav")
buffer.seek(0)
audio_str = base64.b64encode(buffer.getvalue()).decode()
return f"data:audio/wav;base64,{audio_str}"
def handle_array_inputs(self, kwargs):
array_inputs = inputs_that_need_arrays(schema)
for input_name in array_inputs:
@@ -75,13 +108,18 @@ def create_comfyui_node(schema):
def log_input(self, kwargs):
truncated_kwargs = {
k: v[:20] + "..."
if isinstance(v, str) and v.startswith("data:image")
if isinstance(v, str)
and (v.startswith("data:image") or v.startswith("data:audio"))
else v
for k, v in kwargs.items()
}
print(f"Running {replicate_model} with {truncated_kwargs}")
def handle_image_output(self, output):
if output is None:
print("No image output received")
return None
# Handle both string and list outputs
output_list = [output] if isinstance(output, str) else list(output)
if output_list:
@@ -112,19 +150,90 @@ def create_comfyui_node(schema):
print("No output received from the model")
return None
def handle_audio_output(self, output):
if output is None:
print("No audio output received from the model")
return None
if isinstance(output, str):
output_list = [output]
elif isinstance(output, (list, tuple)):
output_list = list(output)
else:
print(f"Unexpected output type: {type(output)}")
return None
if output_list:
audio_data = []
for audio_url in output_list:
if audio_url:
response = requests.get(audio_url)
if response.status_code == 200:
audio_content = BytesIO(response.content)
waveform, sample_rate = torchaudio.load(audio_content)
audio_data.append(
{
"waveform": waveform.unsqueeze(0),
"sample_rate": sample_rate,
}
)
else:
print(
f"Failed to download audio. Status code: {response.status_code}"
)
else:
print("Empty audio URL received")
# If there's only one audio file, return it directly
if len(audio_data) == 1:
return audio_data[0]
# If there are multiple audio files, return them as a list
return audio_data
else:
print("No valid audio URLs in the output")
return None
def remove_falsey_optional_inputs(self, kwargs):
optional_inputs = self.INPUT_TYPES().get("optional", {})
for key in list(kwargs.keys()):
if key in optional_inputs and not kwargs[key]:
del kwargs[key]
def run_replicate_model(self, **kwargs):
self.handle_array_inputs(kwargs)
self.remove_falsey_optional_inputs(kwargs)
self.convert_input_images_to_base64(kwargs)
self.log_input(kwargs)
output = replicate.run(replicate_model, input=kwargs)
kwargs_without_force_rerun = {
k: v for k, v in kwargs.items() if k != "force_rerun"
}
output = replicate.run(replicate_model, input=kwargs_without_force_rerun)
print(f"Output: {output}")
if return_type == "IMAGE":
output = self.handle_image_output(output)
processed_outputs = []
if isinstance(return_type, dict):
for prop_name, prop_type in return_type.items():
if prop_type == "IMAGE":
processed_outputs.append(
self.handle_image_output(output.get(prop_name))
)
elif prop_type == "AUDIO":
processed_outputs.append(
self.handle_audio_output(output.get(prop_name))
)
elif prop_type == "STRING":
processed_outputs.append(
"".join(list(output.get(prop_name, ""))).strip()
)
else:
output = "".join(list(output)).strip()
if return_type == "IMAGE":
processed_outputs.append(self.handle_image_output(output))
elif return_type == "AUDIO":
processed_outputs.append(self.handle_audio_output(output))
else:
processed_outputs.append("".join(list(output)).strip())
return (output,)
return tuple(processed_outputs)
return node_name, ReplicateToComfyUI
@@ -135,7 +244,9 @@ def create_comfyui_nodes_from_schemas(schemas_dir):
schemas_dir_path = os.path.join(current_path, schemas_dir)
for schema_file in os.listdir(schemas_dir_path):
if schema_file.endswith(".json"):
with open(os.path.join(schemas_dir_path, schema_file), "r", encoding="utf-8") as f:
with open(
os.path.join(schemas_dir_path, schema_file), "r", encoding="utf-8"
) as f:
schema = json.load(f)
node_name, node_class = create_comfyui_node(schema)
nodes[node_name] = node_class