Views
No views yet
1import os
2from snac import SNAC
3from pathlib import Path
4import torch
5from transformers import AutoModelForCausalLM, Trainer, TrainingArguments, AutoTokenizer,BitsAndBytesConfig
6from huggingface_hub import snapshot_download
7import librosa
8import numpy as np
9from scipy.io.wavfile import write
10import torchaudio
11from flask import Flask, jsonify, request
12
13modelLocalPath="Cosmobillian/turkish_orpheus_tts"
14
15
16def load_orpheus_tokenizer(model_id: str = modelLocalPath) -> AutoTokenizer:
17 tokenizer = AutoTokenizer.from_pretrained(model_id,local_files_only=True, device_map="cuda")
18 return tokenizer
19
20def load_snac():
21 snac_model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz")
22 return snac_model
23
24def load_orpheus_auto_model(model_id: str = modelLocalPath):
25 model = AutoModelForCausalLM.from_pretrained(model_id, torch_dtype=torch.bfloat16,local_files_only=True, device_map="cuda")
26 model.cuda()
27 return model
28
29
30
31def tokenize_audio(audio_file_path, snac_model):
32 audio_array, sample_rate = librosa.load(audio_file_path, sr=24000)
33 waveform = torch.from_numpy(audio_array).unsqueeze(0)
34 waveform = waveform.to(dtype=torch.float32)
35
36 waveform = waveform.unsqueeze(0)
37
38 with torch.inference_mode():
39 codes = snac_model.encode(waveform)
40
41 all_codes = []
42 for i in range(codes[0].shape[1]):
43 all_codes.append(codes[0][0][i].item() + 128266)
44 all_codes.append(codes[1][0][2 * i].item() + 128266 + 4096)
45 all_codes.append(codes[2][0][4 * i].item() + 128266 + (2 * 4096))
46 all_codes.append(codes[2][0][(4 * i) + 1].item() + 128266 + (3 * 4096))
47 all_codes.append(codes[1][0][(2 * i) + 1].item() + 128266 + (4 * 4096))
48 all_codes.append(codes[2][0][(4 * i) + 2].item() + 128266 + (5 * 4096))
49 all_codes.append(codes[2][0][(4 * i) + 3].item() + 128266 + (6 * 4096))
50
51 return all_codes
52
53
54def prepare_inputs(
55 fpath_audio_ref,
56 audio_ref_transcript: str,
57 text_prompts: list[str],
58 snac_model,
59 tokenizer,
60):
61
62
63 start_tokens = torch.tensor([[128259]], dtype=torch.int64)
64 end_tokens = torch.tensor([[128009, 128260, 128261, 128257]], dtype=torch.int64)
65 final_tokens = torch.tensor([[128258, 128262]], dtype=torch.int64)
66
67
68 all_modified_input_ids = []
69 for prompt in text_prompts:
70 input_ids = tokenizer(prompt, return_tensors="pt").input_ids
71 #second_input_ids = torch.cat([zeroprompt_input_ids, start_tokens, input_ids, end_tokens], dim=1)
72 second_input_ids = torch.cat([start_tokens, input_ids, end_tokens], dim=1)
73 all_modified_input_ids.append(second_input_ids)
74
75 all_padded_tensors = []
76 all_attention_masks = []
77 max_length = max([modified_input_ids.shape[1] for modified_input_ids in all_modified_input_ids])
78
79 for modified_input_ids in all_modified_input_ids:
80 padding = max_length - modified_input_ids.shape[1]
81 padded_tensor = torch.cat([torch.full((1, padding), 128263, dtype=torch.int64), modified_input_ids], dim=1)
82 attention_mask = torch.cat([torch.zeros((1, padding), dtype=torch.int64),
83 torch.ones((1, modified_input_ids.shape[1]), dtype=torch.int64)], dim=1)
84 all_padded_tensors.append(padded_tensor)
85 all_attention_masks.append(attention_mask)
86
87 all_padded_tensors = torch.cat(all_padded_tensors, dim=0)
88 all_attention_masks = torch.cat(all_attention_masks, dim=0)
89
90 input_ids = all_padded_tensors.to("cuda")
91 attention_mask = all_attention_masks.to("cuda")
92 return input_ids, attention_mask
93
94
95
96def inference(model, input_ids, attention_mask):
97 with torch.no_grad():
98 generated_ids = model.generate(
99 input_ids=input_ids,
100 attention_mask=attention_mask,
101 max_new_tokens=2048,
102 do_sample=True,
103 temperature=0.2,
104 top_k=10,
105 top_p=0.9,
106 repetition_penalty=1.9,
107 num_return_sequences=1,
108 eos_token_id=128258,
109
110 )
111
112 generated_ids = torch.cat([generated_ids, torch.tensor([[128262]]).to("cuda")], dim=1) # EOAI
113
114 return generated_ids
115
116
117def convert_tokens_to_speech(generated_ids, snac_model):
118 token_to_find = 128257
119 token_to_remove = 128258
120 token_indices = (generated_ids == token_to_find).nonzero(as_tuple=True)
121
122 if len(token_indices[1]) > 0:
123 last_occurrence_idx = token_indices[1][-1].item()
124 cropped_tensor = generated_ids[:, last_occurrence_idx + 1:]
125 else:
126 cropped_tensor = generated_ids
127
128 _mask = cropped_tensor != token_to_remove
129 processed_rows = []
130 for row in cropped_tensor:
131 masked_row = row[row != token_to_remove]
132 processed_rows.append(masked_row)
133
134 code_lists = []
135 for row in processed_rows:
136 row_length = row.size(0)
137 new_length = (row_length // 7) * 7
138 trimmed_row = row[:new_length]
139 trimmed_row = [t - 128266 for t in trimmed_row]
140 code_lists.append(trimmed_row)
141
142 my_samples = []
143 for code_list in code_lists:
144 samples = redistribute_codes(code_list, snac_model)
145 my_samples.append(samples)
146
147 return my_samples
148
149
150def redistribute_codes(code_list, snac_model):
151 layer_1 = []
152 layer_2 = []
153 layer_3 = []
154
155 for i in range((len(code_list) + 1) // 7):
156 layer_1.append(code_list[7 * i])
157 layer_2.append(code_list[7 * i + 1] - 4096)
158 layer_3.append(code_list[7 * i + 2] - (2 * 4096))
159 layer_3.append(code_list[7 * i + 3] - (3 * 4096))
160 layer_2.append(code_list[7 * i + 4] - (4 * 4096))
161 layer_3.append(code_list[7 * i + 5] - (5 * 4096))
162 layer_3.append(code_list[7 * i + 6] - (6 * 4096))
163
164 codes = [
165 torch.tensor(layer_1).unsqueeze(0),
166 torch.tensor(layer_2).unsqueeze(0),
167 torch.tensor(layer_3).unsqueeze(0)
168 ]
169 audio_hat = snac_model.decode(codes)
170 return audio_hat
171
172
173def to_wav_from(samples: list) -> list[np.ndarray]:
174 """Converts a list of PyTorch tensors (or NumPy arrays) to NumPy arrays."""
175 processed_samples = []
176
177 for s in samples:
178 if isinstance(s, torch.Tensor):
179 s = s.detach().squeeze().to('cpu').numpy()
180 else:
181 s = np.squeeze(s)
182
183 processed_samples.append(s)
184
185 return processed_samples
186
187
188def zero_shot_tts(fpath_audio_ref, audio_ref_transcript, texts: list[str], model, snac_model, tokenizer):
189 print(f"fpath_audio_ref {fpath_audio_ref}")
190 print(f"audio_ref_transcript {audio_ref_transcript}")
191 print(f"texts {texts}")
192 inp_ids, attn_mask = prepare_inputs(fpath_audio_ref, audio_ref_transcript, texts, snac_model, tokenizer)
193 print(f"input_id_len:{len(inp_ids)}")
194 gen_ids = inference(model, inp_ids, attn_mask)
195 samples = convert_tokens_to_speech(gen_ids, snac_model)
196 wav_forms = to_wav_from(samples)
197 return wav_forms
198
199
200def save_wav(samples: list[np.array], sample_rate: int, filenames: list[str]):
201 """ Saves a list of tensors as .wav files.
202
203 Args:
204 samples (list[torch.Tensor]): List of audio tensors.
205 sample_rate (int): Sample rate in Hz.
206 filenames (list[str]): List of filenames to save.
207 """
208 wav_data = to_wav_from(samples)
209
210 for data, filename in zip(wav_data, filenames):
211 write(filename, sample_rate, data.astype(np.float32))
212 print(f"saved to {filename}")
213
214
215def get_ref_audio_and_transcript(root_folder: str):
216 root_path = Path(root_folder)
217 print(f"root_path {root_path}")
218 out = []
219 for speaker_folder in root_path.iterdir():
220 if speaker_folder.is_dir(): # Ensure it's a directory
221 wav_files = list(speaker_folder.glob("*.wav"))
222 txt_files = list(speaker_folder.glob("*.txt"))
223
224 if wav_files and txt_files:
225 ref_audio = wav_files[0] # Assume only one .wav file per folder
226 transcript = txt_files[0].read_text(encoding="utf-8").strip()
227 out.append((ref_audio, transcript))
228
229 return out
230
231app = Flask(__name__)
232
233
234@app.route('/generate', methods=['POST'])
235def generate():
236 content = request.json
237 process_data(content)
238 rresponse = {
239 'received': content,
240 'status': 'success'
241 }
242 response= jsonify(rresponse)
243 response.headers['Content-Type'] = 'application/json; charset=utf-8'
244 return response
245
246
247
248def process_data(jsonText):
249 texts = [f"{jsonText['text']}"]
250 #print(f"texts:{texts}")
251 #print(f"prompt_pairs:{prompt_pairs}")
252 for fpath_audio, audio_transcript in prompt_pairs:
253 print(f"zero shot: {fpath_audio} {audio_transcript}")
254 wav_forms = zero_shot_tts(fpath_audio, audio_transcript, texts, model, snac_model, tokenizer)
255
256 import os
257 from pathlib import Path
258 from datetime import datetime
259 out_dir = Path(fpath_audio).parent / "inference"
260 #print(f"out_dir:{out_dir}")
261 out_dir.mkdir(parents=True, exist_ok=True) #
262 timestamp_str = str(int(datetime.now().timestamp()))
263 file_names = [f"{out_dir.as_posix()}/{Path(fpath_audio).stem}_{i}_{timestamp_str}.wav" for i, t in enumerate(texts)]
264 #print(f"file_names:{file_names}")
265 save_wav(wav_forms, 24000, file_names)
266
267
268
269if __name__ == "__main__":
270 tokenizer = load_orpheus_tokenizer()
271 model = load_orpheus_auto_model()
272 snac_model = load_snac()
273 prompt_pairs = get_ref_audio_and_transcript("D:\\AI_APPS\\Orpheus-TTS\\data")
274 print(f"snac_model loaded")
275 app.run(debug=True,port=5400)