Maarij-Aqeel's picture
Monkey-patch starlette TemplateResponse for gradio 4.44.0 compat
d189374
Raw History Blame Contribute Delete
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()