abubasith86/programming_ai
1
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