Views
No views yet
canopylabs/3b-hi-pretrain-research_releaseIt behaves like a speaker who has learned to speak Pashto fairly recently but understandable overall, but still prone to pronunciation mistakes.
canopylabs/3b-hi-pretrain-research_releaseRTX 3090 24GB4~3.895e-50.030.01116153650010bf16=True, fp16=False32640.0q_projk_projv_projo_projgate_projdown_projup_projlm_headembed_tokensconfig.yaml1# ----------------------------
2# SNAC
3# ----------------------------
4snac_model_name: "hubertsiuzdak/snac_24khz"
5target_sample_rate: 24000
6
7# ----------------------------
8# Tokenizer / model
9# ----------------------------
10tokenizer_name: "canopylabs/3b-hi-pretrain-research_release"
11output_dir: "E:\\ORPHEUS MODEL"
12
13# ----------------------------
14# Precision / resize
15# ----------------------------
16bf16: true
17force_resize_embeddings_if_needed: true
18
19# ----------------------------
20# Special token IDs
21# ----------------------------
22start_of_text: 128000
23end_of_text: 128009
24
25start_of_speech: 128257
26end_of_speech: 128258
27
28start_of_human: 128259
29end_of_human: 128260
30
31start_of_ai: 128261
32end_of_ai: 128262
33pad_token: 128263
34
35audio_tokens_start: 128266
36
37
38## Inference
39
40### Expected folder layout
41
42```text
43your_model_folder/
44├─ README.md
45├─ config.yaml
46├─ infer.py
47└─ merged/
48 ├─ config.json
49 ├─ model.safetensors / pytorch_model.bin
50 ├─ tokenizer files
51 └─ generation config files if presentoutput_dir/mergedinfer.pyconfig.yamlpython infer.pyدا یو ازمایښتي جمله ده چې د پښورۍ پښتو لپاره د وینا جوړولو ازموینه پرې وشي1import os
2import re
3import yaml
4import torch
5import numpy as np
6import soundfile as sf
7
8from snac import SNAC
9from transformers import AutoModelForCausalLM, AutoTokenizer
10
11
12def load_config(path="config.yaml"):
13 with open(path, "r", encoding="utf-8") as f:
14 return yaml.safe_load(f)
15
16
17def load_snac_model(cfg):
18 device = "cuda" if torch.cuda.is_available() else "cpu"
19
20 print(f"Loading SNAC decoder: {cfg['snac_model_name']}")
21 snac_model = SNAC.from_pretrained(cfg["snac_model_name"]).to(device)
22 snac_model.eval()
23 return snac_model
24
25
26def maybe_fix_tokenizer_and_resize(model, tokenizer, cfg):
27 """
28 Keeps tokenizer/model embedding size compatible with reserved audio token IDs.
29 """
30 pad_token = cfg["pad_token"]
31
32 max_audio_token_id = cfg["audio_tokens_start"] + 7 * 4096 - 1
33 special_ids = [
34 cfg["start_of_text"],
35 cfg["end_of_text"],
36 cfg["start_of_speech"],
37 cfg["end_of_speech"],
38 cfg["start_of_human"],
39 cfg["end_of_human"],
40 cfg["start_of_ai"],
41 cfg["end_of_ai"],
42 cfg["pad_token"],
43 ]
44 max_token_id_needed = max(max_audio_token_id, max(special_ids))
45
46 tokenizer_len = len(tokenizer)
47 embedding_rows = model.get_input_embeddings().num_embeddings
48
49 print(f"Tokenizer length before fix: {tokenizer_len}")
50 print(f"Embedding rows before fix: {embedding_rows}")
51 print(f"Max token id needed: {max_token_id_needed}")
52
53 if embedding_rows <= max_token_id_needed:
54 required_vocab_size = max_token_id_needed + 1
55
56 if tokenizer_len < required_vocab_size:
57 extra = required_vocab_size - tokenizer_len
58 print(f"Tokenizer too small. Adding {extra} placeholder tokens...")
59 tokenizer.add_tokens([f"<orpheus_extra_{i}>" for i in range(extra)])
60
61 print(f"Resizing embeddings to {len(tokenizer)}")
62 model.resize_token_embeddings(len(tokenizer))
63
64 print(f"Tokenizer length after fix: {len(tokenizer)}")
65 print(f"Embedding rows after fix: {model.get_input_embeddings().num_embeddings}")
66 else:
67 print("No resize needed; model embeddings already cover required token ids.")
68
69 if tokenizer.pad_token is None:
70 tokenizer.pad_token = tokenizer.eos_token
71
72 model.config.pad_token_id = pad_token
73
74
75def load_merged_model(cfg, tokenizer):
76 merged_dir = os.path.join(cfg["output_dir"], "merged")
77 dtype = torch.bfloat16 if cfg.get("bf16", True) else torch.float16
78
79 if not os.path.isdir(merged_dir):
80 raise FileNotFoundError(
81 f"Merged model folder not found: {merged_dir}\n"
82 f"Run training first so it saves:\n"
83 f" {merged_dir}"
84 )
85
86 try:
87 model = AutoModelForCausalLM.from_pretrained(
88 merged_dir,
89 torch_dtype=dtype,
90 attn_implementation="flash_attention_2",
91 device_map="auto",
92 )
93 except Exception:
94 print("flash_attention_2 not available; falling back to sdpa")
95 model = AutoModelForCausalLM.from_pretrained(
96 merged_dir,
97 torch_dtype=dtype,
98 attn_implementation="sdpa",
99 device_map="auto",
100 )
101
102 if cfg.get("force_resize_embeddings_if_needed", True):
103 maybe_fix_tokenizer_and_resize(model, tokenizer, cfg)
104 else:
105 if tokenizer.pad_token is None:
106 tokenizer.pad_token = tokenizer.eos_token
107 model.config.pad_token_id = cfg["pad_token"]
108
109 model.eval()
110 return model
111
112
113def build_prompt_input_ids(text, tokenizer, cfg):
114 """
115 Must match training prompt format.
116 """
117 start_of_human = cfg["start_of_human"]
118 end_of_human = cfg["end_of_human"]
119 start_of_ai = cfg["start_of_ai"]
120 start_of_speech = cfg["start_of_speech"]
121 end_of_text = cfg["end_of_text"]
122
123 text_ids = tokenizer.encode(text, add_special_tokens=True)
124
125 input_ids = (
126 [start_of_human]
127 + text_ids
128 + [end_of_text]
129 + [end_of_human]
130 + [start_of_ai]
131 + [start_of_speech]
132 )
133 return input_ids
134
135
136def extract_audio_tokens(generated_ids, cfg):
137 """
138 Extracts generated speech tokens and maps them back to SNAC code indices.
139 """
140 start_of_speech = cfg["start_of_speech"]
141 end_of_speech = cfg["end_of_speech"]
142 audio_tokens_start = cfg["audio_tokens_start"]
143
144 try:
145 speech_start_idx = generated_ids.index(start_of_speech) + 1
146 except ValueError:
147 raise ValueError("start_of_speech token not found in generated sequence.")
148
149 speech_tokens = []
150 for tok in generated_ids[speech_start_idx:]:
151 if tok == end_of_speech:
152 break
153 if tok >= audio_tokens_start:
154 speech_tokens.append(tok)
155
156 if not speech_tokens:
157 raise ValueError("No speech/audio tokens were generated.")
158
159 snac_codes = []
160 for i, tok in enumerate(speech_tokens):
161 band = i % 7
162 code = tok - audio_tokens_start - (band * 4096)
163 snac_codes.append(code)
164
165 usable = (len(snac_codes) // 7) * 7
166 snac_codes = snac_codes[:usable]
167
168 if len(snac_codes) < 7:
169 raise ValueError("Too few usable SNAC codes after cleanup.")
170
171 cleaned = []
172 for c in snac_codes:
173 if c < 0:
174 c = 0
175 elif c > 4095:
176 c = 4095
177 cleaned.append(c)
178
179 return cleaned
180
181
182def snac_codes_to_audio(snac_codes, snac_model):
183 """
184 Reverse the 7-code interleaving:
185 frame = [c0, c1a, c2a, c2b, c1b, c2c, c2d]
186 """
187 device = next(snac_model.parameters()).device
188
189 if len(snac_codes) % 7 != 0:
190 raise ValueError("snac_codes length must be divisible by 7.")
191
192 n_frames = len(snac_codes) // 7
193
194 codes_0 = []
195 codes_1 = []
196 codes_2 = []
197
198 for j in range(n_frames):
199 i = 7 * j
200 codes_0.append(snac_codes[i + 0])
201
202 codes_1.append(snac_codes[i + 1])
203 codes_1.append(snac_codes[i + 4])
204
205 codes_2.append(snac_codes[i + 2])
206 codes_2.append(snac_codes[i + 3])
207 codes_2.append(snac_codes[i + 5])
208 codes_2.append(snac_codes[i + 6])
209
210 codes = [
211 torch.tensor(codes_0, dtype=torch.int32, device=device).unsqueeze(0),
212 torch.tensor(codes_1, dtype=torch.int32, device=device).unsqueeze(0),
213 torch.tensor(codes_2, dtype=torch.int32, device=device).unsqueeze(0),
214 ]
215
216 with torch.inference_mode():
217 audio_hat = snac_model.decode(codes)
218
219 audio = audio_hat.squeeze().detach().float().cpu().numpy()
220 audio = np.clip(audio, -1.0, 1.0)
221 return audio
222
223
224def generate_one(
225 model,
226 tokenizer,
227 snac_model,
228 text,
229 cfg,
230 output_wav_path="orpheus_test.wav",
231 temperature=0.7,
232 top_p=0.9,
233 repetition_penalty=1.1,
234 max_new_tokens=2560,
235):
236 input_ids = build_prompt_input_ids(text, tokenizer, cfg)
237
238 device = next(model.parameters()).device
239 input_ids = torch.tensor([input_ids], dtype=torch.long, device=device)
240 attention_mask = torch.ones_like(input_ids)
241
242 end_of_speech = cfg["end_of_speech"]
243 pad_token = cfg["pad_token"]
244
245 with torch.inference_mode():
246 output = model.generate(
247 input_ids=input_ids,
248 attention_mask=attention_mask,
249 max_new_tokens=max_new_tokens,
250 do_sample=True,
251 temperature=temperature,
252 top_p=top_p,
253 repetition_penalty=repetition_penalty,
254 eos_token_id=end_of_speech,
255 pad_token_id=pad_token,
256 )
257
258 generated_ids = output[0].detach().cpu().tolist()
259
260 snac_codes = extract_audio_tokens(generated_ids, cfg)
261 print(f"Generated {len(snac_codes)} SNAC codes")
262
263 audio = snac_codes_to_audio(snac_codes, snac_model)
264
265 sr = int(cfg["target_sample_rate"])
266 sf.write(output_wav_path, audio, sr)
267 print(f"Saved audio to: {output_wav_path}")
268
269 return {
270 "generated_ids": generated_ids,
271 "snac_codes": snac_codes,
272 "wav_path": output_wav_path,
273 "sample_rate": sr,
274 }
275
276
277def make_safe_filename(text, max_len=50):
278 text = re.sub(r"\s+", "_", text.strip())
279 text = re.sub(r'[\\/*?:"<>|]', "", text)
280 text = text[:max_len].strip("_")
281 return text if text else "sample"
282
283
284def main():
285 cfg = load_config()
286
287 test_texts = [
288 "دا یو ازمایښتي جمله ده چې د پښورۍ پښتو لپاره د وینا جوړولو ازموینه پرې وشي"
289 ]
290
291 output_dir = "pashto_test_outputs"
292 os.makedirs(output_dir, exist_ok=True)
293
294 tokenizer = AutoTokenizer.from_pretrained(cfg["tokenizer_name"])
295
296 print(f"Loading merged model from: {os.path.join(cfg['output_dir'], 'merged')}")
297 model = load_merged_model(cfg, tokenizer)
298
299 snac_model = load_snac_model(cfg)
300
301 results = []
302
303 for idx, test_text in enumerate(test_texts, start=1):
304 safe_name = make_safe_filename(test_text)
305 output_wav_path = os.path.join(output_dir, f"{idx:02d}_{safe_name}.wav")
306
307 print(f"\n[{idx}/{len(test_texts)}] Generating for: {test_text}")
308
309 result = generate_one(
310 model=model,
311 tokenizer=tokenizer,
312 snac_model=snac_model,
313 text=test_text,
314 cfg=cfg,
315 output_wav_path=output_wav_path,
316 temperature=0.7,
317 top_p=0.9,
318 repetition_penalty=1.1,
319 max_new_tokens=2560,
320 )
321
322 results.append(result)
323
324 print("\nDone.")
325 print(f"All WAVs saved in: {output_dir}")
326
327
328if __name__ == "__main__":
329 main()temperature = 0.7top_p = 0.9repetition_penalty = 1.1max_new_tokens = 2560temperaturetop_prepetition_penaltymax_new_tokens