A7md47 commited on
Commit
a59e52c
·
1 Parent(s): b7278a8

Add v2 pipeline model (prithivMLmods), route by model param

Browse files
Files changed (1) hide show
  1. app.py +70 -91
app.py CHANGED
@@ -13,7 +13,7 @@ import torch.nn.functional as F
13
  from fastapi import FastAPI, HTTPException, Request
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,10 +25,14 @@ app.add_middleware(
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):
@@ -46,72 +50,67 @@ class SERHead(nn.Module):
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
-
65
- # Strip "wav2vec2." prefix from backbone keys
66
  backbone_prefix = "wav2vec2."
67
- backbone_state = {
68
- k[len(backbone_prefix):]: v
69
- for k, v in state.items()
70
- if k.startswith(backbone_prefix)
71
- }
72
- missing, unexpected = backbone.load_state_dict(backbone_state, strict=False)
73
- if missing:
74
- print(f"[WARN] Backbone missing keys: {missing[:5]}...")
75
- if unexpected:
76
- print(f"[WARN] Backbone unexpected keys: {unexpected[:5]}...")
77
-
78
  head.load_state_dict(state, strict=False)
 
 
79
 
80
- # Sanity check: run a tiny forward pass
81
- dummy_input = torch.randn(1, 16000)
82
- with torch.no_grad():
83
- out = backbone(dummy_input, output_hidden_states=True)
84
- logits = head(out.hidden_states)
85
- probs = F.softmax(logits, dim=-1).squeeze(0)
86
- print(f"[INFO] Sanity check — probs: {probs[:3].tolist()} ... {probs[-3:].tolist()}")
87
- print(f"[INFO] Sanity check — argmax: {MODEL_CLASSES[int(probs.argmax())]}")
88
 
89
- _model = {"backbone": backbone, "head": head, "device": device}
90
- print(f"[INFO] Model loaded on {device}")
 
 
 
 
 
 
 
 
 
91
 
92
 
93
  @app.on_event("startup")
94
  async def startup():
95
- load_model()
 
96
 
97
 
98
- def _extract_features(waveform: np.ndarray, sr: int) -> np.ndarray:
 
99
  import librosa
100
- mfcc = librosa.feature.mfcc(y=waveform, sr=sr, n_mfcc=40, n_fft=512, hop_length=256)
101
- mfcc_delta = librosa.feature.delta(mfcc)
102
- mfcc_delta2 = librosa.feature.delta(mfcc, order=2)
103
- features = np.concatenate([mfcc, mfcc_delta, mfcc_delta2], axis=0)
104
- return features.mean(axis=1)
105
-
106
-
107
- def _predict(waveform: np.ndarray, sr: int) -> dict:
108
- global _model
109
- if _model is None or _model == "dummy":
110
- return _dummy_predict(waveform, sr)
 
 
111
 
112
- backbone = _model["backbone"]
113
- head = _model["head"]
114
- device = _model["device"]
115
 
116
  wav_t = torch.from_numpy(waveform).float().unsqueeze(0).to(device)
117
  with torch.no_grad():
@@ -121,57 +120,29 @@ def _predict(waveform: np.ndarray, sr: int) -> dict:
121
 
122
  probs_np = probs.cpu().numpy()
123
  pred_idx = int(probs_np.argmax())
124
- emotion = MODEL_CLASSES[pred_idx]
125
- prob_map = {c: round(float(probs_np[i]), 4) for i, c in enumerate(MODEL_CLASSES)}
 
126
 
127
- return {
128
- "emotion": emotion,
129
- "confidence": round(float(probs_np[pred_idx]), 4),
130
- "probabilities": prob_map,
131
- }
132
 
 
 
 
133
 
134
- def _dummy_predict(waveform: np.ndarray, sr: int) -> dict:
135
- rng = np.random.default_rng(42)
136
- probs = rng.dirichlet(np.ones(len(MODEL_CLASSES)) * 0.5)
137
- probs = probs / probs.sum()
138
- prob_map = {c: round(float(p), 4) for c, p in zip(MODEL_CLASSES, probs)}
139
  return {
140
- "emotion": "neutral",
141
- "confidence": round(float(probs[MODEL_CLASS_INDEX["neutral"]]), 4),
142
- "probabilities": prob_map,
143
  }
144
 
145
 
