Jianshu001 commited on
Commit
c6bc670
·
verified ·
1 Parent(s): 3af371d

Upload pipeline_v2.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. pipeline_v2.py +667 -0
pipeline_v2.py ADDED
@@ -0,0 +1,667 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Voice Correction Pipeline v2.0 (WavLM Fine-tuned)
3
+ ==================================================
4
+ Phoneme-level English pronunciation assessment using a fine-tuned WavLM-Large
5
+ backbone + MLP scoring head, trained on 11K children's speech samples.
6
+
7
+ Improvements over v1.0:
8
+ - Fine-tuned WavLM-Large backbone (vs frozen wav2vec2 + GOP threshold)
9
+ - Learned phoneme scoring (vs heuristic GOP threshold)
10
+ - AUC 0.870 (vs 0.738), F1 0.595 (vs 0.476), Pearson 0.645 (vs 0.372)
11
+
12
+ Usage:
13
+ # Single file
14
+ python pipeline_v2.py --audio audio.mp3 --text "Hello, Peter."
15
+
16
+ # Batch mode
17
+ python pipeline_v2.py --batch --input eval_log.xlsx --audio-dir audio_files/ --output results.json
18
+
19
+ # Evaluate against ground truth
20
+ python pipeline_v2.py --batch --input eval_log.xlsx --audio-dir audio_files/ --evaluate
21
+ """
22
+
23
+ import argparse
24
+ import json
25
+ import re
26
+ import sys
27
+ import warnings
28
+ from pathlib import Path
29
+
30
+ import numpy as np
31
+ import torch
32
+ import torch.nn as nn
33
+ import torchaudio
34
+ import torchaudio.functional as F_audio
35
+ from g2p_en import G2p
36
+ from huggingface_hub import hf_hub_download
37
+ from transformers import Wav2Vec2FeatureExtractor, Wav2Vec2ForCTC, WavLMModel
38
+
39
+ warnings.filterwarnings("ignore")
40
+
41
+ # ============================================================
42
+ # Constants
43
+ # ============================================================
44
+ SAMPLE_RATE = 16000
45
+ DEFAULT_PHERR_THRESHOLD = 0.70 # Calibrated on validation set
46
+
47
+ ARPABET_TO_IPA = {
48
+ "aa": ["ɑː", "ɑ", "ɒ", "a"], "ae": ["æ"], "ah": ["ʌ", "ə", "ɐ"],
49
+ "ao": ["ɔː", "ɔ", "ɒ"], "aw": ["aʊ"], "ax": ["ə", "ɐ", "ʌ"], "ay": ["aɪ"],
50
+ "b": ["b"], "ch": ["tʃ"], "d": ["d"], "dh": ["ð"],
51
+ "eh": ["ɛ", "e"], "er": ["ɜː", "ɝ", "ɚ", "ɜ"], "ey": ["eɪ"],
52
+ "f": ["f"], "g": ["ɡ", "g"], "hh": ["h"],
53
+ "ih": ["ɪ", "ᵻ"], "iy": ["iː", "i"],
54
+ "ir": ["ɪɹ"], "jh": ["dʒ"], "k": ["k"], "l": ["l"],
55
+ "m": ["m"], "n": ["n"], "ng": ["ŋ"],
56
+ "ow": ["oʊ", "o", "əʊ"], "oy": ["ɔɪ"],
57
+ "p": ["p"], "r": ["ɹ", "r"], "s": ["s"], "sh": ["ʃ"],
58
+ "t": ["t"], "th": ["θ"], "uh": ["ʊ"], "uw": ["uː", "u"],
59
+ "ur": ["ʊɹ"], "v": ["v"], "w": ["w"], "y": ["j"],
60
+ "z": ["z"], "zh": ["ʒ"],
61
+ "ar": ["ɑːɹ"], "oo": ["ʊ", "uː"], "dr": ["dɹ"], "tr": ["tɹ"],
62
+ "ts": ["ts"], "dz": ["dz"],
63
+ }
64
+
65
+ ALL_PHONES = sorted(ARPABET_TO_IPA.keys())
66
+ PHONE_TO_ID = {ph: i for i, ph in enumerate(ALL_PHONES)}
67
+ N_PHONE_TYPES = len(ALL_PHONES)
68
+
69
+
70
+ # ============================================================
71
+ # MLP Scoring Head (must match training architecture)
72
+ # ============================================================
73
+ class PhoneScorerHead(nn.Module):
74
+ def __init__(self, hidden_dim=1024, n_phone_types=N_PHONE_TYPES,
75
+ phone_emb_dim=32, mlp_dim=512):
76
+ super().__init__()
77
+ self.phone_emb = nn.Embedding(n_phone_types, phone_emb_dim)
78
+ input_dim = hidden_dim + phone_emb_dim + 2 # +2 for GOP and n_frames
79
+
80
+ self.shared = nn.Sequential(
81
+ nn.Linear(input_dim, mlp_dim),
82
+ nn.BatchNorm1d(mlp_dim),
83
+ nn.GELU(),
84
+ nn.Dropout(0.3),
85
+ nn.Linear(mlp_dim, mlp_dim),
86
+ nn.BatchNorm1d(mlp_dim),
87
+ nn.GELU(),
88
+ nn.Dropout(0.3),
89
+ nn.Linear(mlp_dim, 256),
90
+ nn.BatchNorm1d(256),
91
+ nn.GELU(),
92
+ nn.Dropout(0.2),
93
+ )
94
+ self.score_head = nn.Sequential(nn.Linear(256, 64), nn.GELU(), nn.Linear(64, 1))
95
+ self.pherr_head = nn.Sequential(nn.Linear(256, 64), nn.GELU(), nn.Linear(64, 1))
96
+
97
+ def forward(self, h, phone_id, gop, n_frames):
98
+ emb = self.phone_emb(phone_id)
99
+ x = torch.cat([h, emb, gop.unsqueeze(-1), n_frames.unsqueeze(-1)], dim=-1)
100
+ shared = self.shared(x)
101
+ score = self.score_head(shared).squeeze(-1)
102
+ pherr_logit = self.pherr_head(shared).squeeze(-1)
103
+ return score, pherr_logit
104
+
105
+
106
+ # ============================================================
107
+ # Main Pipeline
108
+ # ============================================================
109
+ class PronunciationAssessorV2:
110
+ """Phoneme-level pronunciation assessment using fine-tuned WavLM backbone."""
111
+
112
+ def __init__(self, checkpoint_path=None, device=None, pherr_threshold=DEFAULT_PHERR_THRESHOLD):
113
+ self.device = device or torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
114
+ self.pherr_threshold = pherr_threshold
115
+ self._checkpoint_path = checkpoint_path or self._default_checkpoint()
116
+ self._backbone = None
117
+ self._scorer = None
118
+ self._fe_backbone = None
119
+ self._ctc_model = None
120
+ self._fe_ctc = None
121
+ self._vocab = None
122
+ self._blank_idx = None
123
+ self._g2p = None
124
+
125
+ @staticmethod
126
+ def _default_checkpoint():
127
+ return str(Path(__file__).parent / "wavlm_finetuned.pt")
128
+
129
+ def _load_models(self):
130
+ if self._backbone is not None:
131
+ return
132
+
133
+ print("Loading WavLM backbone + scoring head...", file=sys.stderr)
134
+
135
+ # Load checkpoint
136
+ ckpt = torch.load(self._checkpoint_path, map_location="cpu", weights_only=False)
137
+ state_dict = ckpt["model_state"]
138
+
139
+ # Separate backbone and head state dicts
140
+ backbone_state = {}
141
+ head_state = {}
142
+ for k, v in state_dict.items():
143
+ if k.startswith("backbone."):
144
+ backbone_state[k[len("backbone."):]] = v
145
+ else:
146
+ head_state[k] = v
147
+
148
+ # Load WavLM backbone
149
+ self._backbone = WavLMModel.from_pretrained(
150
+ "microsoft/wavlm-large",
151
+ output_hidden_states=False,
152
+ mask_time_prob=0.0,
153
+ )
154
+ self._backbone.load_state_dict(backbone_state, strict=False)
155
+ self._backbone.to(self.device)
156
+ self._backbone.eval()
157
+
158
+ self._fe_backbone = Wav2Vec2FeatureExtractor.from_pretrained("microsoft/wavlm-large")
159
+
160
+ # Load scoring head
161
+ self._scorer = PhoneScorerHead()
162
+ self._scorer.load_state_dict(head_state)
163
+ self._scorer.to(self.device)
164
+ self._scorer.eval()
165
+
166
+ # Load CTC model for alignment
167
+ print("Loading CTC alignment model...", file=sys.stderr)
168
+ ctc_name = "facebook/wav2vec2-xlsr-53-espeak-cv-ft"
169
+ self._ctc_model = Wav2Vec2ForCTC.from_pretrained(ctc_name).to(self.device)
170
+ self._ctc_model.eval()
171
+ self._fe_ctc = Wav2Vec2FeatureExtractor.from_pretrained(ctc_name)
172
+
173
+ vocab_path = hf_hub_download(ctc_name, "vocab.json")
174
+ with open(vocab_path) as f:
175
+ self._vocab = json.load(f)
176
+ self._blank_idx = self._vocab.get("<pad>", 0)
177
+
178
+ self._g2p = G2p()
179
+ print(f"Models loaded (device={self.device})", file=sys.stderr)
180
+
181
+ # --------------------------------------------------------
182
+ # Audio
183
+ # --------------------------------------------------------
184
+ @staticmethod
185
+ def load_audio(audio_path):
186
+ waveform, sr = torchaudio.load(audio_path)
187
+ if waveform.shape[0] > 1:
188
+ waveform = waveform.mean(0, keepdim=True)
189
+ if sr != SAMPLE_RATE:
190
+ waveform = F_audio.resample(waveform, sr, SAMPLE_RATE)
191
+ return waveform
192
+
193
+ # --------------------------------------------------------
194
+ # G2P
195
+ # --------------------------------------------------------
196
+ def _text_to_phonemes(self, text):
197
+ words = re.sub(r"[^\w' ]", " ", text).split()
198
+ result = []
199
+ for word_idx, word in enumerate(words):
200
+ phones_raw = self._g2p(word)
201
+ for ph in phones_raw:
202
+ if ph == " ":
203
+ continue
204
+ clean = re.sub(r"\d", "", ph).lower()
205
+ if clean:
206
+ result.append({"phone": clean, "word": word, "word_idx": word_idx})
207
+ return result
208
+
209
+ def _arpabet_to_model_idx(self, ph):
210
+ for ipa in ARPABET_TO_IPA.get(ph, []):
211
+ if ipa in self._vocab:
212
+ return self._vocab[ipa]
213
+ return self._vocab.get(ph, -1)
214
+
215
+ # --------------------------------------------------------
216
+ # Viterbi forced alignment
217
+ # --------------------------------------------------------
218
+ @staticmethod
219
+ def _viterbi_align(emissions, phone_indices, blank_idx):
220
+ T, C = emissions.shape
221
+ S = len(phone_indices)
222
+ if S == 0 or T < S:
223
+ return []
224
+
225
+ extended = [blank_idx]
226
+ for p in phone_indices:
227
+ extended.append(p)
228
+ extended.append(blank_idx)
229
+ S_ext = len(extended)
230
+
231
+ NEG_INF = float("-inf")
232
+ dp = np.full((T, S_ext), NEG_INF, dtype=np.float64)
233
+ bp = np.zeros((T, S_ext), dtype=np.int32)
234
+ dp[0][0] = emissions[0, extended[0]].item()
235
+ if S_ext > 1:
236
+ dp[0][1] = emissions[0, extended[1]].item()
237
+
238
+ for t in range(1, T):
239
+ for s in range(S_ext):
240
+ emit = emissions[t, extended[s]].item()
241
+ best, best_s = dp[t - 1][s], s
242
+ if s > 0 and dp[t - 1][s - 1] > best:
243
+ best, best_s = dp[t - 1][s - 1], s - 1
244
+ if s > 1 and extended[s] != blank_idx and extended[s] != extended[s - 2]:
245
+ if dp[t - 1][s - 2] > best:
246
+ best, best_s = dp[t - 1][s - 2], s - 2
247
+ dp[t][s] = best + emit
248
+ bp[t][s] = best_s
249
+
250
+ s = S_ext - 1 if (S_ext >= 2 and dp[T - 1][S_ext - 1] >= dp[T - 1][S_ext - 2]) else max(S_ext - 2, 0)
251
+ path = []
252
+ for t in range(T - 1, -1, -1):
253
+ path.append((t, extended[s]))
254
+ s = bp[t][s]
255
+ path.reverse()
256
+ return path
257
+
258
+ # --------------------------------------------------------
259
+ # Core: align + extract features + score
260
+ # --------------------------------------------------------
261
+ @torch.no_grad()
262
+ def _score_phonemes(self, waveform, phone_indices):
263
+ """
264
+ Given waveform and expected phone indices, returns per-phoneme scores.
265
+ Steps:
266
+ 1. CTC emissions → Viterbi alignment → frame segments
267
+ 2. WavLM backbone → hidden states per segment
268
+ 3. GOP from CTC emissions
269
+ 4. MLP head → score + pherr probability
270
+ """
271
+ wav_np = waveform.squeeze(0).numpy()
272
+
273
+ # CTC emissions for alignment
274
+ ctc_inputs = self._fe_ctc(wav_np, sampling_rate=SAMPLE_RATE, return_tensors="pt", padding=True)
275
+ ctc_logits = self._ctc_model(ctc_inputs.input_values.to(self.device)).logits
276
+ emissions = torch.log_softmax(ctc_logits, dim=-1).squeeze(0).cpu()
277
+
278
+ # Viterbi alignment
279
+ path = self._viterbi_align(emissions, phone_indices, self._blank_idx)
280
+ if not path:
281
+ return None
282
+
283
+ # Group frames by phoneme
284
+ segments = []
285
+ cur_tok, cur_frames = None, []
286
+ for f, tok in path:
287
+ if tok == self._blank_idx:
288
+ if cur_tok is not None:
289
+ segments.append((cur_tok, cur_frames))
290
+ cur_tok, cur_frames = None, []
291
+ continue
292
+ if tok != cur_tok:
293
+ if cur_tok is not None:
294
+ segments.append((cur_tok, cur_frames))
295
+ cur_tok, cur_frames = tok, [f]
296
+ else:
297
+ cur_frames.append(f)
298
+ if cur_tok is not None:
299
+ segments.append((cur_tok, cur_frames))
300
+
301
+ # WavLM backbone hidden states
302
+ bb_inputs = self._fe_backbone(wav_np, sampling_rate=SAMPLE_RATE, return_tensors="pt", padding=True)
303
+ hidden = self._backbone(bb_inputs.input_values.to(self.device)).last_hidden_state.squeeze(0) # (T_h, 1024)
304
+ T_h = hidden.shape[0]
305
+ T_ctc = emissions.shape[0]
306
+ scale = T_h / T_ctc
307
+
308
+ # Per-phoneme: pool hidden states, compute GOP, run MLP
309
+ results = []
310
+ for i, expected_idx in enumerate(phone_indices):
311
+ if i >= len(segments):
312
+ results.append({"gop": -20.0, "score": 0.0, "pherr_prob": 1.0})
313
+ continue
314
+
315
+ _, frames = segments[i]
316
+
317
+ # Pool hidden states
318
+ h_start = max(0, int(min(frames) * scale))
319
+ h_end = min(T_h, int((max(frames) + 1) * scale))
320
+ if h_end <= h_start:
321
+ h_end = h_start + 1
322
+ h_pooled = hidden[h_start:h_end].mean(dim=0) # (1024,)
323
+
324
+ # GOP
325
+ seg_em = emissions[frames]
326
+ target_lp = seg_em[:, expected_idx].mean().item()
327
+ mask = torch.ones(emissions.shape[1], dtype=torch.bool)
328
+ mask[self._blank_idx] = False
329
+ mask[expected_idx] = False
330
+ best_other = seg_em[:, mask].max(dim=-1).values.mean().item()
331
+ gop = target_lp - best_other
332
+
333
+ results.append({
334
+ "h": h_pooled,
335
+ "gop": gop,
336
+ "n_frames": len(frames),
337
+ })
338
+
339
+ # Batch MLP inference
340
+ valid_indices = [i for i, r in enumerate(results) if "h" in r]
341
+ if valid_indices:
342
+ h_batch = torch.stack([results[i]["h"] for i in valid_indices]).to(self.device)
343
+ gop_batch = torch.tensor([results[i]["gop"] for i in valid_indices],
344
+ dtype=torch.float32).to(self.device)
345
+ nf_batch = torch.tensor([results[i]["n_frames"] for i in valid_indices],
346
+ dtype=torch.float32).to(self.device)
347
+
348
+ # Need phone_ids for the valid phonemes
349
+ pid_batch = torch.tensor([PHONE_TO_ID.get(
350
+ # We need to pass phone names - will be set from caller
351
+ "_placeholder_", 0) for _ in valid_indices],
352
+ dtype=torch.long).to(self.device)
353
+
354
+ # Return raw data for batch processing in assess()
355
+ return results, valid_indices, h_batch, gop_batch, nf_batch
356
+
357
+ return results, [], None, None, None
358
+
359
+ # --------------------------------------------------------
360
+ # Public API
361
+ # --------------------------------------------------------
362
+ def assess(self, audio_path, text):
363
+ """
364
+ Assess pronunciation of an audio file against reference text.
365
+
366
+ Returns dict with overall_score, words (with per-phoneme scores and errors).
367
+ """
368
+ self._load_models()
369
+
370
+ # G2P
371
+ phone_info = self._text_to_phonemes(text)
372
+ if not phone_info:
373
+ return {"text": text, "overall_score": 0, "words": [], "error": "No phonemes extracted"}
374
+
375
+ # Map to model indices
376
+ indices = []
377
+ valid_info = []
378
+ for pi in phone_info:
379
+ idx = self._arpabet_to_model_idx(pi["phone"])
380
+ if idx >= 0:
381
+ indices.append(idx)
382
+ valid_info.append(pi)
383
+ if not indices:
384
+ return {"text": text, "overall_score": 0, "words": [], "error": "No phonemes mapped"}
385
+
386
+ # Load audio & score
387
+ waveform = self.load_audio(audio_path)
388
+ result = self._score_phonemes(waveform, indices)
389
+ if result is None:
390
+ return {"text": text, "overall_score": 0, "words": [], "error": "Alignment failed"}
391
+
392
+ raw_results, valid_idx, h_batch, gop_batch, nf_batch = result
393
+
394
+ # Run MLP scoring
395
+ if h_batch is not None:
396
+ pid_list = [PHONE_TO_ID.get(valid_info[i]["phone"], 0) for i in valid_idx]
397
+ pid_batch = torch.tensor(pid_list, dtype=torch.long).to(self.device)
398
+
399
+ pred_score, pred_pherr_logit = self._scorer(h_batch, pid_batch, gop_batch, nf_batch)
400
+ pred_score = pred_score.detach().cpu().numpy()
401
+ pred_pherr = torch.sigmoid(pred_pherr_logit).detach().cpu().numpy()
402
+
403
+ for j, i in enumerate(valid_idx):
404
+ raw_results[i]["score"] = float(np.clip(pred_score[j], 0, 100))
405
+ raw_results[i]["pherr_prob"] = float(pred_pherr[j])
406
+
407
+ # Assemble per-phoneme results
408
+ phoneme_results = []
409
+ for i, (info, raw) in enumerate(zip(valid_info, raw_results)):
410
+ score = raw.get("score", 0.0)
411
+ pherr_prob = raw.get("pherr_prob", 1.0)
412
+ phoneme_results.append({
413
+ "phone": info["phone"],
414
+ "word": info["word"],
415
+ "word_idx": info["word_idx"],
416
+ "score": round(score, 1),
417
+ "gop": round(raw["gop"], 3),
418
+ "pherr_prob": round(pherr_prob, 3),
419
+ "error": pherr_prob >= self.pherr_threshold,
420
+ })
421
+
422
+ # Group by word
423
+ words_dict = {}
424
+ for pr in phoneme_results:
425
+ widx = pr["word_idx"]
426
+ if widx not in words_dict:
427
+ words_dict[widx] = {"word": pr["word"], "phonemes": []}
428
+ words_dict[widx]["phonemes"].append({
429
+ "phone": pr["phone"],
430
+ "score": pr["score"],
431
+ "gop": pr["gop"],
432
+ "pherr_prob": pr["pherr_prob"],
433
+ "error": pr["error"],
434
+ })
435
+
436
+ # Word-level scores
437
+ words_list = []
438
+ for widx in sorted(words_dict.keys()):
439
+ wd = words_dict[widx]
440
+ scores = [p["score"] for p in wd["phonemes"]]
441
+ n_errors = sum(1 for p in wd["phonemes"] if p["error"])
442
+ mean_score = np.mean(scores)
443
+ words_list.append({
444
+ "word": wd["word"],
445
+ "score": round(float(mean_score), 1),
446
+ "n_errors": n_errors,
447
+ "n_phonemes": len(wd["phonemes"]),
448
+ "has_error": n_errors > 0,
449
+ "phonemes": wd["phonemes"],
450
+ })
451
+
452
+ # Overall score
453
+ all_scores = [pr["score"] for pr in phoneme_results]
454
+ overall_score = np.mean(all_scores)
455
+ n_total_errors = sum(1 for pr in phoneme_results if pr["error"])
456
+
457
+ return {
458
+ "text": text,
459
+ "overall_score": round(float(overall_score), 1),
460
+ "n_phonemes": len(phoneme_results),
461
+ "n_errors": n_total_errors,
462
+ "error_rate": round(n_total_errors / len(phoneme_results) * 100, 1),
463
+ "words": words_list,
464
+ }
465
+
466
+ def assess_batch(self, items, show_progress=True):
467
+ """Assess a list of (audio_path, text) pairs."""
468
+ self._load_models()
469
+ results = []
470
+ for i, (audio_path, text) in enumerate(items):
471
+ try:
472
+ result = self.assess(audio_path, text)
473
+ results.append(result)
474
+ except Exception as e:
475
+ results.append({"text": text, "overall_score": 0, "error": str(e)})
476
+ if show_progress and (i + 1) % 50 == 0:
477
+ print(f" Processed {i + 1}/{len(items)}...", file=sys.stderr)
478
+ return results
479
+
480
+
481
+ # ============================================================
482
+ # Pretty print
483
+ # ============================================================
484
+ def print_result(result):
485
+ print(f"\n{'='*60}")
486
+ print(f"Text: \"{result['text']}\"")
487
+ print(f"Overall Score: {result['overall_score']:.1f}/100 "
488
+ f"(errors: {result.get('n_errors', '?')}/{result.get('n_phonemes', '?')})")
489
+ print(f"{'='*60}")
490
+
491
+ for wd in result.get("words", []):
492
+ status = "\u2717" if wd["has_error"] else "\u2713"
493
+ print(f"\n {status} {wd['word']:<15s} score={wd['score']:5.1f} "
494
+ f"errors={wd['n_errors']}/{wd['n_phonemes']}")
495
+
496
+ for ph in wd["phonemes"]:
497
+ marker = " \u2190 ERROR" if ph["error"] else ""
498
+ print(f" /{ph['phone']:<4s}/ score={ph['score']:5.1f} "
499
+ f"GOP={ph['gop']:+6.2f} pherr={ph['pherr_prob']:.2f}{marker}")
500
+
501
+
502
+ # ============================================================
503
+ # CLI
504
+ # ============================================================
505
+ def main():
506
+ parser = argparse.ArgumentParser(description="Voice Correction Pipeline v2.0 (WavLM Fine-tuned)")
507
+ parser.add_argument("--audio", type=str, help="Path to audio file")
508
+ parser.add_argument("--text", type=str, help="Reference text")
509
+ parser.add_argument("--checkpoint", type=str, default=None,
510
+ help="Path to fine-tuned model checkpoint (default: wavlm_finetuned.pt)")
511
+ parser.add_argument("--threshold", type=float, default=DEFAULT_PHERR_THRESHOLD,
512
+ help=f"Pherr probability threshold (default: {DEFAULT_PHERR_THRESHOLD})")
513
+ parser.add_argument("--batch", action="store_true", help="Batch mode")
514
+ parser.add_argument("--input", type=str, help="Input xlsx file (batch mode)")
515
+ parser.add_argument("--audio-dir", type=str, help="Audio directory (batch mode)")
516
+ parser.add_argument("--output", type=str, help="Output JSON file")
517
+ parser.add_argument("--limit", type=int, default=0, help="Limit number of samples (0=all)")
518
+ parser.add_argument("--evaluate", action="store_true", help="Evaluate against ground truth")
519
+ parser.add_argument("--json", action="store_true", help="Output raw JSON")
520
+ parser.add_argument("--device", type=str, default=None, help="Device (cuda:0, cpu)")
521
+ args = parser.parse_args()
522
+
523
+ device = torch.device(args.device) if args.device else None
524
+ assessor = PronunciationAssessorV2(
525
+ checkpoint_path=args.checkpoint, device=device, pherr_threshold=args.threshold
526
+ )
527
+
528
+ if args.batch:
529
+ if not args.input:
530
+ parser.error("--input required for batch mode")
531
+
532
+ import openpyxl
533
+ audio_dir = Path(args.audio_dir) if args.audio_dir else Path(args.input).parent / "audio_files"
534
+
535
+ wb = openpyxl.load_workbook(args.input, read_only=True)
536
+ ws = wb.active
537
+ items = []
538
+ for row in ws.iter_rows(min_row=2, values_only=True):
539
+ fname, content = row[0], row[1]
540
+ if not fname or not content:
541
+ continue
542
+ audio_path = audio_dir / fname
543
+ if audio_path.exists():
544
+ items.append((str(audio_path), content))
545
+ wb.close()
546
+
547
+ if args.limit > 0:
548
+ items = items[:args.limit]
549
+
550
+ print(f"Processing {len(items)} files...", file=sys.stderr)
551
+ results = assessor.assess_batch(items)
552
+
553
+ if args.output:
554
+ with open(args.output, "w") as f:
555
+ json.dump(results, f, indent=2, ensure_ascii=False)
556
+ print(f"\nResults saved to {args.output}", file=sys.stderr)
557
+ else:
558
+ for r in results:
559
+ if args.json:
560
+ print(json.dumps(r, ensure_ascii=False))
561
+ else:
562
+ print_result(r)
563
+
564
+ if args.evaluate:
565
+ _run_evaluation(results, args.input)
566
+
567
+ else:
568
+ if not args.audio or not args.text:
569
+ parser.error("--audio and --text are required")
570
+
571
+ result = assessor.assess(args.audio, args.text)
572
+ if args.json:
573
+ print(json.dumps(result, indent=2, ensure_ascii=False))
574
+ else:
575
+ print_result(result)
576
+
577
+
578
+ def _run_evaluation(results, xlsx_path):
579
+ """Compare pipeline output against ground truth."""
580
+ import openpyxl
581
+ from sklearn.metrics import precision_recall_fscore_support, roc_auc_score
582
+ from scipy import stats
583
+
584
+ wb = openpyxl.load_workbook(xlsx_path, read_only=True)
585
+ ws = wb.active
586
+
587
+ gt_records = []
588
+ for row in ws.iter_rows(min_row=2, values_only=True):
589
+ raw = row[2]
590
+ if not raw:
591
+ continue
592
+ s = raw
593
+ for _ in range(5):
594
+ s = s.replace("\\\\", "\\")
595
+ s = s.replace('\\"', '"')
596
+ acc = re.search(r'"accuracy":\s*([\d.]+)', s)
597
+ phones = re.findall(
598
+ r'"char":"([^"]+)","ph2alpha":"([^"]*)".*?"pherr":(\d+).*?"score":([\d.]+)', s
599
+ )
600
+ gt_records.append({
601
+ "accuracy": float(acc.group(1)) if acc else 0,
602
+ "phones": [{"pherr": int(e), "score": float(sc)} for _, _, e, sc in phones],
603
+ })
604
+ wb.close()
605
+
606
+ # Overall accuracy correlation
607
+ n = min(len(gt_records), len(results))
608
+ gt_acc = np.array([r["accuracy"] for r in gt_records[:n]])
609
+ pred_acc = np.array([r.get("overall_score", 0) for r in results[:n]])
610
+ mae_acc = np.mean(np.abs(gt_acc - pred_acc))
611
+ corr_acc, _ = stats.pearsonr(gt_acc, pred_acc)
612
+
613
+ # Phoneme-level metrics
614
+ all_gt_pherr, all_pred_pherr = [], []
615
+ all_gt_score, all_pred_score = [], []
616
+
617
+ for i in range(n):
618
+ if "error" in results[i]:
619
+ continue
620
+ gt_phones = gt_records[i]["phones"]
621
+ pred_words = results[i].get("words", [])
622
+ pred_phones = []
623
+ for w in pred_words:
624
+ pred_phones.extend(w.get("phonemes", []))
625
+
626
+ m = min(len(gt_phones), len(pred_phones))
627
+ for j in range(m):
628
+ all_gt_pherr.append(gt_phones[j]["pherr"])
629
+ all_pred_pherr.append(pred_phones[j]["pherr_prob"])
630
+ all_gt_score.append(gt_phones[j]["score"])
631
+ all_pred_score.append(pred_phones[j]["score"])
632
+
633
+ gt_pherr = np.array(all_gt_pherr)
634
+ pred_pherr = np.array(all_pred_pherr)
635
+ gt_score = np.array(all_gt_score)
636
+ pred_score = np.array(all_pred_score)
637
+
638
+ auc = roc_auc_score(gt_pherr, pred_pherr) if len(np.unique(gt_pherr)) > 1 else 0
639
+
640
+ best_f1, best_th = 0, 0.5
641
+ for th in np.arange(0.1, 0.9, 0.05):
642
+ pb = (pred_pherr >= th).astype(int)
643
+ _, _, f1, _ = precision_recall_fscore_support(gt_pherr, pb, average="binary", zero_division=0)
644
+ if f1 > best_f1:
645
+ best_f1, best_th = f1, th
646
+
647
+ pb = (pred_pherr >= best_th).astype(int)
648
+ prec, rec, f1, _ = precision_recall_fscore_support(gt_pherr, pb, average="binary")
649
+
650
+ corr_phone, _ = stats.pearsonr(gt_score, pred_score)
651
+ mae_phone = np.mean(np.abs(gt_score - pred_score))
652
+
653
+ print(f"\n{'='*60}", file=sys.stderr)
654
+ print(f"EVALUATION ({n} samples, {len(gt_pherr)} phonemes)", file=sys.stderr)
655
+ print(f"{'='*60}", file=sys.stderr)
656
+ print(f" Overall accuracy MAE: {mae_acc:.2f}", file=sys.stderr)
657
+ print(f" Overall accuracy Pearson: {corr_acc:.3f}", file=sys.stderr)
658
+ print(f"\n Phoneme error AUC-ROC: {auc:.3f}", file=sys.stderr)
659
+ print(f" Phoneme error F1: {f1:.3f} (threshold={best_th:.2f})", file=sys.stderr)
660
+ print(f" Phoneme error Precision: {prec:.3f}", file=sys.stderr)
661
+ print(f" Phoneme error Recall: {rec:.3f}", file=sys.stderr)
662
+ print(f"\n Phone score Pearson: {corr_phone:.3f}", file=sys.stderr)
663
+ print(f" Phone score MAE: {mae_phone:.2f}", file=sys.stderr)
664
+
665
+
666
+ if __name__ == "__main__":
667
+ main()