Spaces:
Running
Running
Download app.py from Maarij-Aqeel/urdu-voice-clone: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/spaces/Maarij-Aqeel/urdu-voice-clone/resolve/main/app.py
- Command line
-
hf download hf://spaces/Maarij-Aqeel/urdu-voice-clone/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Maarij-Aqeel/urdu-voice-clone/resolve/main/app.py
10.6 kB
| import os | |
| # Starlette >= 1.0 removed the deprecated positional TemplateResponse(name, context) | |
| # call style. gradio 4.44.0 still calls templates.TemplateResponse(template_str, | |
| # context_dict), which then hits the new-only signature TemplateResponse(request, | |
| # name, ...) and passes the context dict as `name`, exploding inside jinja2's cache | |
| # with "unhashable type: 'dict'". Translate old-style calls before gradio imports. | |
| import starlette.templating as _st_templating | |
| _orig_template_response = _st_templating.Jinja2Templates.TemplateResponse | |
| def _compat_template_response(self, *args, **kwargs): | |
| if args and isinstance(args[0], str): | |
| name = args[0] | |
| context = args[1] if len(args) > 1 else kwargs.pop("context", None) or {} | |
| request = context.get("request") if isinstance(context, dict) else None | |
| if request is not None: | |
| return _orig_template_response(self, request, name, context, *args[2:], **kwargs) | |
| return _orig_template_response(self, *args, **kwargs) | |
| _st_templating.Jinja2Templates.TemplateResponse = _compat_template_response | |
| import torch | |
| import torchaudio | |
| import gradio as gr | |
| import shutil | |
| from huggingface_hub import hf_hub_download | |
| from pydub import AudioSegment | |
| # ============================================================ | |
| # STEP 1: Download model files from the Urdu FT repo | |
| # ============================================================ | |
| MODEL_REPO = "suhaibrashid17/XTTS-v2-Urdu-FT" | |
| CACHE_DIR = "./model_cache" | |
| os.makedirs(CACHE_DIR, exist_ok=True) | |
| print("Downloading model files...") | |
| model_path = hf_hub_download(repo_id=MODEL_REPO, filename="model.pth", cache_dir=CACHE_DIR) | |
| config_path = hf_hub_download(repo_id=MODEL_REPO, filename="config.json", cache_dir=CACHE_DIR) | |
| vocab_path = hf_hub_download(repo_id=MODEL_REPO, filename="vocab.json", cache_dir=CACHE_DIR) | |
| # Note: file in repo is "tokenizer.py" (singular) | |
| # but we patch the TTS library's "tokenizers.py" (plural) | |
| tokenizer_path = hf_hub_download(repo_id=MODEL_REPO, filename="tokenizer.py", cache_dir=CACHE_DIR) | |
| print("Downloaded all files.") | |
| # ============================================================ | |
| # STEP 2: Find and patch the TTS library's tokenizer | |
| # ============================================================ | |
| import TTS | |
| tts_path = os.path.dirname(TTS.__file__) | |
| print(f"TTS package location: {tts_path}") | |
| # Diagnostic: find where the tokenizer actually lives in this TTS version | |
| print("Searching for tokenizer files in TTS package...") | |
| found_tokenizers = [] | |
| for root, dirs, files in os.walk(tts_path): | |
| for f in files: | |
| if "token" in f.lower() and f.endswith(".py"): | |
| full_path = os.path.join(root, f) | |
| found_tokenizers.append(full_path) | |
| print(f" Found: {full_path}") | |
| # Try the most likely locations in order | |
| candidate_paths = [ | |
| os.path.join(tts_path, "tts", "layers", "xtts", "tokenizers.py"), | |
| os.path.join(tts_path, "tts", "layers", "xtts", "tokenizer.py"), | |
| ] | |
| target_tokenizer = None | |
| for candidate in candidate_paths: | |
| if os.path.exists(candidate): | |
| target_tokenizer = candidate | |
| break | |
| # If neither standard path exists, use whatever we found in the search | |
| if target_tokenizer is None and found_tokenizers: | |
| # Prefer any xtts-related tokenizer | |
| xtts_tokenizers = [p for p in found_tokenizers if "xtts" in p] | |
| if xtts_tokenizers: | |
| target_tokenizer = xtts_tokenizers[0] | |
| else: | |
| target_tokenizer = found_tokenizers[0] | |
| if target_tokenizer is None: | |
| raise FileNotFoundError( | |
| f"No tokenizer.py found anywhere in {tts_path}. " | |
| f"TTS package may not be installed correctly." | |
| ) | |
| print(f"Patching tokenizer at: {target_tokenizer}") | |
| shutil.copy(tokenizer_path, target_tokenizer) | |
| print("Tokenizer patched successfully.") | |
| # ============================================================ | |
| # STEP 3: Load the model (AFTER patching the tokenizer) | |
| # ============================================================ | |
| # Patch torch.load to allow full pickle loading (XTTS checkpoints need this) | |
| _original_load = torch.load | |
| def patched_load(*args, **kwargs): | |
| kwargs['weights_only'] = False | |
| return _original_load(*args, **kwargs) | |
| torch.load = patched_load | |
| from TTS.tts.configs.xtts_config import XttsConfig | |
| from TTS.tts.models.xtts import Xtts | |
| device = "cuda:0" if torch.cuda.is_available() else "cpu" | |
| print(f"Using device: {device}") | |
| config = XttsConfig() | |
| config.load_json(config_path) | |
| XTTS_MODEL = Xtts.init_from_config(config) | |
| # The newer coqui-tts expects a checkpoint_dir with a speakers_xtts.pth file. | |
| # The Urdu finetuned model doesn't ship one, so we set up a directory with | |
| # all the model files in one place and pass checkpoint_dir explicitly. | |
| checkpoint_dir = os.path.dirname(model_path) | |
| print(f"Checkpoint directory: {checkpoint_dir}") | |
| print(f"Files in checkpoint dir: {os.listdir(checkpoint_dir)}") | |
| # Create an empty speakers_xtts.pth if missing (the Urdu FT model doesn't include one) | |
| speakers_file = os.path.join(checkpoint_dir, "speakers_xtts.pth") | |
| if not os.path.exists(speakers_file): | |
| print(f"Creating empty speakers file at: {speakers_file}") | |
| torch.save({}, speakers_file) | |
| # Ensure vocab.json is in the checkpoint_dir too (it might be in a different cache subfolder) | |
| import shutil | |
| target_vocab = os.path.join(checkpoint_dir, "vocab.json") | |
| if not os.path.exists(target_vocab): | |
| print(f"Copying vocab to checkpoint dir: {target_vocab}") | |
| shutil.copy(vocab_path, target_vocab) | |
| # Convert vocab.json to the schema that tokenizers==0.15.2 expects. | |
| # The Urdu FT vocab.json was saved by tokenizers >= 0.20, which: | |
| # - added the BPE "ignore_merges" field (unknown to 0.15.2) | |
| # - changed "merges" from ["a b", ...] strings to [["a","b"], ...] arrays | |
| # Both break Rust serde with "data did not match any variant of untagged enum ModelWrapper". | |
| import json | |
| patched_vocab = os.path.join(checkpoint_dir, "vocab_patched.json") | |
| with open(target_vocab, "r", encoding="utf-8") as f: | |
| vocab_data = json.load(f) | |
| if isinstance(vocab_data.get("model"), dict): | |
| model = vocab_data["model"] | |
| removed = [k for k in ("ignore_merges",) if model.pop(k, None) is not None] | |
| if removed: | |
| print(f"Stripped unsupported fields from vocab.json model: {removed}") | |
| merges = model.get("merges") | |
| if isinstance(merges, list) and merges and isinstance(merges[0], list): | |
| model["merges"] = [f"{a} {b}" for a, b in merges] | |
| print(f"Converted {len(merges)} merges from [a,b] arrays to 'a b' strings") | |
| with open(patched_vocab, "w", encoding="utf-8") as f: | |
| json.dump(vocab_data, f, ensure_ascii=False) | |
| vocab_path = patched_vocab | |
| # Ensure config.json is there too | |
| target_config = os.path.join(checkpoint_dir, "config.json") | |
| if not os.path.exists(target_config): | |
| shutil.copy(config_path, target_config) | |
| XTTS_MODEL.load_checkpoint( | |
| config, | |
| checkpoint_path=model_path, | |
| checkpoint_dir=checkpoint_dir, | |
| vocab_path=vocab_path, | |
| use_deepspeed=False, | |
| eval=True, | |
| ) | |
| XTTS_MODEL.to(device) | |
| print("Model loaded successfully.") | |
| # ============================================================ | |
| # STEP 4: Helper functions | |
| # ============================================================ | |
| def convert_audio_to_wav(input_path): | |
| """Convert any audio format to 22050Hz mono WAV.""" | |
| output_path = input_path.rsplit(".", 1)[0] + "_converted.wav" | |
| audio = AudioSegment.from_file(input_path) | |
| audio = audio.set_channels(1).set_frame_rate(22050) | |
| audio.export(output_path, format="wav") | |
| return output_path | |
| def generate_speech(text, reference_audio, temperature, top_p): | |
| """Generate Urdu speech in the reference voice.""" | |
| if not text or not text.strip(): | |
| return None, "Please enter some Urdu text." | |
| if reference_audio is None: | |
| return None, "Please upload a reference voice clip." | |
| try: | |
| # Convert reference to WAV if needed | |
| if not reference_audio.endswith(".wav"): | |
| reference_audio = convert_audio_to_wav(reference_audio) | |
| # Extract speaker conditioning from reference clip | |
| gpt_cond_latent, speaker_embedding = XTTS_MODEL.get_conditioning_latents( | |
| audio_path=[reference_audio], | |
| gpt_cond_len=XTTS_MODEL.config.gpt_cond_len, | |
| max_ref_length=XTTS_MODEL.config.max_ref_len, | |
| sound_norm_refs=XTTS_MODEL.config.sound_norm_refs, | |
| ) | |
| # Run inference | |
| result = XTTS_MODEL.inference( | |
| text=text, | |
| language="ur", | |
| gpt_cond_latent=gpt_cond_latent, | |
| speaker_embedding=speaker_embedding, | |
| temperature=temperature, | |
| length_penalty=0.1, | |
| repetition_penalty=10.0, | |
| top_k=10, | |
| top_p=top_p, | |
| ) | |
| wav = torch.tensor(result["wav"]).unsqueeze(0).cpu() | |
| output_path = "output.wav" | |
| torchaudio.save(output_path, wav, 24000) | |
| return output_path, "Generated successfully." | |
| except Exception as e: | |
| import traceback | |
| traceback.print_exc() | |
| return None, f"Error: {str(e)}" | |
| # ============================================================ | |
| # STEP 5: Gradio UI | |
| # ============================================================ | |
| with gr.Blocks(title="Urdu Voice Clone") as demo: | |
| gr.Markdown("# Urdu Voice Cloning with XTTS-v2") | |
| gr.Markdown( | |
| "Upload a 30-60 second clean audio clip of the target voice, " | |
| "then enter Urdu text in **Nastaliq script** (not Roman Urdu)." | |
| ) | |
| with gr.Row(): | |
| with gr.Column(): | |
| text_input = gr.Textbox( | |
| label="Urdu text (Nastaliq)", | |
| placeholder="کیا حال ہے؟ آج موسم بہت اچھا ہے۔", | |
| lines=4, | |
| rtl=True, | |
| ) | |
| reference_audio = gr.Audio( | |
| label="Reference voice clip (30-60s, clean audio)", | |
| type="filepath", | |
| ) | |
| with gr.Accordion("Advanced settings", open=False): | |
| temperature = gr.Slider(0.1, 1.0, value=0.3, step=0.05, label="Temperature") | |
| top_p = gr.Slider(0.1, 1.0, value=0.3, step=0.05, label="Top-p") | |
| generate_btn = gr.Button("Generate", variant="primary") | |
| with gr.Column(): | |
| output_audio = gr.Audio(label="Generated speech") | |
| status = gr.Textbox(label="Status", interactive=False) | |
| generate_btn.click( | |
| fn=generate_speech, | |
| inputs=[text_input, reference_audio, temperature, top_p], | |
| outputs=[output_audio, status], | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() |