crowncode-backend / app /services /wav2vec2_detector.py
Rthur2003's picture
feat: aI analysis servisleri routes detectors ve test suites eklendi
06b738c
Raw History Blame Contribute Delete
5.39 kB
"""
wav2vec2 inference service for AURIS (Tower 1).
Loads the fine-tuned wav2vec2 model and provides
real-time predictions. Falls back gracefully if
model file is not available.
"""
from __future__ import annotations
import io
from dataclasses import dataclass
from pathlib import Path
from typing import Optional, Union
import numpy as np
from .logging_config import get_logger
logger = get_logger(__name__)
_MODEL_PATH = (
Path(__file__).resolve().parents[2] / "models" / "wav2vec2_auris_v1.pt"
)
_SAMPLE_RATE = 16000
_MAX_SEC = 30.0
@dataclass
class Wav2Vec2Result:
"""Result from wav2vec2 tower."""
available: bool
p_ai: float = 0.5
# Hidden state statistics for meta-classifier
hidden_mean: float = 0.0
hidden_std: float = 0.0
hidden_max: float = 0.0
hidden_min: float = 0.0
hidden_kurtosis: float = 0.0
error: Optional[str] = None
class Wav2Vec2DetectorService:
"""
Tower 1: wav2vec2-based AI music detection.
Loads the fine-tuned model from disk or HuggingFace Hub.
On CPU, inference takes ~0.5s for 30s audio.
"""
def __init__(self, model_path: Optional[Path] = None) -> None:
"""Initialize wav2vec2 detector with optional model path."""
self._model = None
self._device = None
self._initialized = False
self._model_path = model_path or _MODEL_PATH
def status(self) -> dict:
"""Weights on disk / model in memory (it loads on first use)."""
return {
"weights": self._model_path.exists(),
"loaded": self._model is not None,
"initialized": self._initialized,
}
def _ensure_loaded(self) -> bool:
"""Lazy-load model on first use."""
if self._initialized:
return self._model is not None
self._initialized = True
if not self._model_path.exists():
logger.warning(
f"wav2vec2 model not found: {self._model_path}. "
"Run training pipeline first."
)
return False
try:
import torch
from app.training.wav2vec2_classifier import (
Wav2Vec2MusicClassifier,
Wav2Vec2Config,
)
self._device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
config = Wav2Vec2Config()
self._model = Wav2Vec2MusicClassifier(config)
state = torch.load(
self._model_path,
map_location=self._device,
weights_only=True,
)
self._model.load_state_dict(state)
self._model.to(self._device)
self._model.eval()
logger.info(
f"wav2vec2 model loaded from {self._model_path} "
f"on {self._device}"
)
return True
except Exception as e:
logger.error(f"Failed to load wav2vec2 model: {e}")
self._model = None
return False
def predict(
self, source: Union[Path, bytes, io.BytesIO]
) -> Wav2Vec2Result:
"""
Run wav2vec2 inference on audio.
Args:
source: Audio file path, raw bytes, or BytesIO.
Returns:
Wav2Vec2Result with prediction and hidden stats.
"""
if not self._ensure_loaded():
return Wav2Vec2Result(
available=False,
error="model_not_loaded",
)
try:
import torch
import librosa
# Load audio at 16kHz
if isinstance(source, (bytes, io.BytesIO)):
if isinstance(source, bytes):
source = io.BytesIO(source)
y, _ = librosa.load(source, sr=_SAMPLE_RATE, mono=True)
else:
y, _ = librosa.load(
str(source), sr=_SAMPLE_RATE, mono=True
)
# Truncate or pad
max_samples = int(_MAX_SEC * _SAMPLE_RATE)
if len(y) > max_samples:
y = y[:max_samples]
elif len(y) < _SAMPLE_RATE:
return Wav2Vec2Result(
available=False,
error="audio_too_short",
)
# Inference
input_tensor = torch.tensor(
y, dtype=torch.float32
).unsqueeze(0).to(self._device)
with torch.no_grad():
logits, hidden = self._model(input_tensor)
p_ai = float(
torch.sigmoid(logits).cpu().item()
)
# Hidden state statistics for meta-classifier
h = hidden.cpu().numpy().flatten()
from scipy.stats import kurtosis
return Wav2Vec2Result(
available=True,
p_ai=round(p_ai, 4),
hidden_mean=float(np.mean(h)),
hidden_std=float(np.std(h)),
hidden_max=float(np.max(h)),
hidden_min=float(np.min(h)),
hidden_kurtosis=float(kurtosis(h)),
)
except Exception as e:
logger.warning(f"wav2vec2 prediction failed: {e}")
return Wav2Vec2Result(
available=False,
error=str(e),
)