1import json
2import numpy as np
3import onnxruntime as ort
4from transformers import PreTrainedTokenizerFast
5
6# --- Load models (use _int8 / _fp16 / _int4 suffix for quantized variants) ---
7prompt_sess = ort.InferenceSession("onnx/model_prompt.onnx")
8decode_sess = ort.InferenceSession("onnx/model_decode.onnx")
9speaker_sess = ort.InferenceSession("onnx/speaker_proj.onnx")
10embed_table = np.load("embed_tokens.npy")
11
12# --- Load tokenizer (AutoTokenizer won't work due to custom tokenizer_class) ---
13with open("tokenizer_config.json") as f:
14 tok_cfg = json.load(f)
15tokenizer = PreTrainedTokenizerFast(
16 tokenizer_file="tokenizer.json",
17 eos_token=tok_cfg.get("eos_token", "<eos>"),
18 pad_token=tok_cfg.get("pad_token"),
19)
20TEXT_TAG = tokenizer.convert_tokens_to_ids("<text>")
21AUDIO_TAG = tokenizer.convert_tokens_to_ids("<audio>")
22AUDIO_START = tokenizer.convert_tokens_to_ids("<audio_0>")
23AUDIO_END = tokenizer.convert_tokens_to_ids("<audio_12799>")
24EOS_ID = 0 # <eos> token, MUST be 0, not 2
25
26# --- Load speaker embedding ---
27with open("speakers.json") as f:
28 speakers = json.load(f)
29speaker_emb = np.array(speakers["dv"], dtype=np.float32) # or "ida"
30
31# --- Sampling parameters (these are critical!) ---
32TEMPERATURE = 0.6
33TOP_K = 50
34TOP_P = 0.95
35REPETITION_PENALTY = 1.3
36MAX_TOKENS = 500
37
38# --- 1. Project speaker embedding to hidden size ---
39speaker_hidden = speaker_sess.run(None, {"input": speaker_emb[np.newaxis, :]})[0]
40
41# --- 2. Tokenize and build inputs_embeds ---
42text = "Hej, hvordan har du det?"
43text_ids = tokenizer.encode(text, add_special_tokens=False)
44prompt_ids = [TEXT_TAG] + text_ids + [AUDIO_TAG]
45token_embeds = embed_table[prompt_ids] # [seq_len, H]
46
47inputs_embeds = np.concatenate([
48 speaker_hidden[:, np.newaxis, :], # [1, 1, H]
49 token_embeds[np.newaxis, :, :], # [1, seq_len, H]
50], axis=1).astype(np.float32) # [1, 1+seq_len, H]
51
52total_len = inputs_embeds.shape[1]
53attention_mask = np.ones((1, total_len), dtype=np.int64)
54
55# --- 3. Prompt pass ---
56outputs = prompt_sess.run(None, {
57 "inputs_embeds": inputs_embeds,
58 "attention_mask": attention_mask,
59})
60logits = outputs[0]
61past_kvs = outputs[1:]
62
63# Read KV cache input names from decode session
64kv_names = [inp.name for inp in decode_sess.get_inputs()][3:]
65
66# --- 4. Autoregressive decode loop ---
67generated = []
68expected_audio = len(text_ids) * 10
69eos_boost_threshold = max(int(expected_audio * 1.5), 50)
70audio_count = 0
71
72for step in range(MAX_TOKENS):
73 current_logits = logits[0, -1, :].copy()
74
75 # EOS boost: after expected duration, progressively boost EOS logit
76 if audio_count > eos_boost_threshold:
77 current_logits[EOS_ID] += (audio_count - eos_boost_threshold) * 2.0
78
79 # Repetition penalty: penalize already-generated tokens
80 if REPETITION_PENALTY > 1.0 and generated:
81 for tid in set(generated):
82 if current_logits[tid] > 0:
83 current_logits[tid] /= REPETITION_PENALTY
84 else:
85 current_logits[tid] *= REPETITION_PENALTY
86
87 # Temperature + top-k + top-p sampling
88 current_logits /= TEMPERATURE
89 if TOP_K > 0:
90 top_k_idx = np.argsort(current_logits)[-TOP_K:]
91 mask = np.full_like(current_logits, -np.inf)
92 mask[top_k_idx] = current_logits[top_k_idx]
93 current_logits = mask
94 current_logits -= np.max(current_logits)
95 probs = np.exp(current_logits) / np.sum(np.exp(current_logits))
96 if TOP_P < 1.0:
97 sorted_idx = np.argsort(probs)[::-1]
98 cumsum = np.cumsum(probs[sorted_idx])
99 keep = sorted_idx[:np.searchsorted(cumsum, TOP_P) + 1]
100 filtered = np.zeros_like(probs)
101 filtered[keep] = probs[keep]
102 probs = filtered / filtered.sum()
103
104 next_token = int(np.random.choice(len(probs), p=probs))
105
106 # Stop on EOS
107 if next_token == EOS_ID:
108 break
109 generated.append(next_token)
110
111 # Track audio tokens
112 if AUDIO_START <= next_token <= AUDIO_END:
113 audio_count += 1
114
115 # Loop detection: 8 identical tokens = degenerate loop, stop
116 if len(generated) >= 8 and len(set(generated[-8:])) == 1:
117 break
118
119 # Prepare next decode step
120 new_embed = embed_table[next_token][np.newaxis, np.newaxis, :].astype(np.float32)
121 total_len += 1
122 attention_mask = np.ones((1, total_len), dtype=np.int64)
123 cache_pos = np.array([total_len - 1], dtype=np.int64)
124
125 feed = {"inputs_embeds": new_embed, "attention_mask": attention_mask, "cache_position": cache_pos}
126 for name, kv in zip(kv_names, past_kvs):
127 feed[name] = kv
128
129 outputs = decode_sess.run(None, feed)
130 logits = outputs[0]
131 past_kvs = outputs[1:]
132
133# --- 5. Extract audio tokens ---
134audio_tokens = [t - AUDIO_START for t in generated if AUDIO_START <= t <= AUDIO_END]
135print(f"Generated {len(generated)} tokens ({len(audio_tokens)} audio)")
136
137# --- 6. Decode to audio (requires PyTorch + kanade-tokenizer) ---
138import torch
139from kanade_tokenizer import KanadeModel, load_vocoder, vocode
140
141kanade = KanadeModel.from_pretrained("frothywater/kanade-25hz-clean").eval()
142vocoder = load_vocoder(kanade.config.vocoder_name)
143
144with torch.no_grad():
145 mel = kanade.decode(
146 content_token_indices=torch.tensor(audio_tokens),
147 global_embedding=torch.tensor(speaker_emb),
148 )
149 waveform = vocode(vocoder, mel.unsqueeze(0)).squeeze().numpy()
150
151import soundfile as sf
152sf.write("output.wav", waveform, 24000)