import onnxruntime import numpy as np import librosa import soundfile as sf from transformers import AutoTokenizer from huggingface_hub import hf_hub_download import gradio as gr from tqdm import trange import tempfile import os MODEL_ID = "ResembleAI/chatterbox-turbo-ONNX" SAMPLE_RATE = 24000 START_SPEECH_TOKEN = 6561 STOP_SPEECH_TOKEN = 6562 SILENCE_TOKEN = 4299 NUM_KV_HEADS = 16 HEAD_DIM = 64 # ------------------------- # Utility Classes # ------------------------- class RepetitionPenaltyLogitsProcessor: def __init__(self, penalty: float): self.penalty = penalty def __call__(self, input_ids: np.ndarray, scores: np.ndarray) -> np.ndarray: score = np.take_along_axis(scores, input_ids, axis=1) score = np.where(score < 0, score * self.penalty, score / self.penalty) scores_processed = scores.copy() np.put_along_axis(scores_processed, input_ids, score, axis=1) return scores_processed def download_model(name: str, dtype: str = "fp32") -> str: filename = f"{name}{'' if dtype == 'fp32' else f'_{dtype}'}.onnx" graph = hf_hub_download(MODEL_ID, subfolder="onnx", filename=filename) hf_hub_download(MODEL_ID, subfolder="onnx", filename=f"{filename}_data") return graph # ------------------------- # Load Models (once) # ------------------------- conditional_decoder_path = download_model("conditional_decoder") speech_encoder_path = download_model("speech_encoder") embed_tokens_path = download_model("embed_tokens") language_model_path = download_model("language_model") speech_encoder_session = onnxruntime.InferenceSession(speech_encoder_path) embed_tokens_session = onnxruntime.InferenceSession(embed_tokens_path) language_model_session = onnxruntime.InferenceSession(language_model_path) cond_decoder_session = onnxruntime.InferenceSession(conditional_decoder_path) tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) # ------------------------- # Core Generation Function # ------------------------- def generate_speech(text, audio_file): max_new_tokens = 1024 repetition_penalty = 1.2 # Load audio audio_values, _ = librosa.load(audio_file, sr=SAMPLE_RATE) audio_values = audio_values[np.newaxis, :].astype(np.float32) # Tokenize text input_ids = tokenizer(text, return_tensors="np")["input_ids"].astype(np.int64) repetition_penalty_processor = RepetitionPenaltyLogitsProcessor(repetition_penalty) generate_tokens = np.array([[START_SPEECH_TOKEN]], dtype=np.int64) for i in range(max_new_tokens): inputs_embeds = embed_tokens_session.run(None, {"input_ids": input_ids})[0] if i == 0: cond_emb, prompt_token, speaker_embeddings, speaker_features = speech_encoder_session.run( None, {"audio_values": audio_values} ) inputs_embeds = np.concatenate((cond_emb, inputs_embeds), axis=1) batch_size, seq_len, _ = inputs_embeds.shape past_key_values = { i.name: np.zeros( [batch_size, NUM_KV_HEADS, 0, HEAD_DIM], dtype=np.float32 ) for i in language_model_session.get_inputs() if "past_key_values" in i.name } attention_mask = np.ones((batch_size, seq_len), dtype=np.int64) position_ids = np.arange(seq_len).reshape(1, -1) logits, *present_key_values = language_model_session.run( None, dict( inputs_embeds=inputs_embeds, attention_mask=attention_mask, position_ids=position_ids, **past_key_values, ), ) logits = logits[:, -1, :] next_token_logits = repetition_penalty_processor(generate_tokens, logits) input_ids = np.argmax(next_token_logits, axis=-1, keepdims=True) generate_tokens = np.concatenate((generate_tokens, input_ids), axis=-1) if (input_ids.flatten() == STOP_SPEECH_TOKEN).all(): break # update attention_mask = np.concatenate( [attention_mask, np.ones((batch_size, 1), dtype=np.int64)], axis=1 ) position_ids = position_ids[:, -1:] + 1 for j, key in enumerate(past_key_values): past_key_values[key] = present_key_values[j] # Decode audio speech_tokens = generate_tokens[:, 1:-1] silence_tokens = np.full((speech_tokens.shape[0], 3), SILENCE_TOKEN) speech_tokens = np.concatenate( [prompt_token, speech_tokens, silence_tokens], axis=1 ) wav = cond_decoder_session.run( None, dict( speech_tokens=speech_tokens, speaker_embeddings=speaker_embeddings, speaker_features=speaker_features, ), )[0].squeeze() # Save temp file output_path = os.path.join(tempfile.gettempdir(), "output.wav") sf.write(output_path, wav, SAMPLE_RATE) return output_path # ------------------------- # Gradio UI # ------------------------- with gr.Blocks() as demo: gr.Markdown("# 🎤 Voice Cloning TTS (ONNX)") gr.Markdown("Upload a voice sample and enter text to clone speech.") with gr.Row(): text_input = gr.Textbox(label="Text", value="Hello, how are you today?") audio_input = gr.Audio(type="filepath", label="Voice Sample") generate_btn = gr.Button("Generate") audio_output = gr.Audio(label="Generated Speech") generate_btn.click( fn=generate_speech, inputs=[text_input, audio_input], outputs=audio_output, ) # Launch demo.launch()