Clap / app.py
Radiotomy's picture
Add CLAP embedding service
5cecf0e verified
Raw History Blame Contribute Delete
3.73 kB
"""Clap — CLAP embedding service for BASE Station semantic audio search.
Text and audio land in the SAME 512-d space, which is the whole point: a typed
query can be compared to a stored audio embedding by plain cosine similarity,
with no separate ranking model.
Embeddings are L2-NORMALIZED here, at the source. Normalizing once on this side
means every consumer can use a plain dot product and can never accidentally
compare a normalized vector against an unnormalized one.
"""
import io
import logging
import librosa
import numpy as np
import requests
import torch
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from transformers import ClapModel, ClapProcessor
MODEL_ID = "laion/larger_clap_music_and_speech"
SAMPLE_RATE = 48000 # CLAP's expected input rate — resampling is not optional
MAX_SECONDS = 30 # cap the decode; a long file adds cost without adding meaning
DOWNLOAD_TIMEOUT = 60
MAX_BYTES = 60 * 1024 * 1024
logging.basicConfig(level=logging.INFO)
log = logging.getLogger("clap")
app = FastAPI(title="Clap")
_model = None
_processor = None
def get_model():
"""Load lazily on first request so the container boots fast and a cold
Space still answers /health immediately."""
global _model, _processor
if _model is None:
log.info("loading %s", MODEL_ID)
_processor = ClapProcessor.from_pretrained(MODEL_ID)
_model = ClapModel.from_pretrained(MODEL_ID).eval()
log.info("model ready")
return _model, _processor
def normalize(vec):
v = np.asarray(vec, dtype=np.float32).reshape(-1)
n = float(np.linalg.norm(v))
if n == 0.0:
raise HTTPException(500, "model returned a zero vector")
return (v / n).tolist()
class TextIn(BaseModel):
text: str
class AudioIn(BaseModel):
audio_url: str
@app.get("/health")
def health():
return {"ok": True, "model": MODEL_ID, "dim": 512}
@app.post("/embed/text")
def embed_text(body: TextIn):
text = (body.text or "").strip()
if not text:
raise HTTPException(400, "text is required")
model, processor = get_model()
inputs = processor(text=[text], return_tensors="pt", padding=True)
with torch.no_grad():
feats = model.get_text_features(**inputs)
return {"embedding": normalize(feats[0].numpy()), "dim": 512}
@app.post("/embed/audio")
def embed_audio(body: AudioIn):
url = (body.audio_url or "").strip()
if not url.startswith(("http://", "https://")):
raise HTTPException(400, "audio_url must be an http(s) URL")
try:
r = requests.get(url, timeout=DOWNLOAD_TIMEOUT, stream=True)
r.raise_for_status()
raw = r.raw.read(MAX_BYTES + 1, decode_content=True)
except Exception as e:
raise HTTPException(502, f"could not fetch audio: {e}")
if len(raw) > MAX_BYTES:
raise HTTPException(413, "audio file too large")
if not raw:
raise HTTPException(502, "audio url returned no data")
try:
# mono at CLAP's rate; librosa handles wav/flac/mp3/ogg via soundfile+ffmpeg
wav, _ = librosa.load(io.BytesIO(raw), sr=SAMPLE_RATE, mono=True,
duration=MAX_SECONDS)
except Exception as e:
raise HTTPException(422, f"could not decode audio: {e}")
if wav.size == 0:
raise HTTPException(422, "decoded audio was empty")
model, processor = get_model()
inputs = processor(audios=[wav], sampling_rate=SAMPLE_RATE, return_tensors="pt")
with torch.no_grad():
feats = model.get_audio_features(**inputs)
return {
"embedding": normalize(feats[0].numpy()),
"dim": 512,
"seconds": round(float(wav.size) / SAMPLE_RATE, 2),
}