sao: serialize inference under an asyncio.Lock
StableAudioPipeline isn't reentrant — concurrent requests share the scheduler's step_index counter and corrupt each other mid-run (observed: IndexError in cosine_dpmsolver_multistep when two requests overlap). Wrap the pipeline call + audio decode in a single asyncio.Lock created at startup, and run the (sync, GPU-bound) pipeline call via asyncio.to_thread so the event loop stays responsive. Concurrent requests now queue cleanly instead of racing. Verified: 5 parallel POSTs at steps=50 all return 200, clear ~4s serialization spacing (4, 8, 12, 16, 20s wall time), distinct output hashes per seed.
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
# FastAPI shim around diffusers' StableAudioPipeline.
|
# FastAPI shim around diffusers' StableAudioPipeline.
|
||||||
# Single endpoint POST /v1/audio/sfx returns a WAV blob.
|
# Single endpoint POST /v1/audio/sfx returns a WAV blob.
|
||||||
# Model is loaded once on startup and held in process memory.
|
# Model is loaded once on startup and held in process memory.
|
||||||
|
import asyncio
|
||||||
import io
|
import io
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
@@ -27,6 +28,11 @@ async def lifespan(app: FastAPI):
|
|||||||
pipe = StableAudioPipeline.from_pretrained(MODEL_ID, torch_dtype=DTYPE)
|
pipe = StableAudioPipeline.from_pretrained(MODEL_ID, torch_dtype=DTYPE)
|
||||||
pipe = pipe.to(DEVICE)
|
pipe = pipe.to(DEVICE)
|
||||||
state["pipe"] = pipe
|
state["pipe"] = pipe
|
||||||
|
# Single-GPU diffusers pipelines aren't reentrant — concurrent calls
|
||||||
|
# share the same scheduler step_index counter and corrupt each
|
||||||
|
# other (observed: IndexError in cosine_dpmsolver_multistep when
|
||||||
|
# two requests overlap mid-run). Serialize at the request boundary.
|
||||||
|
state["lock"] = asyncio.Lock()
|
||||||
print(f"[sao] loaded in {time.time() - t0:.1f}s", flush=True)
|
print(f"[sao] loaded in {time.time() - t0:.1f}s", flush=True)
|
||||||
yield
|
yield
|
||||||
state.clear()
|
state.clear()
|
||||||
@@ -55,26 +61,33 @@ def health():
|
|||||||
|
|
||||||
|
|
||||||
@app.post("/v1/audio/sfx")
|
@app.post("/v1/audio/sfx")
|
||||||
def sfx(req: SfxRequest):
|
async def sfx(req: SfxRequest):
|
||||||
pipe = state.get("pipe")
|
pipe = state.get("pipe")
|
||||||
if pipe is None:
|
lock = state.get("lock")
|
||||||
|
if pipe is None or lock is None:
|
||||||
raise HTTPException(503, "model not loaded yet")
|
raise HTTPException(503, "model not loaded yet")
|
||||||
|
|
||||||
generator = None
|
generator = None
|
||||||
if req.seed is not None:
|
if req.seed is not None:
|
||||||
generator = torch.Generator(DEVICE).manual_seed(req.seed)
|
generator = torch.Generator(DEVICE).manual_seed(req.seed)
|
||||||
|
|
||||||
audio = pipe(
|
# Hold the lock for the whole inference + encode. Concurrent
|
||||||
req.prompt,
|
# requests queue cleanly instead of racing the scheduler's
|
||||||
negative_prompt=req.negative_prompt,
|
# step_index counter into an IndexError.
|
||||||
num_inference_steps=req.steps,
|
async with lock:
|
||||||
audio_end_in_s=req.duration,
|
audio = await asyncio.to_thread(
|
||||||
num_waveforms_per_prompt=1,
|
lambda: pipe(
|
||||||
guidance_scale=req.cfg_scale,
|
req.prompt,
|
||||||
generator=generator,
|
negative_prompt=req.negative_prompt,
|
||||||
).audios
|
num_inference_steps=req.steps,
|
||||||
|
audio_end_in_s=req.duration,
|
||||||
|
num_waveforms_per_prompt=1,
|
||||||
|
guidance_scale=req.cfg_scale,
|
||||||
|
generator=generator,
|
||||||
|
).audios
|
||||||
|
)
|
||||||
|
waveform = audio[0].T.float().cpu().numpy()
|
||||||
|
|
||||||
waveform = audio[0].T.float().cpu().numpy()
|
|
||||||
buf = io.BytesIO()
|
buf = io.BytesIO()
|
||||||
sf.write(buf, waveform, pipe.vae.sampling_rate, format="WAV")
|
sf.write(buf, waveform, pipe.vae.sampling_rate, format="WAV")
|
||||||
buf.seek(0)
|
buf.seek(0)
|
||||||
|
|||||||
Reference in New Issue
Block a user