originalTimi/Hypa-Orpheus-Step-latest-16bit
022
1"""2HF Inference Endpoint handler โ Hypa Orpheus TTS + Voice Cloning (merged 16-bit).3Task matrix (routed by `parameters`):4 task="tts", mode="vanilla" : {speaker}: text -> speech5 task="tts", mode="translate" : {speaker} - {Language}: text -> speech in Language6 task="vc", mode="vanilla" : reference (text+audio) + target text -> speech in reference voice7 task="vc", mode="translate" : + language tag on target text -> cross-lingual cloning8 VC method="m1" (in-context) | method="m2" (continue-speaking)9Output parity with the legacy endpoint: `audio_b64` is base64 of the RAW10float32 little-endian mono PCM buffer at 24000 Hz (NO WAV/RIFF container),11so existing products decode with: np.frombuffer(base64.b64decode(s), dtype=np.int16)12[If the legacy endpoint used int16, change RAW_DTYPE to np.int16 below.]13Prompts are byte-identical to Step-III training (_encode_text / build_tts /14build_vc_both), reference codes are frame-deduped, and prompts reach vLLM as15token ids (never a decoded string).16"""17 18import io19import os20import base6421import tempfile22import traceback23 24import numpy as np25import torch26import soundfile as sf27import librosa28 29from transformers import AutoTokenizer30from snac import SNAC31from vllm import LLM, SamplingParams32 33 34class EndpointHandler:35 # ---- Orpheus special tokens (fixed by the model) ----36 TOKENISER_LEN = 12825637 START_OF_TEXT = 12800038 END_OF_TEXT = 12800939 START_OF_SPEECH = TOKENISER_LEN + 1 # 12825740 END_OF_SPEECH = TOKENISER_LEN + 2 # 12825841 START_OF_HUMAN = TOKENISER_LEN + 3 # 12825942 END_OF_HUMAN = TOKENISER_LEN + 4 # 12826043 START_OF_AI = TOKENISER_LEN + 5 # 12826144 END_OF_AI = TOKENISER_LEN + 6 # 12826245 AUDIO_OFFSET = 12826646 47 # NOTE: fine-tune data capped at 2048 tokens; 4096 kept so M1-VC prompts48 # (ref codes + two texts, often 1000-2000 tokens) retain a generation49 # budget. Base Llama-3 RoPE supports these positions natively; expect the50 # best quality when prompt+generation stays near the trained ~2048.51 MAX_MODEL_LEN = 409652 MAX_REF_SECONDS = 3053 SNAC_SR = 2400054 RAW_DTYPE = np.int16 # legacy raw-PCM dtype (see docstring)55 56 LANG_DISPLAY = {57 "en": "English", "es": "Spanish", "fr": "French", "ha": "Hausa",58 "yo": "Yoruba", "sw": "Swahili", "ar": "Arabic", "pt": "Portuguese",59 "ann": "Annang", "ebi": "Ebira", "efi": "Efik", "ego": "Eggon",60 "urh": "Urhobo", "ibb": "Ibibio", "idm": "Idoma", "igl": "Igala",61 "ig": "Igbo", "nup": "Nupe", "tiv": "Tiv", "pg": "Pidgin",62 }63 64 # ------------------------------------------------------------------ init65 def __init__(self, path=""):66 self.device = "cuda" if torch.cuda.is_available() else "cpu"67 # SNAC first (tiny, ~80 MB) so it never contends with vLLM's reservation.68 self.snac = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").to(self.device).eval()69 self.model = LLM(70 path,71 max_model_len=self.MAX_MODEL_LEN,72 gpu_memory_utilization=0.75,73 model_impl="transformers",74 )75 self.tokenizer = AutoTokenizer.from_pretrained(path)76 77 # ------------------------------------------------------- text encoding78 def _lang_display(self, x):79 if x is None:80 return None81 k = str(x).strip().lower()82 return self.LANG_DISPLAY.get(k, k.capitalize() if k else None)83 84 def _encode_text(self, text, speaker=None, lang_tag=None, add_bos=True):85 text = "" if text is None else str(text).strip()86 spk = speaker if (speaker and str(speaker).strip().lower() not in ("", "random", "none")) else None87 if spk and lang_tag:88 prompt = f"{spk} - {lang_tag}: {text}"89 elif spk:90 prompt = f"{spk}: {text}"91 elif lang_tag:92 prompt = f"{lang_tag}: {text}"93 else:94 prompt = text95 ids = self.tokenizer.encode(prompt, add_special_tokens=add_bos)96 ids.append(self.END_OF_TEXT)97 return ids98 99 # ------------------------------------------------------ audio encoding100 def _b64_to_wave(self, b64_str):101 raw = base64.b64decode(b64_str)102 if not raw:103 raise ValueError("reference_audio is empty.")104 try:105 arr, sr = sf.read(io.BytesIO(raw), dtype="float32")106 except Exception:107 # temp-file fallback: librosa/audioread handles containers108 # libsndfile can't open, but needs a real file path for some codecs.109 tmp = None110 try:111 with tempfile.NamedTemporaryFile(delete=False, suffix=".audio") as f:112 f.write(raw)113 tmp = f.name114 arr, sr = librosa.load(tmp, sr=None, mono=False)115 arr = np.asarray(arr, dtype=np.float32)116 if arr.ndim > 1:117 arr = arr.T # librosa returns (ch, n)118 finally:119 if tmp and os.path.exists(tmp):120 os.remove(tmp)121 if arr.ndim > 1:122 arr = arr.mean(axis=1)123 if arr.size == 0 or not np.isfinite(arr).all():124 raise ValueError("Reference audio is empty or contains invalid samples.")125 if sr != self.SNAC_SR:126 arr = librosa.resample(arr.astype(np.float32), orig_sr=sr, target_sr=self.SNAC_SR)127 dur = len(arr) / self.SNAC_SR128 if dur > self.MAX_REF_SECONDS:129 raise ValueError(f"Reference audio is {dur:.1f}s; max is {self.MAX_REF_SECONDS}s. "130 f"Send a shorter clip.")131 return arr.astype(np.float32)132 133 @torch.inference_mode()134 def _audio_to_codes(self, arr):135 wav = torch.from_numpy(arr).to(self.device)[None, None]136 codes = self.snac.encode(wav)137 c0, c1, c2 = codes[0][0].tolist(), codes[1][0].tolist(), codes[2][0].tolist()138 n = min(len(c0), len(c1) // 2, len(c2) // 4)139 out = []140 for i in range(n):141 out += [142 c0[i] + self.AUDIO_OFFSET,143 c1[2 * i] + self.AUDIO_OFFSET + 4096,144 c2[4 * i] + self.AUDIO_OFFSET + 2 * 4096,145 c2[4 * i + 1] + self.AUDIO_OFFSET + 3 * 4096,146 c1[2 * i + 1] + self.AUDIO_OFFSET + 4 * 4096,147 c2[4 * i + 2] + self.AUDIO_OFFSET + 5 * 4096,148 c2[4 * i + 3] + self.AUDIO_OFFSET + 6 * 4096,149 ]150 return out151 152 @staticmethod153 def _dedup_frames(codes):154 if not codes:155 return codes156 codes = codes[: (len(codes) // 7) * 7]157 if len(codes) < 7:158 return codes159 result = codes[:7]160 for i in range(7, len(codes), 7):161 if codes[i] != result[-7]:162 result.extend(codes[i:i + 7])163 return result164 165 # ------------------------------------------------------ prompt builders166 def build_tts_prompt(self, text, speaker, mode, language):167 lang_tag = self._lang_display(language) if mode == "translate" else None168 tt = self._encode_text(text, speaker, lang_tag, add_bos=True)169 return [self.START_OF_HUMAN] + tt + [self.END_OF_HUMAN,170 self.START_OF_AI, self.START_OF_SPEECH]171 172 def build_vc_prompt(self, ref_text, ref_codes, target_text, mode, language, method):173 tag2 = self._lang_display(language) if mode == "translate" else None174 tt1 = self._encode_text(ref_text, None, None, add_bos=True)175 tt2 = self._encode_text(target_text, None, tag2, add_bos=False)176 if method == "m1":177 return ([self.START_OF_HUMAN] + tt1 + [self.END_OF_HUMAN,178 self.START_OF_AI, self.START_OF_SPEECH] + ref_codes +179 [self.END_OF_SPEECH, self.END_OF_AI,180 self.START_OF_HUMAN] + tt2 + [self.END_OF_HUMAN,181 self.START_OF_AI, self.START_OF_SPEECH])182 return ([self.START_OF_HUMAN] + tt1 + tt2 + [self.END_OF_HUMAN,183 self.START_OF_AI, self.START_OF_SPEECH] + ref_codes)184 185 # --------------------------------------------------------- generation186 def _generate(self, prompt_ids, params):187 sampling = SamplingParams(188 temperature = params["temperature"],189 top_p = params["top_p"],190 top_k = params["top_k"],191 max_tokens = params["max_new_tokens"],192 repetition_penalty = params["repetition_penalty"],193 stop_token_ids = [self.END_OF_SPEECH, self.END_OF_AI],194 detokenize = False,195 )196 outputs = self.model.generate({"prompt_token_ids": prompt_ids}, sampling)197 return list(outputs[0].outputs[0].token_ids)198 199 # ----------------------------------------------------------- decoding200 @torch.inference_mode()201 def _codes_to_wave(self, gen_ids):202 """Frame-validating SNAC decode: accepts only well-formed 7-token frames203 (token k in slot-k range), resyncs on malformed spans."""204 frames, i, n, resyncs = [], 0, len(gen_ids), 0205 while i <= n - 7:206 vals, ok = [], True207 for k in range(7):208 lo = self.AUDIO_OFFSET + k * 4096209 t = gen_ids[i + k]210 if not (lo <= t < lo + 4096):211 ok = False212 break213 vals.append(t - lo)214 if ok:215 frames.append(vals)216 i += 7217 else:218 i += 1219 resyncs += 1220 self._last_resyncs = resyncs221 if not frames:222 return None, 0223 l1 = [f[0] for f in frames]224 l2, l3 = [], []225 for f in frames:226 l2.append(f[1]); l3.append(f[2]); l3.append(f[3])227 l2.append(f[4]); l3.append(f[5]); l3.append(f[6])228 tensors = [229 torch.tensor(l1)[None].to(self.device),230 torch.tensor(l2)[None].to(self.device),231 torch.tensor(l3)[None].to(self.device),232 ]233 wav = self.snac.decode(tensors).squeeze().detach().cpu().numpy()234 return wav, len(frames)235 236 def _wave_to_b64_raw(self, wav):237 wav = np.clip(wav, -1.0, 1.0)238 pcm16 = (wav * 32767.0).astype(self.RAW_DTYPE)239 return base64.b64encode(np.ascontiguousarray(pcm16).tobytes()).decode("utf-8")240 241 # -------------------------------------------------------------- entry242 def __call__(self, data):243 try:244 target_text = data.get("inputs")245 if not target_text:246 return {"error": "Missing 'inputs' (target text)."}247 248 p = data.get("parameters", {}) or {}249 task = str(p.get("task", "tts")).lower()250 mode = str(p.get("mode", "vanilla")).lower()251 method = str(p.get("method", "m2")).lower()252 if mode in ("translation", "trans"):253 mode = "translate"254 255 if task not in ("tts", "vc"):256 return {"error": "parameters.task must be 'tts' or 'vc'."}257 if mode not in ("vanilla", "translate"):258 return {"error": "parameters.mode must be 'vanilla' or 'translate'."}259 if mode == "translate" and not p.get("language"):260 return {"error": "parameters.language is required for translate mode."}261 262 gen_params = {263 "temperature": float(p.get("temperature", 0.6)),264 "top_p": float(p.get("top_p", 0.95)),265 "top_k": int(p.get("top_k", 50)),266 "max_new_tokens": int(p.get("max_new_tokens", 1200)),267 "repetition_penalty": float(p.get("repetition_penalty", 1.1)),268 }269 if not 0 < gen_params["top_p"] <= 1:270 return {"error": "top_p must be within (0, 1]."}271 if not (gen_params["top_k"] == -1 or gen_params["top_k"] > 0):272 return {"error": "top_k must be -1 (disabled) or a positive integer."}273 if not 0 < gen_params["repetition_penalty"] <= 2:274 return {"error": "repetition_penalty must be within (0, 2]."}275 if gen_params["max_new_tokens"] <= 0:276 return {"error": "max_new_tokens must be positive."}277 278 if task == "vc":279 ref_text = p.get("reference_text")280 ref_audio = p.get("reference_audio")281 if not ref_text or not ref_audio:282 return {"error": "VC requires parameters.reference_text and "283 "parameters.reference_audio (base64)."}284 if method not in ("m1", "m2"):285 return {"error": "parameters.method must be 'm1' or 'm2'."}286 ref_wave = self._b64_to_wave(ref_audio)287 ref_codes = self._dedup_frames(self._audio_to_codes(ref_wave))288 if not ref_codes:289 return {"error": "Reference audio produced no SNAC codes."}290 prompt_ids = self.build_vc_prompt(291 ref_text, ref_codes, target_text, mode, p.get("language"), method)292 else:293 prompt_ids = self.build_tts_prompt(294 target_text, p.get("voice") or p.get("speaker"),295 mode, p.get("language"))296 297 budget = self.MAX_MODEL_LEN - gen_params["max_new_tokens"]298 if len(prompt_ids) > budget:299 return {"error": f"Prompt is {len(prompt_ids)} tokens; exceeds budget "300 f"{budget} (max_model_len - max_new_tokens). "301 f"Shorten the reference clip or text."}302 303 gen_ids = self._generate(prompt_ids, gen_params)304 wav, n_frames = self._codes_to_wave(gen_ids)305 if wav is None:306 return {"error": "Model generated no audio tokens.",307 "input_tokens": len(prompt_ids),308 "generated_tokens": len(gen_ids)}309 310 return {311 "audio_b64": self._wave_to_b64_raw(wav), # RAW float32 PCM (legacy parity)312 "audio_dtype": np.dtype(self.RAW_DTYPE).name,313 "sample_rate": self.SNAC_SR,314 "duration_seconds": round(len(wav) / self.SNAC_SR, 3),315 "audio_frames": n_frames,316 "input_tokens": len(prompt_ids),317 "generated_tokens": len(gen_ids),318 "task": task, "mode": mode,319 "method": method if task == "vc" else None,320 "decode_resyncs": getattr(self, "_last_resyncs", 0),321 }322 323 except ValueError as e:324 return {"error": str(e)}325 except Exception as e:326 traceback.print_exc()327 return {"error": str(e)}