146
- def _predict_from_bytes(wav_bytes: bytes) -> dict:
147
- try:
148
- import soundfile as sf
149
- import librosa
150
-
151
- buf = io.BytesIO(wav_bytes)
152
- waveform, sr = sf.read(buf)
153
- if waveform.ndim > 1:
154
- waveform = waveform.mean(axis=1)
155
- if sr != 16000:
156
- waveform = librosa.resample(y=waveform, orig_sr=sr, target_sr=16000)
157
- sr = 16000
158
- except Exception as e:
159
- print(f"[WARN] Failed to decode audio: {e}")
160
- return {"emotion": "neutral", "confidence": 0.0, "probabilities": {}}
161
-
162
- return _predict(waveform, sr)
163
-
164
-
165
  @app.get("/")
166
  @app.get("/health")
167
  async def health():
168
- return {
169
- "status": "ok",
170
- "model_loaded": _model is not None and _model != "dummy",
171
- }
172
-
173
-
174
-
175
 
176
 
177
  @app.post("/predict_b64")
@@ -183,8 +154,8 @@ async def predict_b64(request: Request):
183
  if "application/json" in content_type or body.startswith(b"{"):
184
  payload = json.loads(body)
185
  b64_str = payload.get("audio") or payload.get("image", "")
 
186
  else:
187
- # form-encoded with key "data"
188
  import urllib.parse
189
  parsed = urllib.parse.parse_qs(body.decode())
190
  raw = parsed.get("data", [None])[0]
@@ -192,12 +163,20 @@ async def predict_b64(request: Request):
192
  raise HTTPException(status_code=400, detail="Missing 'data' field")
193
  payload = json.loads(raw)
194
  b64_str = payload.get("audio") or payload.get("image", "") or raw
 
195
 
196
  if not b64_str:
197
  raise HTTPException(status_code=400, detail="No audio data found")
198
 
199
  wav_bytes = base64.b64decode(b64_str)
200
- result = _predict_from_bytes(wav_bytes)
 
 
 
 
 
 
 
201
  return result
202
  except HTTPException:
203
  raise
 
13
  from fastapi import FastAPI, HTTPException, Request
14
  from fastapi.middleware.cors import CORSMiddleware
15
  from safetensors.torch import load_file
16
+ from transformers import Wav2Vec2Model, Wav2Vec2Config, pipeline
17
 
18
  app = FastAPI(title="Speech Emotion Recognition")
19
 
 
25
  allow_headers=["*"],
26
  )
27
 
28
+ # ── v1: Custom Wav2Vec2 + SERHead (7 classes) ──────────────────────────
29
+ MODEL_CLASSES_V1 = ["angry", "disgust", "fear", "happy", "neutral", "pleasant_surprise", "sad"]
30
 
31
+ # ── v2: HuggingFace pipeline (8 classes) ────────────────────────────────
32
+ V2_LABELS = ["ANG", "CAL", "DIS", "FEA", "HAP", "NEU", "SAD", "SUR"]
33
+
34
+ _model_v1 = None
35
+ _model_v2 = None
36
 
37
 
38
  class SERHead(nn.Module):
 
50
  return self.classifier(F.relu(self.projector(pooled)))
51
 
52
 
53
+ def load_model_v1():
54
+ global _model_v1
55
  device = "cuda" if torch.cuda.is_available() else "cpu"
 
56
  config = Wav2Vec2Config.from_pretrained("facebook/wav2vec2-base")
57
  backbone = Wav2Vec2Model(config).to(device).eval()
58
  head = SERHead().to(device).eval()
59
 
60
  model_path = "model.safetensors"
61
  if not os.path.exists(model_path):
62
+ print("[WARN] model.safetensors not found — v1 unavailable")
63
+ _model_v1 = "unavailable"
64
  return
65
 
66
  state = load_file(model_path)
 
 
67
  backbone_prefix = "wav2vec2."
68
+ backbone_state = {k[len(backbone_prefix):]: v for k, v in state.items() if k.startswith(backbone_prefix)}
69
+ backbone.load_state_dict(backbone_state, strict=False)
 
 
 
 
 
 
 
 
 
70
  head.load_state_dict(state, strict=False)
71
+ _model_v1 = {"backbone": backbone, "head": head, "device": device}
72
+ print("[INFO] v1 (Wav2Vec2) loaded")
73
 
 
 
 
 
 
 
 
 
