CoolFace
Modelpublic

originalTimi/Hypa-Orpheus-Step-latest-16bit

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes22downloads
handler.py327 linesDownload Raw Back to root
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)}