Spaces:
Running
Running
Real Wav2Vec2 SER model inference from model.safetensors
Browse files- app.py +60 -42
- requirements.txt +3 -0
app.py
CHANGED
|
@@ -7,8 +7,13 @@ import os
|
|
| 7 |
from typing import Optional
|
| 8 |
|
| 9 |
import numpy as np
|
|
|
|
|
|
|
|
|
|
| 10 |
from fastapi import FastAPI, File, Form, UploadFile, HTTPException
|
| 11 |
from fastapi.middleware.cors import CORSMiddleware
|
|
|
|
|
|
|
| 12 |
|
| 13 |
app = FastAPI(title="Speech Emotion Recognition")
|
| 14 |
|
|
@@ -20,23 +25,47 @@ app.add_middleware(
|
|
| 20 |
allow_headers=["*"],
|
| 21 |
)
|
| 22 |
|
| 23 |
-
# ββ Model output classes (7-class SER model) ββββββββββββββββββββββββββ
|
| 24 |
MODEL_CLASSES = ["angry", "disgust", "fear", "happy", "neutral", "pleasant_surprise", "sad"]
|
| 25 |
MODEL_CLASS_INDEX = {c: i for i, c in enumerate(MODEL_CLASSES)}
|
| 26 |
|
| 27 |
-
# ββ Model βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 28 |
_model = None
|
| 29 |
|
| 30 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 31 |
def load_model():
|
| 32 |
-
"""Load the speech emotion model. Called once at startup."""
|
| 33 |
global _model
|
| 34 |
-
|
| 35 |
-
|
| 36 |
-
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
|
| 41 |
|
| 42 |
@app.on_event("startup")
|
|
@@ -45,48 +74,42 @@ async def startup():
|
|
| 45 |
|
| 46 |
|
| 47 |
def _extract_features(waveform: np.ndarray, sr: int) -> np.ndarray:
|
| 48 |
-
"""Extract MFCC features for model input."""
|
| 49 |
import librosa
|
| 50 |
mfcc = librosa.feature.mfcc(y=waveform, sr=sr, n_mfcc=40, n_fft=512, hop_length=256)
|
| 51 |
mfcc_delta = librosa.feature.delta(mfcc)
|
| 52 |
mfcc_delta2 = librosa.feature.delta(mfcc, order=2)
|
| 53 |
features = np.concatenate([mfcc, mfcc_delta, mfcc_delta2], axis=0)
|
| 54 |
-
# Mean-pool over time β fixed-size vector
|
| 55 |
return features.mean(axis=1)
|
| 56 |
|
| 57 |
|
| 58 |
def _predict(waveform: np.ndarray, sr: int) -> dict:
|
| 59 |
-
|
| 60 |
-
Run the loaded model on waveform.
|
| 61 |
-
Replace this function with actual inference when _model is a real model.
|
| 62 |
-
"""
|
| 63 |
if _model is None or _model == "dummy":
|
| 64 |
return _dummy_predict(waveform, sr)
|
| 65 |
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
|
| 85 |
-
|
| 86 |
|
| 87 |
|
| 88 |
def _dummy_predict(waveform: np.ndarray, sr: int) -> dict:
|
| 89 |
-
"""Dummy predictor β returns neutral with random noise."""
|
| 90 |
rng = np.random.default_rng(42)
|
| 91 |
probs = rng.dirichlet(np.ones(len(MODEL_CLASSES)) * 0.5)
|
| 92 |
probs = probs / probs.sum()
|
|
@@ -99,7 +122,6 @@ def _dummy_predict(waveform: np.ndarray, sr: int) -> dict:
|
|
| 99 |
|
| 100 |
|
| 101 |
def _predict_from_bytes(wav_bytes: bytes) -> dict:
|
| 102 |
-
"""Convert raw WAV bytes to features and run the model."""
|
| 103 |
try:
|
| 104 |
import soundfile as sf
|
| 105 |
import librosa
|
|
@@ -118,20 +140,17 @@ def _predict_from_bytes(wav_bytes: bytes) -> dict:
|
|
| 118 |
return _predict(waveform, sr)
|
| 119 |
|
| 120 |
|
| 121 |
-
# ββ Endpoints βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 122 |
-
|
| 123 |
@app.get("/")
|
| 124 |
@app.get("/health")
|
| 125 |
async def health():
|
| 126 |
return {
|
| 127 |
"status": "ok",
|
| 128 |
-
"model_loaded": _model is not None,
|
| 129 |
}
|
| 130 |
|
| 131 |
|
| 132 |
@app.post("/predict")
|
| 133 |
async def predict(audio: UploadFile = File(...)):
|
| 134 |
-
"""Predict emotion from an uploaded WAV file."""
|
| 135 |
try:
|
| 136 |
content = await audio.read()
|
| 137 |
if not content:
|
|
@@ -146,7 +165,6 @@ async def predict(audio: UploadFile = File(...)):
|
|
| 146 |
|
| 147 |
@app.post("/predict_b64")
|
| 148 |
async def predict_b64(data: str = Form(...)):
|
| 149 |
-
"""Predict emotion from base64-encoded WAV bytes."""
|
| 150 |
try:
|
| 151 |
payload = json.loads(data) if isinstance(data, str) else data
|
| 152 |
if isinstance(payload, dict):
|
|
|
|
| 7 |
from typing import Optional
|
| 8 |
|
| 9 |
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
import torch.nn as nn
|
| 12 |
+
import torch.nn.functional as F
|
| 13 |
from fastapi import FastAPI, File, Form, UploadFile, HTTPException
|
| 14 |
from fastapi.middleware.cors import CORSMiddleware
|
| 15 |
+
from safetensors.torch import load_file
|
| 16 |
+
from transformers import Wav2Vec2Model, Wav2Vec2Config
|
| 17 |
|
| 18 |
app = FastAPI(title="Speech Emotion Recognition")
|
| 19 |
|
|
|
|
| 25 |
allow_headers=["*"],
|
| 26 |
)
|
| 27 |
|
|
|
|
| 28 |
MODEL_CLASSES = ["angry", "disgust", "fear", "happy", "neutral", "pleasant_surprise", "sad"]
|
| 29 |
MODEL_CLASS_INDEX = {c: i for i, c in enumerate(MODEL_CLASSES)}
|
| 30 |
|
|
|
|
| 31 |
_model = None
|
| 32 |
|
| 33 |
|
| 34 |
+
class SERHead(nn.Module):
|
| 35 |
+
def __init__(self):
|
| 36 |
+
super().__init__()
|
| 37 |
+
self.projector = nn.Linear(768, 256)
|
| 38 |
+
self.classifier = nn.Linear(256, 7)
|
| 39 |
+
self.layer_weights = nn.Parameter(torch.ones(13) / 13)
|
| 40 |
+
|
| 41 |
+
def forward(self, hidden_states):
|
| 42 |
+
stacked = torch.stack(list(hidden_states), dim=0)
|
| 43 |
+
w = F.softmax(self.layer_weights, dim=0)
|
| 44 |
+
weighted = (stacked * w.view(-1, 1, 1, 1)).sum(dim=0)
|
| 45 |
+
pooled = weighted.mean(dim=1)
|
| 46 |
+
return self.classifier(F.relu(self.projector(pooled)))
|
| 47 |
+
|
| 48 |
+
|
| 49 |
def load_model():
|
|
|
|
| 50 |
global _model
|
| 51 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 52 |
+
|
| 53 |
+
config = Wav2Vec2Config.from_pretrained("facebook/wav2vec2-base")
|
| 54 |
+
backbone = Wav2Vec2Model(config).to(device).eval()
|
| 55 |
+
head = SERHead().to(device).eval()
|
| 56 |
+
|
| 57 |
+
model_path = "model.safetensors"
|
| 58 |
+
if not os.path.exists(model_path):
|
| 59 |
+
print(f"[WARN] {model_path} not found β using dummy predictor")
|
| 60 |
+
_model = "dummy"
|
| 61 |
+
return
|
| 62 |
+
|
| 63 |
+
state = load_file(model_path)
|
| 64 |
+
backbone.load_state_dict(state, strict=False)
|
| 65 |
+
head.load_state_dict(state, strict=False)
|
| 66 |
+
|
| 67 |
+
_model = {"backbone": backbone, "head": head, "device": device}
|
| 68 |
+
print(f"[INFO] Model loaded on {device}")
|
| 69 |
|
| 70 |
|
| 71 |
@app.on_event("startup")
|
|
|
|
| 74 |
|
| 75 |
|
| 76 |
def _extract_features(waveform: np.ndarray, sr: int) -> np.ndarray:
|
|
|
|
| 77 |
import librosa
|
| 78 |
mfcc = librosa.feature.mfcc(y=waveform, sr=sr, n_mfcc=40, n_fft=512, hop_length=256)
|
| 79 |
mfcc_delta = librosa.feature.delta(mfcc)
|
| 80 |
mfcc_delta2 = librosa.feature.delta(mfcc, order=2)
|
| 81 |
features = np.concatenate([mfcc, mfcc_delta, mfcc_delta2], axis=0)
|
|
|
|
| 82 |
return features.mean(axis=1)
|
| 83 |
|
| 84 |
|
| 85 |
def _predict(waveform: np.ndarray, sr: int) -> dict:
|
| 86 |
+
global _model
|
|
|
|
|
|
|
|
|
|
| 87 |
if _model is None or _model == "dummy":
|
| 88 |
return _dummy_predict(waveform, sr)
|
| 89 |
|
| 90 |
+
backbone = _model["backbone"]
|
| 91 |
+
head = _model["head"]
|
| 92 |
+
device = _model["device"]
|
| 93 |
+
|
| 94 |
+
wav_t = torch.from_numpy(waveform).float().unsqueeze(0).to(device)
|
| 95 |
+
with torch.no_grad():
|
| 96 |
+
outputs = backbone(wav_t, output_hidden_states=True)
|
| 97 |
+
logits = head(outputs.hidden_states)
|
| 98 |
+
probs = F.softmax(logits, dim=-1).squeeze(0)
|
| 99 |
+
|
| 100 |
+
probs_np = probs.cpu().numpy()
|
| 101 |
+
pred_idx = int(probs_np.argmax())
|
| 102 |
+
emotion = MODEL_CLASSES[pred_idx]
|
| 103 |
+
prob_map = {c: round(float(probs_np[i]), 4) for i, c in enumerate(MODEL_CLASSES)}
|
| 104 |
+
|
| 105 |
+
return {
|
| 106 |
+
"emotion": emotion,
|
| 107 |
+
"confidence": round(float(probs_np[pred_idx]), 4),
|
| 108 |
+
"probabilities": prob_map,
|
| 109 |
+
}
|
| 110 |
|
| 111 |
|
| 112 |
def _dummy_predict(waveform: np.ndarray, sr: int) -> dict:
|
|
|
|
| 113 |
rng = np.random.default_rng(42)
|
| 114 |
probs = rng.dirichlet(np.ones(len(MODEL_CLASSES)) * 0.5)
|
| 115 |
probs = probs / probs.sum()
|
|
|
|
| 122 |
|
| 123 |
|
| 124 |
def _predict_from_bytes(wav_bytes: bytes) -> dict:
|
|
|
|
| 125 |
try:
|
| 126 |
import soundfile as sf
|
| 127 |
import librosa
|
|
|
|
| 140 |
return _predict(waveform, sr)
|
| 141 |
|
| 142 |
|
|
|
|
|
|
|
| 143 |
@app.get("/")
|
| 144 |
@app.get("/health")
|
| 145 |
async def health():
|
| 146 |
return {
|
| 147 |
"status": "ok",
|
| 148 |
+
"model_loaded": _model is not None and _model != "dummy",
|
| 149 |
}
|
| 150 |
|
| 151 |
|
| 152 |
@app.post("/predict")
|
| 153 |
async def predict(audio: UploadFile = File(...)):
|
|
|
|
| 154 |
try:
|
| 155 |
content = await audio.read()
|
| 156 |
if not content:
|
|
|
|
| 165 |
|
| 166 |
@app.post("/predict_b64")
|
| 167 |
async def predict_b64(data: str = Form(...)):
|
|
|
|
| 168 |
try:
|
| 169 |
payload = json.loads(data) if isinstance(data, str) else data
|
| 170 |
if isinstance(payload, dict):
|
requirements.txt
CHANGED
|
@@ -4,3 +4,6 @@ python-multipart>=0.0.6
|
|
| 4 |
soundfile>=0.12.1
|
| 5 |
librosa>=0.10.0
|
| 6 |
numpy>=1.24.0
|
|
|
|
|
|
|
|
|
|
|
|
| 4 |
soundfile>=0.12.1
|
| 5 |
librosa>=0.10.0
|
| 6 |
numpy>=1.24.0
|
| 7 |
+
torch>=2.0.0
|
| 8 |
+
transformers>=4.30.0
|
| 9 |
+
safetensors>=0.3.0
|