Views
No views yet
q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj1import torch
2import torch.nn as nn
3from huggingface_hub import hf_hub_download
4from chatterbox.mtl_tts import ChatterboxMultilingualTTS
5
6# 1. Define LoRA mapping helper
7class LoRALayer(nn.Module):
8 def __init__(self, original, rank=32, alpha=64.0, dropout=0.05):
9 super().__init__()
10 self.original_layer = original
11 self.scaling = alpha / rank
12 in_f, out_f = original.in_features, original.out_features
13 dev, dt = original.weight.device, original.weight.dtype
14 self.lora_A = nn.Parameter(torch.zeros(rank, in_f, device=dev, dtype=dt))
15 self.lora_B = nn.Parameter(torch.zeros(out_f, rank, device=dev, dtype=dt))
16 self.lora_dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
17 for p in self.original_layer.parameters():
18 p.requires_grad = False
19
20 def forward(self, x):
21 return self.original_layer(x) + (self.lora_dropout(x) @ self.lora_A.T @ self.lora_B.T * self.scaling)
22
23# 2. Load Base Model
24device = "cuda" if torch.cuda.is_available() else "cpu"
25model = ChatterboxMultilingualTTS.from_pretrained(device=device)
26
27# 3. Download and Inject LoRA
28lora_path = hf_hub_download(repo_id="amanuelbyte/chatterbox-fr-lora-v2", filename="best_lora_adapter.pt")
29payload = torch.load(lora_path, map_location=device, weights_only=True)
30lora_sd = payload.get("lora_state_dict", payload)
31
32targets = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]
33for name, module in model.t3.named_modules():
34 if isinstance(module, nn.Linear) and any(x in name for x in targets):
35 parent_name, child_name = ".".join(name.split(".")[:-1]), name.split(".")[-1]
36 parent = model.t3.get_submodule(parent_name)
37 setattr(parent, child_name, LoRALayer(module))
38
39# Load the weights into the injected layers
40current_sd = model.t3.state_dict()
41for key, value in lora_sd.items():
42 if key in current_sd:
43 current_sd[key] = value.to(device)
44model.t3.load_state_dict(current_sd, strict=False)
45model.t3.eval()
46
47# 4. Generate Audio
48import torchaudio
49text = "Bonjour, bienvenue à cette démonstration de clonage vocal cross-lingue!"
50ref_audio_path = "path/to/english_voice.wav"
51
52with torch.inference_mode():
53 wav = model.generate(text, audio_prompt_path=ref_audio_path, language_id="fr")
54
55torchaudio.save("output_french.wav", wav.cpu().unsqueeze(0), 16000)