1import torch
2from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
3from snac import SNAC
4import soundfile as sf
5
6# Model configuration for 4-bit inference
7quantization_config = BitsAndBytesConfig(
8 load_in_4bit=True,
9 bnb_4bit_quant_type="nf4",
10 bnb_4bit_compute_dtype=torch.bfloat16,
11 bnb_4bit_use_double_quant=True,
12)
13
14# Load model and tokenizer
15model = AutoModelForCausalLM.from_pretrained(
16 "maya-research/veena-tts",
17 quantization_config=quantization_config,
18 device_map="auto",
19 trust_remote_code=True,
20)
21tokenizer = AutoTokenizer.from_pretrained("maya-research/veena-tts", trust_remote_code=True)
22
23# Initialize SNAC decoder
24snac_model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").eval().cuda()
25
26# Control token IDs (fixed for Veena)
27START_OF_SPEECH_TOKEN = 128257
28END_OF_SPEECH_TOKEN = 128258
29START_OF_HUMAN_TOKEN = 128259
30END_OF_HUMAN_TOKEN = 128260
31START_OF_AI_TOKEN = 128261
32END_OF_AI_TOKEN = 128262
33AUDIO_CODE_BASE_OFFSET = 128266
34
35# Available speakers
36speakers = ["kavya", "agastya", "maitri", "vinaya"]
37
38def generate_speech(text, speaker="kavya", temperature=0.4, top_p=0.9):
39 """Generate speech from text using specified speaker voice"""
40
41 # Prepare input with speaker token
42 prompt = f"<spk_{speaker}> {text}"
43 prompt_tokens = tokenizer.encode(prompt, add_special_tokens=False)
44
45 # Construct full sequence: [HUMAN] <spk_speaker> text [/HUMAN] [AI] [SPEECH]
46 input_tokens = [
47 START_OF_HUMAN_TOKEN,
48 *prompt_tokens,
49 END_OF_HUMAN_TOKEN,
50 START_OF_AI_TOKEN,
51 START_OF_SPEECH_TOKEN
52 ]
53
54 input_ids = torch.tensor([input_tokens], device=model.device)
55
56 # Calculate max tokens based on text length
57 max_tokens = min(int(len(text) * 1.3) * 7 + 21, 700)
58
59 # Generate audio tokens
60 with torch.no_grad():
61 output = model.generate(
62 input_ids,
63 max_new_tokens=max_tokens,
64 do_sample=True,
65 temperature=temperature,
66 top_p=top_p,
67 repetition_penalty=1.05,
68 pad_token_id=tokenizer.pad_token_id,
69 eos_token_id=[END_OF_SPEECH_TOKEN, END_OF_AI_TOKEN]
70 )
71
72 # Extract SNAC tokens
73 generated_ids = output[0][len(input_tokens):].tolist()
74 snac_tokens = [
75 token_id for token_id in generated_ids
76 if AUDIO_CODE_BASE_OFFSET <= token_id < (AUDIO_CODE_BASE_OFFSET + 7 * 4096)
77 ]
78
79 if not snac_tokens:
80 raise ValueError("No audio tokens generated")
81
82 # Decode audio
83 audio = decode_snac_tokens(snac_tokens, snac_model)
84 return audio
85
86def decode_snac_tokens(snac_tokens, snac_model):
87 """De-interleave and decode SNAC tokens to audio"""
88 if not snac_tokens or len(snac_tokens) % 7 != 0:
89 return None
90
91 # Get the device of the SNAC model. Fixed by Shresth to run on colab notebook :)
92 snac_device = next(snac_model.parameters()).device
93
94 # De-interleave tokens into 3 hierarchical levels
95 codes_lvl = [[] for _ in range(3)]
96 llm_codebook_offsets = [AUDIO_CODE_BASE_OFFSET + i * 4096 for i in range(7)]
97
98 for i in range(0, len(snac_tokens), 7):
99 # Level 0: Coarse (1 token)
100 codes_lvl[0].append(snac_tokens[i] - llm_codebook_offsets[0])
101 # Level 1: Medium (2 tokens)
102 codes_lvl[1].append(snac_tokens[i+1] - llm_codebook_offsets[1])
103 codes_lvl[1].append(snac_tokens[i+4] - llm_codebook_offsets[4])
104 # Level 2: Fine (4 tokens)
105 codes_lvl[2].append(snac_tokens[i+2] - llm_codebook_offsets[2])
106 codes_lvl[2].append(snac_tokens[i+3] - llm_codebook_offsets[3])
107 codes_lvl[2].append(snac_tokens[i+5] - llm_codebook_offsets[5])
108 codes_lvl[2].append(snac_tokens[i+6] - llm_codebook_offsets[6])
109
110 # Convert to tensors for SNAC decoder
111 hierarchical_codes = []
112 for lvl_codes in codes_lvl:
113 tensor = torch.tensor(lvl_codes, dtype=torch.int32, device=snac_device).unsqueeze(0)
114 if torch.any((tensor < 0) | (tensor > 4095)):
115 raise ValueError("Invalid SNAC token values")
116 hierarchical_codes.append(tensor)
117
118 # Decode with SNAC
119 with torch.no_grad():
120 audio_hat = snac_model.decode(hierarchical_codes)
121
122 return audio_hat.squeeze().clamp(-1, 1).cpu().numpy()
123
124# --- Example Usage ---
125
126# Hindi
127text_hindi = "आज मैंने एक नई तकनीक के बारे में सीखा जो कृत्रिम बुद्धिमत्ता का उपयोग करके मानव जैसी आवाज़ उत्पन्न कर सकती है।"
128audio = generate_speech(text_hindi, speaker="kavya")
129sf.write("output_hindi_kavya.wav", audio, 24000)
130
131# English
132text_english = "Today I learned about a new technology that uses artificial intelligence to generate human-like voices."
133audio = generate_speech(text_english, speaker="agastya")
134sf.write("output_english_agastya.wav", audio, 24000)
135
136# Code-mixed
137text_mixed = "मैं तो पूरा presentation prepare कर चुका हूं! कल रात को ही मैंने पूरा code base चेक किया।"
138audio = generate_speech(text_mixed, speaker="maitri")
139sf.write("output_mixed_maitri.wav", audio, 24000)