From 83bad820a8a9a63b3d1d1b1a373e62dd82a8a511 Mon Sep 17 00:00:00 2001 From: Alexander Piskun Date: Fri, 14 Aug 2026 15:13:52 +0300 Subject: [PATCH] [Partner Nodes] feat(FishAudio): implement basic nodes Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/fishaudio.py | 49 ++++ comfy_api_nodes/nodes_fishaudio.py | 454 +++++++++++++++++++++++++++++ 2 files changed, 503 insertions(+) create mode 100644 comfy_api_nodes/apis/fishaudio.py create mode 100644 comfy_api_nodes/nodes_fishaudio.py diff --git a/comfy_api_nodes/apis/fishaudio.py b/comfy_api_nodes/apis/fishaudio.py new file mode 100644 index 000000000..f615d3d24 --- /dev/null +++ b/comfy_api_nodes/apis/fishaudio.py @@ -0,0 +1,49 @@ +from pydantic import BaseModel, Field + + +class FishAudioProsody(BaseModel): + speed: float = Field(1.0, description="Speaking rate multiplier, 0.5-2.0") + volume: float = Field(0.0, description="Volume adjustment in decibels") + + +class FishAudioTTSRequest(BaseModel): + text: str = Field(..., description="Text to synthesize") + reference_id: str | list[str] | None = Field(None, description="Voice model ID or list of IDs") + temperature: float = Field(0.7, description="Expressiveness, 0-1") + top_p: float = Field(0.7, description="Nucleus sampling diversity, (0, 1]") + prosody: FishAudioProsody = Field(..., description="Speed and volume adjustments") + normalize: bool = Field(True, description="Normalize numbers and text for English and Chinese") + format: str = Field("wav", description="Output audio format") + + +class FishAudioASRRequest(BaseModel): + language: str | None = Field(None, description="Optional ISO 639-1 language hint") + ignore_timestamps: bool = Field(True, description="Skip precise timestamp computation") + + +class FishAudioASRSegment(BaseModel): + text: str | None = Field(None, description="Segment text") + start: float | None = Field(None, description="Segment start time in seconds") + end: float | None = Field(None, description="Segment end time in seconds") + + +class FishAudioASRResponse(BaseModel): + text: str | None = Field(None, description="Transcribed text") + duration: float | None = Field(None, description="Audio duration in seconds") + segments: list[FishAudioASRSegment] | None = Field(None, description="Timestamped transcript segments") + language_code: str | None = Field(None, description="Detected language as ISO 639-1 code") + language: str | None = Field(None, description="Detected language display name") + + +class FishAudioCreateModelRequest(BaseModel): + type: str = Field("tts", description="Model type") + title: str = Field(..., description="Voice model name") + train_mode: str = Field("fast", description="Training mode; fast is instantly available") + visibility: str = Field("private", description="Model visibility") + enhance_audio_quality: bool = Field(..., description="Enhance reference audio quality") + + +class FishAudioCreateModelResponse(BaseModel): + id: str = Field(..., alias="_id", description="Voice model ID for use as reference_id") + state: str | None = Field(None, description="Training state") + visibility: str | None = Field(None, description="Model visibility") diff --git a/comfy_api_nodes/nodes_fishaudio.py b/comfy_api_nodes/nodes_fishaudio.py new file mode 100644 index 000000000..d40fc2740 --- /dev/null +++ b/comfy_api_nodes/nodes_fishaudio.py @@ -0,0 +1,454 @@ +import json +import re +import uuid + +from typing_extensions import override + +from comfy_api.latest import IO, ComfyExtension, Input +from comfy_api_nodes.apis.fishaudio import ( + FishAudioASRRequest, + FishAudioASRResponse, + FishAudioCreateModelRequest, + FishAudioCreateModelResponse, + FishAudioProsody, + FishAudioTTSRequest, +) +from comfy_api_nodes.util import ( + ApiEndpoint, + audio_bytes_to_audio_input, + audio_ndarray_to_bytesio, + audio_tensor_to_contiguous_ndarray, + sync_op, + sync_op_raw, + validate_string, +) + +FISHAUDIO_VOICE = "FISHAUDIO_VOICE" + +FISHAUDIO_VOICES = [ + ("802e3bc2b27e49c2995d23ef70e6ac89", "Energetic Male (en)"), + ("b545c585f631496c914815291da4e893", "Friendly Women (en)"), + ("933563129e564b19a115bedd57b7406a", "Sarah (en)"), + ("8d21b053e2804e2a890e1cf62f267b6f", "Verity (en)"), + ("f48d143a59a946ab87c0130fd081f349", "Polo (en)"), + ("bf322df2096a46f18c579d0baa36f41d", "Adrian (en)"), + ("98655a12fa944e26b274c535e5e03842", "E-girl (en)"), + ("0327fdb5da9e4fd782899a8058c8ae2b", "Narrator (en)"), + ("5212eb29e500460391d03af42af6552e", "Warm Conversational Voice (en)"), + ("5c8dc6a69c0b4edfb32634db6384bf34", "Warm Storyteller (en)"), + ("7a18a1851d2649108c48ec9f2c80eb2c", "Dramatic Character Male (en)"), + ("59cb5986671546eaa6ca8ae6f29f6d22", "News Narrator (zh)"), + ("bf6c479f5a384b8d857310030035824b", "Lively Female (zh)"), + ("faccba1a8ac54016bcfc02761285e67f", "Gentle Female (zh)"), + ("5161d41404314212af1254556477c17d", "Energetic Female (ja)"), + ("0089dce5fefb4c6ba9b9f2f0debe1ddc", "Calm Female (ja)"), + ("45c5d3723c9c42f598e4776dcfd5f02d", "Calm Male (ja)"), +] + +FISHAUDIO_VOICE_MAP = {label: voice_id for voice_id, label in FISHAUDIO_VOICES} + +MAX_REFERENCE_AUDIO_SECONDS = 270 + + +def _rewrite_voice_tags(text: str, voice_count: int) -> tuple[str, set[int]]: + referenced: set[int] = set() + + def repl(match: re.Match) -> str: + index = int(match.group(1)) + if index < 1 or index > voice_count: + raise ValueError( + f"@Voice{index} does not match any connected voice ({voice_count} connected)." + ) + referenced.add(index) + return f"<|speaker:{index - 1}|>" + + rewritten = re.sub(r"(? list: + return [ + IO.Float.Input( + "temperature", + default=0.7, + min=0.0, + max=1.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Expressiveness. Higher values are more varied, lower values are more consistent.", + ), + IO.Float.Input( + "top_p", + default=0.7, + min=0.01, + max=1.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Diversity via nucleus sampling.", + ), + IO.Float.Input( + "speed", + default=1.0, + min=0.5, + max=2.0, + step=0.01, + display_mode=IO.NumberDisplay.slider, + tooltip="Speaking rate. 1.0 is normal, <1.0 slower, >1.0 faster.", + ), + IO.Float.Input( + "volume", + default=0.0, + min=-10.0, + max=10.0, + step=0.5, + display_mode=IO.NumberDisplay.slider, + tooltip="Volume adjustment in decibels. 0 is no change.", + ), + IO.Boolean.Input( + "normalize", + default=True, + tooltip="Normalize numbers and text for English and Chinese, " + "improving stability for numbers and dates.", + ), + ] + + +def _multi_speaker_inputs() -> list: + return [ + IO.Autogrow.Input( + "voices", + template=IO.Autogrow.TemplatePrefix( + IO.Custom(FISHAUDIO_VOICE).Input("voice"), + prefix="voice", + min=0, + max=5, + ), + tooltip="Voices for synthesis. Leave empty for the default voice. " + "With two or more voices, mark speaker changes in the text with @Voice1, @Voice2, etc.", + ), + *_tts_option_inputs(), + ] + + +class FishAudioVoiceSelector(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioVoiceSelector", + display_name="Fish Audio Voice Selector", + category="partner/audio/Fish Audio", + description="Select a voice from the Fish Audio library for text-to-speech generation.", + inputs=[ + IO.DynamicCombo.Input( + "voice", + options=[ + *(IO.DynamicCombo.Option(label, []) for _, label in FISHAUDIO_VOICES), + IO.DynamicCombo.Option( + "custom", + [ + IO.String.Input( + "voice_id", + default="", + tooltip="Voice model ID from fish.audio, e.g. the ID in " + "https://fish.audio/m//.", + ), + ], + ), + ], + tooltip="Choose a voice, or 'custom' to enter any fish.audio voice model ID.", + ), + ], + outputs=[ + IO.Custom(FISHAUDIO_VOICE).Output(display_name="voice"), + ], + is_api_node=False, + ) + + @classmethod + def execute(cls, voice: dict) -> IO.NodeOutput: + selected = voice["voice"] + if selected == "custom": + voice_id = voice["voice_id"].strip() + if not voice_id: + raise ValueError("Custom voice ID is empty.") + return IO.NodeOutput(voice_id) + voice_id = FISHAUDIO_VOICE_MAP.get(selected) + if not voice_id: + raise ValueError(f"Unknown voice: {selected}") + return IO.NodeOutput(voice_id) + + +class FishAudioTextToSpeech(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioTextToSpeech", + display_name="Fish Audio Text to Speech", + category="partner/audio/Fish Audio", + description="Convert text to speech. Supports emotion cues in the text " + "([happy], [whispering] on s2.1-pro; (happy) on s1) and multi-speaker dialogue " + "via @Voice1/@Voice2 tags with multiple connected voices.", + inputs=[ + IO.String.Input( + "text", + multiline=True, + default="", + tooltip="The text to convert to speech. With two or more voices connected, " + "mark speaker changes with @Voice1, @Voice2, etc.", + ), + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option("s2.1-pro", _multi_speaker_inputs()), + IO.DynamicCombo.Option( + "s1", + [ + IO.Custom(FISHAUDIO_VOICE).Input( + "voice", + optional=True, + tooltip="Voice for synthesis. Leave unconnected for the default voice.", + ), + *_tts_option_inputs(), + ], + ), + ], + tooltip="Model to use for text-to-speech.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed controls whether the node should re-run; " + "results are non-deterministic regardless of seed.", + ), + ], + outputs=[ + IO.Audio.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["text"]), + expr=""" + ( + $t := widgets.text; + $type($t) = "string" + ? ( + $bytes := $length($t) + 2 * $count($match($t, /[^\\x00-\\x7F]/)); + {"type":"usd","usd": $bytes * 21.45 / 1000000, "format":{"approximate":true}} + ) + : {"type":"usd","usd": 0.02145, "format":{"approximate":true, "suffix":"/1K bytes"}} + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + text: str, + model: dict, + seed: int, + ) -> IO.NodeOutput: + validate_string(text, field_name="text", min_length=1) + model_name = model["model"] + if model_name == "s1": + voices = [model["voice"]] if model.get("voice") else [] + else: + voices = [model["voices"][key] for key in model["voices"]] + rewritten, referenced = _rewrite_voice_tags(text, len(voices)) + if len(voices) >= 2: + missing = [i for i in range(1, len(voices) + 1) if i not in referenced] + if missing: + raise ValueError( + "With multiple voices, the text must mark speaker changes with tags for " + "each connected voice; missing: " + ", ".join(f"@Voice{i}" for i in missing) + ) + reference_id: str | list[str] | None = None + if len(voices) == 1: + reference_id = voices[0] + elif voices: + reference_id = voices + request = FishAudioTTSRequest( + text=rewritten, + reference_id=reference_id, + temperature=model["temperature"], + top_p=model["top_p"], + prosody=FishAudioProsody(speed=model["speed"], volume=model["volume"]), + normalize=model["normalize"], + ) + response = await sync_op_raw( + cls, + ApiEndpoint( + path="/proxy/fishaudio/v1/tts", + method="POST", + headers={"model": model_name}, + ), + data=request, + as_binary=True, + ) + return IO.NodeOutput(audio_bytes_to_audio_input(response)) + + +class FishAudioSpeechToText(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioSpeechToText", + display_name="Fish Audio Speech to Text", + category="partner/audio/Fish Audio", + description="Transcribe audio to text with automatic language detection.", + inputs=[ + IO.Audio.Input( + "audio", + tooltip="Audio to transcribe.", + ), + IO.String.Input( + "language", + default="", + tooltip="ISO 639-1 language hint (e.g. 'en', 'zh'). " + "The language is auto-detected regardless.", + ), + IO.Boolean.Input( + "precise_timestamps", + default=False, + tooltip="Return word-level timestamped segments.", + ), + ], + outputs=[ + IO.String.Output(id="text", display_name="text"), + IO.String.Output(id="language_code", display_name="language_code"), + IO.String.Output(id="segments_json", display_name="segments_json"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd":0.00858,"format":{"approximate":true,"suffix":"/minute"}}""", + ), + ) + + @classmethod + async def execute( + cls, + audio: Input.Audio, + language: str, + precise_timestamps: bool, + ) -> IO.NodeOutput: + audio_data_np = audio_tensor_to_contiguous_ndarray(audio["waveform"]) + audio_bytes_io = audio_ndarray_to_bytesio(audio_data_np, audio["sample_rate"], "mp4", "aac") + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/fishaudio/v1/asr", method="POST"), + response_model=FishAudioASRResponse, + data=FishAudioASRRequest( + language=language.strip() or None, + ignore_timestamps=not precise_timestamps, + ), + files={"audio": ("audio.mp4", audio_bytes_io, "audio/mp4")}, + content_type="multipart/form-data", + ) + segments_json = json.dumps( + [s.model_dump(exclude_none=True) for s in (response.segments or [])], + indent=2, + ) + return IO.NodeOutput(response.text or "", response.language_code or "", segments_json) + + +class FishAudioInstantVoiceClone(IO.ComfyNode): + @classmethod + def define_schema(cls) -> IO.Schema: + return IO.Schema( + node_id="FishAudioInstantVoiceClone", + display_name="Fish Audio Instant Voice Clone", + category="partner/audio/Fish Audio", + description="Create a private cloned voice from audio samples, instantly usable " + "for text-to-speech. Provide 1-20 recordings, 10-30 seconds each recommended, " + "under 270 seconds in total.", + inputs=[ + IO.Autogrow.Input( + "files", + template=IO.Autogrow.TemplatePrefix( + IO.Audio.Input("audio"), + prefix="audio", + min=1, + max=20, + ), + tooltip="Audio recordings for voice cloning.", + ), + IO.Boolean.Input( + "enhance_audio_quality", + default=True, + tooltip="Enhance reference audio quality before training.", + ), + ], + outputs=[ + IO.Custom(FISHAUDIO_VOICE).Output(display_name="voice"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge(expr="""{"type":"usd","usd":0}"""), + ) + + @classmethod + async def execute( + cls, + files: IO.Autogrow.Type, + enhance_audio_quality: bool, + ) -> IO.NodeOutput: + total_seconds = 0.0 + for key in files: + audio = files[key] + total_seconds += audio["waveform"].shape[-1] / audio["sample_rate"] + if total_seconds >= MAX_REFERENCE_AUDIO_SECONDS: + raise ValueError( + f"Total reference audio is {total_seconds:.0f} seconds; " + f"it must be under {MAX_REFERENCE_AUDIO_SECONDS} seconds." + ) + file_tuples: list[tuple[str, tuple[str, bytes, str]]] = [] + for key in files: + audio = files[key] + audio_data_np = audio_tensor_to_contiguous_ndarray(audio["waveform"]) + audio_bytes_io = audio_ndarray_to_bytesio(audio_data_np, audio["sample_rate"], "mp4", "aac") + file_tuples.append(("voices", (f"{key}.mp4", audio_bytes_io.getvalue(), "audio/mp4"))) + response = await sync_op( + cls, + ApiEndpoint(path="/proxy/fishaudio/model", method="POST"), + response_model=FishAudioCreateModelResponse, + data=FishAudioCreateModelRequest( + title=str(uuid.uuid4()), + enhance_audio_quality=enhance_audio_quality, + ), + files=file_tuples, + content_type="multipart/form-data", + ) + return IO.NodeOutput(response.id) + + +class FishAudioExtension(ComfyExtension): + @override + async def get_node_list(self) -> list[type[IO.ComfyNode]]: + return [ + FishAudioVoiceSelector, + FishAudioTextToSpeech, + FishAudioSpeechToText, + FishAudioInstantVoiceClone, + ] + + +async def comfy_entrypoint() -> FishAudioExtension: + return FishAudioExtension()