A7md47 commited on
Commit
7ad1978
Β·
1 Parent(s): 503d798

Real Wav2Vec2 SER model inference from model.safetensors

Browse files
Files changed (2) hide show
  1. app.py +60 -42
  2. 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
- # TODO: Replace with actual model loading, e.g. PyTorch:
35
- # import torch
36
- # _model = torch.jit.load("model.pt")
37
- # _model.eval()
38
- print("[INFO] No model loaded β€” using dummy predictor")
39
- _model = "dummy"
 
 
 
 
 
 
 
 
 
 
 
 
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
- # ── TODO: Replace with real model inference ───────────────────────
67
- # Example for a PyTorch model that takes MFCC features:
68
- #
69
- # feats = _extract_features(waveform, sr) # shape: (120,)
70
- # import torch
71
- # with torch.no_grad():
72
- # logits = _model(torch.from_numpy(feats).unsqueeze(0))
73
- # probs = torch.softmax(logits, dim=-1).squeeze(0).numpy()
74
- #
75
- # For a TensorFlow/Keras model:
76
- #
77
- # feats = _extract_features(waveform, sr).reshape(1, -1)
78
- # probs = _model.predict(feats, verbose=0)[0]
79
- #
80
- # Then construct response:
81
- # pred_idx = int(probs.argmax())
82
- # emotion = MODEL_CLASSES[pred_idx]
83
- # prob_map = {c: round(float(probs[i]), 4) for i, c in enumerate(MODEL_CLASSES)}
84
-
85
- return _dummy_predict(waveform, sr)
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