diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..6cf4ed5 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +/__pycache__/ +/node/__pycache__/ diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..8eb525e --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .node.orpheus_lmstudio_tts import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/node/decoder.py b/node/decoder.py new file mode 100644 index 0000000..fdc0137 --- /dev/null +++ b/node/decoder.py @@ -0,0 +1,72 @@ +import torch +import numpy as np +from snac import SNAC + +# One global SNAC, placed on GPU if available. +_SNAC = None +_DEVICE = None + +def _init_snac(): + global _SNAC, _DEVICE + if _SNAC is not None: + return + if torch.cuda.is_available(): + _DEVICE = torch.device("cuda") + elif getattr(torch.backends, "mps", None) and torch.backends.mps.is_available(): + _DEVICE = torch.device("cpu") # MPS path not stable for SNAC; keep CPU. + else: + _DEVICE = torch.device("cpu") + + _SNAC = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").eval().to(_DEVICE) + +def get_snac_device(): + _init_snac() + return str(_DEVICE) + +def convert_to_audio(multiframe, count): + """ + multiframe: list[int] flat, length multiple of 7 (codebook layout) + Returns: bytes of int16 PCM at 24kHz, or None if invalid. + """ + _init_snac() + + if len(multiframe) < 7: + return None + + # Lay out codebooks: c0 (1), c1 (2), c2 (4) per 7-tuple frame + num_frames = len(multiframe) // 7 + frame = multiframe[: num_frames * 7] + + c0 = torch.tensor([], device=_DEVICE, dtype=torch.int32) + c1 = torch.tensor([], device=_DEVICE, dtype=torch.int32) + c2 = torch.tensor([], device=_DEVICE, dtype=torch.int32) + + for j in range(num_frames): + i = 7 * j + # c0: [i] + c0 = torch.cat([c0, torch.tensor([frame[i]], device=_DEVICE, dtype=torch.int32)]) + # c1: [i+1, i+4] + c1 = torch.cat([c1, + torch.tensor([frame[i + 1]], device=_DEVICE, dtype=torch.int32), + torch.tensor([frame[i + 4]], device=_DEVICE, dtype=torch.int32)]) + # c2: [i+2, i+3, i+5, i+6] + c2 = torch.cat([c2, + torch.tensor([frame[i + 2]], device=_DEVICE, dtype=torch.int32), + torch.tensor([frame[i + 3]], device=_DEVICE, dtype=torch.int32), + torch.tensor([frame[i + 5]], device=_DEVICE, dtype=torch.int32), + torch.tensor([frame[i + 6]], device=_DEVICE, dtype=torch.int32)]) + + # Bounds check + for cb in (c0, c1, c2): + if torch.any(cb < 0) or torch.any(cb > 4096): + return None + + codes = [c0.unsqueeze(0), c1.unsqueeze(0), c2.unsqueeze(0)] + + with torch.inference_mode(): + audio_hat = _SNAC.decode(codes) # (B, 1, n) + # Empirically use middle slice for stable chunking (as in upstream) + audio_slice = audio_hat[:, :, 2048:4096] + audio_np = audio_slice.detach().cpu().numpy() + pcm16 = (audio_np * 32767.0).astype(np.int16).tobytes() + return pcm16 diff --git a/node/orpheus_lmstudio_tts.py b/node/orpheus_lmstudio_tts.py new file mode 100644 index 0000000..96fdab8 --- /dev/null +++ b/node/orpheus_lmstudio_tts.py @@ -0,0 +1,232 @@ +import re +import time +import hashlib +import numpy as np +import torch +import lmstudio as lms +from comfy.utils import ProgressBar + +try: + from transformers import AutoTokenizer +except ImportError: + print("Transformers library not found. Please install it with: pip install transformers") + AutoTokenizer = None + +from .decoder import ( + convert_to_audio, + get_snac_device, +) + +CUSTOM_RE = re.compile(r"") +SAMPLE_RATE = 24000 + +VOICES = [ + "tara","leah","jess","leo","dan","mia","zac","zoe", # English +] + +TOKENIZER = None +def get_tokenizer(): + global TOKENIZER + if TOKENIZER is None and AutoTokenizer is not None: + try: + TOKENIZER = AutoTokenizer.from_pretrained("gpt2") + except Exception as e: + print(f"Warning: Could not load GPT-2 tokenizer for estimation: {e}") + return TOKENIZER + + +def _format_prompt(text, voice): + return f"<|audio|>{voice}: {text}<|eot_id|>" + +class OrpheusLMStudioTTS: + """ + Streams text completions from LM Studio, collects all audio tokens, + then decodes with SNAC using a two-phase progress bar that accurately + reflects the workload of the LLM (70%) and the SNAC decoder (30%). + Includes performance timers for debugging. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "text": ("STRING", {"default": "Hello there. This should sound natural.", "multiline": True}), + "voice": (VOICES, ), + "model_key": ("STRING", {"default": ""}), + "max_tokens": ("INT", {"default": 4096, "min": 256, "max": 131072, "step": 1}), + "temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.01}), + "top_p": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}), + "repeat_penalty": ("FLOAT", {"default": 1.1, "min": 1.0, "max": 2.0, "step": 0.01}), + "timeout_seconds": ("INT", {"default": 300, "min": 5, "max": 3600, "step": 1}), + "auto_unload": (["True", "False"], {"default": "True"}), + "unload_delay": ("INT", {"default": 0, "min": 0, "max": 3600, "step": 1}), + "stop_at_eot": ("BOOLEAN", {"default": True}), + "seed": ("INT", {"default": -1, "min": -1, "max": 0xffffffffffffffff}), + "debug": ("BOOLEAN", {"default": False}), + }, + "optional": { "custom_stop": ("STRING", {"default": ""}) } + } + + RETURN_TYPES = ("AUDIO",) + RETURN_NAMES = ("audio",) + FUNCTION = "run" + CATEGORY = "Orpheus/LM Studio" + + @classmethod + def IS_CHANGED(cls, **kwargs): + m = hashlib.sha256() + for k, v in kwargs.items(): + m.update(str(v).encode()) + return m.hexdigest() + + def _get_model(self, client, model_key, auto_unload, unload_delay): + if model_key and model_key.strip(): + return client.llm.model(model_key, ttl=unload_delay) if (auto_unload == "True" and unload_delay > 0) else client.llm.model(model_key) + return client.llm.model() + + def run( + self, text, voice, model_key, max_tokens, temperature, top_p, + repeat_penalty, timeout_seconds, auto_unload, unload_delay, + stop_at_eot, seed, debug, custom_stop="", + ): + if seed == -1: + seed = int(time.time() * 1000) & 0xffffffffffffffff + + prompt = _format_prompt(text, voice) + + stop_strings = ["<|eot_id|>"] if stop_at_eot else [] + if custom_stop.strip(): + stop_strings.append(custom_stop.strip()) + + config = { + "temperature": float(temperature), "topP": float(top_p), + "maxTokens": int(max_tokens), "seed": int(seed), + "repeatPenalty": float(repeat_penalty), + } + if stop_strings: + config["stopStrings"] = stop_strings + + if debug: + print(f"[Orpheus-LMS] SNAC device: {get_snac_device()}") + print(f"[Orpheus-LMS] Using model_key='{model_key or '(default)'}'") + print(f"[Orpheus-LMS] Config: {config}") + + pbar = ProgressBar(100) + + start_time = time.time() + llm_end_time, snac_prep_end_time, snac_decode_end_time = 0, 0, 0 + + # --- Phase 1: LLM Token Generation (0% -> 70%) --- + estimated_total_audio_tokens = 100 + tokenizer = get_tokenizer() + if tokenizer: + input_token_count = len(tokenizer.encode(text)) + estimated_total_audio_tokens = max(1, input_token_count * 25) + if debug: + print(f"[Orpheus-LMS] Input text tokens: {input_token_count}, estimated audio tokens for progress bar: {estimated_total_audio_tokens}") + + llm_output_fragments = [] + model = None + current_audio_tokens = 0 + + with lms.Client() as client: + try: + model = self._get_model(client, model_key, auto_unload, unload_delay) + stream = model.complete_stream(prompt, config=config) + deadline = time.time() + timeout_seconds + + for frag in stream: + if time.time() > deadline: + try: stream.cancel() + except: pass + raise TimeoutError(f"LLM prediction exceeded {timeout_seconds}s") + + content = frag.content or "" + llm_output_fragments.append(content) + + if content.startswith(' 100%) --- + collected_pcm = bytearray() + + all_tokens = [int(m.group(1)) for m in CUSTOM_RE.finditer(llm_output_str)] + + total_accepted_tokens = 0 + mapped_tokens = [] + for n in all_tokens: + mapped = n - 10 - ((total_accepted_tokens % 7) * 4096) + if mapped > 0: + mapped_tokens.append(mapped) + total_accepted_tokens += 1 + + total_snac_calls = max(0, ((total_accepted_tokens - 28) // 7) + 1) if total_accepted_tokens > 27 else 0 + snac_prep_end_time = time.time() + + if debug: + # ADDED BACK: Log to compare the estimate with the actual result. + print(f"[Orpheus-LMS] Audio Token Estimate vs. Actual: {estimated_total_audio_tokens} vs. {total_accepted_tokens}") + print(f"[Orpheus-LMS] Expecting {total_snac_calls} SNAC decoding calls.") + + if total_snac_calls > 0: + for i in range(total_snac_calls): + start_idx = i * 7 + end_idx = start_idx + 28 + sub = mapped_tokens[start_idx:end_idx] + + audio_chunk = convert_to_audio(sub, i) + if audio_chunk: + collected_pcm.extend(audio_chunk) + + progress = 70 + int(((i + 1) / total_snac_calls) * 30) + pbar.update_absolute(progress) + + snac_decode_end_time = time.time() + pbar.update_absolute(100) + + if debug: + llm_duration = llm_end_time - start_time + prep_duration = snac_prep_end_time - llm_end_time + decode_duration = snac_decode_end_time - snac_prep_end_time + total_duration = snac_decode_end_time - start_time + secs = len(collected_pcm) / 2 / SAMPLE_RATE + + print("--- Performance ---") + print(f" LLM Stream Collection: {llm_duration:.2f}s") + print(f" SNAC Pre-calculation: {prep_duration:.4f}s") + print(f" SNAC Decoding: {decode_duration:.2f}s") + print(f" Total Generation Time: {total_duration:.2f}s") + print(f" Final Audio Duration: {secs:.2f}s") + print("-------------------") + + if not collected_pcm: + silence = np.zeros(int(SAMPLE_RATE / 10), dtype=np.int16).tobytes() + collected_pcm.extend(silence) + + pcm_i16 = np.frombuffer(bytes(collected_pcm), dtype=np.int16) + audio_f32 = (pcm_i16.astype(np.float32) / 32767.0).reshape(1, -1) + waveform = torch.from_numpy(audio_f32) + + audio = {"waveform": waveform.unsqueeze(0), "sample_rate": SAMPLE_RATE} + return (audio,) + +NODE_CLASS_MAPPINGS = {"OrpheusLMStudioTTS": OrpheusLMStudioTTS} +NODE_DISPLAY_NAME_MAPPINGS = {"OrpheusLMStudioTTS": "Orpheus TTS (LM Studio)"} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..2cfc19e --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +lmstudio>=1.5.0 +snac>=1.2.1 +torch>=2.1 +numpy>=1.23