Team Ai
Apppublic

abubasith86/programming_ai

sourceHugging Faceupdated 2mo agoView on Hugging Face
1likes
app.py204 linesDownload Raw Back to root
1import os2 3# Must be set before torch's OpenMP/MKL thread pools initialize on first use.4# This is the most reliable way to make sure BLAS/OpenMP actually uses all5# cores instead of a conservative default.6def _detect_cpu_threads() -> int:7    # Respect a value the platform/container already set -- on cgroup-limited8    # containers (like HF Spaces) this is often the *correct* real quota,9    # whereas os.cpu_count() reports the host machine's full core count and10    # will cause thread oversubscription if trusted blindly.11    for var in ("OMP_NUM_THREADS", "MKL_NUM_THREADS"):12        val = os.environ.get(var)13        if val and val.isdigit() and int(val) > 0:14            return int(val)15    # sched_getaffinity reflects CPU-affinity/cgroup restrictions more16    # accurately than os.cpu_count() on Linux; fall back to cpu_count if17    # unavailable (e.g. non-Linux).18    try:19        return len(os.sched_getaffinity(0))20    except AttributeError:21        return os.cpu_count() or 422 23 24_CPU_THREADS = str(_detect_cpu_threads())25os.environ.setdefault("OMP_NUM_THREADS", _CPU_THREADS)26os.environ.setdefault("MKL_NUM_THREADS", _CPU_THREADS)27 28import tempfile29import time30import traceback31from typing import Optional32 33import soundfile as sf34import torch35from fastapi import FastAPI, HTTPException, Query36from fastapi.responses import FileResponse37from starlette.background import BackgroundTask38 39from qwen_tts import Qwen3TTSModel40 41app = FastAPI()42 43MODEL = None44SUPPORTED_SPEAKERS = []45SUPPORTED_LANGUAGES = []46 47MODEL_NAME = "Qwen/Qwen3-TTS-12Hz-1.7B-CustomVoice"48DEFAULT_SPEAKER = "Ryan"       # English male voice; pick any from SUPPORTED_SPEAKERS49DEFAULT_LANGUAGE = "English"   # or "Auto" to let the model detect it50 51# bf16 has no native hardware acceleration on most x86 CPUs (it gets52# emulated via upcasting), which is often slower than plain fp32. fp32 is53# the reliable baseline for generic CPU inference.54CPU_DTYPE = torch.float3255 56 57@app.on_event("startup")58def startup_event():59    global MODEL, SUPPORTED_SPEAKERS, SUPPORTED_LANGUAGES60    try:61        print("=" * 80)62        print("Starting application...")63        print(f"PyTorch version: {torch.__version__}")64        cuda_available = torch.cuda.is_available()65        print(f"CUDA available: {cuda_available}")66 67        if not cuda_available:68            n_threads = int(_CPU_THREADS)69            # Intra-op parallelism (parallelizes ops like matmul/conv internally).70            torch.set_num_threads(n_threads)71            # Inter-op parallelism (runs independent ops concurrently). Can72            # only be set once, before any parallel work starts, so this73            # must stay early and wrapped defensively.74            try:75                torch.set_num_interop_threads(max(1, n_threads // 2))76            except RuntimeError as e:77                print(f"Could not set interop threads (already initialized): {e}")78            # Denormal floats are handled by a slow FP path on most CPUs;79            # flushing them to zero avoids random latency spikes during80            # generation. Negligible effect on audio quality.81            torch.set_flush_denormal(True)82            print(f"No GPU detected. CPU threads: {n_threads} "83                  f"(OMP_NUM_THREADS={os.environ.get('OMP_NUM_THREADS')})")84 85        print(f"Loading Qwen3-TTS model: {MODEL_NAME}")86        load_start = time.perf_counter()87 88        device_map = "cuda:0" if cuda_available else "cpu"89        dtype = torch.bfloat16 if cuda_available else CPU_DTYPE90        load_kwargs = dict(device_map=device_map, dtype=dtype)91 92        if cuda_available:93            try:94                MODEL = Qwen3TTSModel.from_pretrained(95                    MODEL_NAME,96                    attn_implementation="flash_attention_2",97                    **load_kwargs,98                )99            except Exception as flash_err:100                print(f"flash_attention_2 unavailable ({flash_err}); "101                      f"falling back to default attention implementation.")102                MODEL = Qwen3TTSModel.from_pretrained(MODEL_NAME, **load_kwargs)103        else:104            MODEL = Qwen3TTSModel.from_pretrained(MODEL_NAME, **load_kwargs)105 106        SUPPORTED_SPEAKERS = MODEL.get_supported_speakers()107        SUPPORTED_LANGUAGES = MODEL.get_supported_languages()108 109        load_elapsed = time.perf_counter() - load_start110        print(f"Supported speakers: {SUPPORTED_SPEAKERS}")111        print(f"Supported languages: {SUPPORTED_LANGUAGES}")112        print(f"Model loaded successfully in {load_elapsed:.1f}s.")113        print("=" * 80)114    except Exception as exc:115        print(traceback.format_exc())116        MODEL = None117        print(f"Startup failed: {exc}")118 119 120@app.get("/")121def health():122    return {123        "status": "running",124        "model_loaded": MODEL is not None,125        "cpu_threads": _CPU_THREADS,126        "supported_speakers": SUPPORTED_SPEAKERS,127        "supported_languages": SUPPORTED_LANGUAGES,128    }129 130 131@app.get("/tts")132def generate(133    text: str,134    speaker: str = Query(DEFAULT_SPEAKER, description="Voice to use, e.g. Ryan, Vivian, Aiden"),135    language: str = Query(DEFAULT_LANGUAGE, description="Target language, or 'Auto' to detect"),136    instruct: Optional[str] = Query(137        None, description="Optional natural-language style instruction, e.g. 'Speak happily'"138    ),139    max_new_tokens: Optional[int] = Query(140        None,141        description="Optional cap on generated audio tokens. Lower = faster on CPU, "142                    "but can truncate longer sentences. Omit to use the model default.",143    ),144):145    output_file = None146    try:147        print("=" * 80)148        print(f"Incoming text: {text}")149        print(f"speaker={speaker} language={language} instruct={instruct!r} "150              f"max_new_tokens={max_new_tokens}")151 152        if MODEL is None:153            raise HTTPException(status_code=500, detail="Model was not loaded.")154 155        if SUPPORTED_SPEAKERS and speaker not in SUPPORTED_SPEAKERS:156            raise HTTPException(157                status_code=400,158                detail=f"Unsupported speaker '{speaker}'. Choose from: {SUPPORTED_SPEAKERS}",159            )160 161        generate_kwargs = {}162        if max_new_tokens is not None:163            generate_kwargs["max_new_tokens"] = max_new_tokens164 165        gen_start = time.perf_counter()166        # inference_mode disables autograd bookkeeping entirely (faster and167        # lower memory than no_grad) -- pure inference, no training ever168        # happens on this path.169        with torch.inference_mode():170            wavs, sample_rate = MODEL.generate_custom_voice(171                text=text,172                language=language,173                speaker=speaker,174                instruct=instruct or "",175                **generate_kwargs,176            )177        gen_elapsed = time.perf_counter() - gen_start178        print(f"Generation took {gen_elapsed:.2f}s")179 180        output_file = tempfile.NamedTemporaryFile(suffix=".wav", delete=False).name181        sf.write(output_file, wavs[0], sample_rate)182 183        size = os.path.getsize(output_file)184        print(f"Output file: {output_file}")185        print(f"File size: {size} bytes")186        if size == 0:187            raise RuntimeError("Generated WAV file is empty.")188        print("=" * 80)189 190        return FileResponse(191            output_file,192            media_type="audio/wav",193            filename="speech.wav",194            headers={"X-Generation-Seconds": f"{gen_elapsed:.2f}"},195            # Clean up the temp file once the response has been sent.196            background=BackgroundTask(lambda: os.remove(output_file) if os.path.exists(output_file) else None),197        )198    except HTTPException:199        raise200    except Exception:201        print(traceback.format_exc())202        if output_file and os.path.exists(output_file):203            os.remove(output_file)204        raise