74
 
75
+ def load_model_v2():
76
+ global _model_v2
77
+ try:
78
+ _model_v2 = pipeline(
79
+ "audio-classification",
80
+ model="prithivMLmods/Speech-Emotion-Classification",
81
+ )
82
+ print("[INFO] v2 (prithivMLmods) loaded")
83
+ except Exception as e:
84
+ print(f"[WARN] v2 failed to load: {e}")
85
+ _model_v2 = "unavailable"
86
 
87
 
88
  @app.on_event("startup")
89
  async def startup():
90
+ load_model_v1()
91
+ load_model_v2()
92
 
93
 
94
+ def _decode_audio(wav_bytes: bytes):
95
+ import soundfile as sf
96
  import librosa
97
+ buf = io.BytesIO(wav_bytes)
98
+ waveform, sr = sf.read(buf)
99
+ if waveform.ndim > 1:
100
+ waveform = waveform.mean(axis=1)
101
+ if sr != 16000:
102
+ waveform = librosa.resample(y=waveform, orig_sr=sr, target_sr=16000)
103
+ sr = 16000
104
+ return waveform, sr
105
+
106
+
107
+ def _predict_v1(waveform: np.ndarray) -> dict:
108
+ if _model_v1 is None or _model_v1 == "unavailable":
109
+ return {"emotion": "neutral", "confidence": 0.0, "probabilities": {}}
110
 
111
+ backbone = _model_v1["backbone"]
112
+ head = _model_v1["head"]
113
+ device = _model_v1["device"]
114
 
115
  wav_t = torch.from_numpy(waveform).float().unsqueeze(0).to(device)
116
  with torch.no_grad():
 
120
 
121
  probs_np = probs.cpu().numpy()
122
  pred_idx = int(probs_np.argmax())
123
+ emotion = MODEL_CLASSES_V1[pred_idx]
124
+ prob_map = {c: round(float(probs_np[i]), 4) for i, c in enumerate(MODEL_CLASSES_V1)}
125
+ return {"emotion": emotion, "confidence": round(float(probs_np[pred_idx]), 4), "probabilities": prob_map}
126
 
 
 
 
 
 
127
 
128
+ def _predict_v2(waveform: np.ndarray, sr: int) -> dict:
129
+ if _model_v2 is None or _model_v2 == "unavailable":
130
+ return {"emotion": "neutral", "confidence": 0.0, "probabilities": {}}
131
 
132
+ result = _model_v2(waveform, top_k=8)
133
+ probs = {r["label"]: r["score"] for r in result}
134
+ top = result[0]
 
 
135
  return {
136
+ "emotion": top["label"],
137
+ "confidence": round(float(top["score"]), 4),
138
+ "probabilities": probs,
139
  }
140
 
141
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
142
  @app.get("/")
143
  @app.get("/health")
144
  async def health():
145
+ return {"status": "ok", "v1_loaded": _model_v1 is not None and _model_v1 != "unavailable", "v2_loaded": _model_v2 is not None and _model_v2 != "unavailable"}
 
 
 
 
 
 
146
 
147
 
148
  @app.post("/predict_b64")
 
154
  if "application/json" in content_type or body.startswith(b"{"):
155
  payload = json.loads(body)
156
  b64_str = payload.get("audio") or payload.get("image", "")
157
+ model_ver = payload.get("model", "v1")
158
  else:
 
159
  import urllib.parse
160
  parsed = urllib.parse.parse_qs(body.decode())
161
  raw = parsed.get("data", [None])[0]
 
163
  raise HTTPException(status_code=400, detail="Missing 'data' field")
164
  payload = json.loads(raw)
165
  b64_str = payload.get("audio") or payload.get("image", "") or raw
166
+ model_ver = payload.get("model", "v1")
167
 
168
  if not b64_str:
169
  raise HTTPException(status_code=400, detail="No audio data found")
170
 
171
  wav_bytes = base64.b64decode(b64_str)
172
+ waveform, sr = _decode_audio(wav_bytes)
173
+
174
+ if model_ver == "v2":
175
+ result = _predict_v2(waveform, sr)
176
+ else:
177
+ result = _predict_v1(waveform)
178
+
179
+ result["model"] = model_ver
180
  return result
181
  except HTTPException:
182
  raise