Jianshu001 commited on
Commit
7b9b23e
·
verified ·
1 Parent(s): c6bc670

Upload finetune_wavlm.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. finetune_wavlm.py +748 -0
finetune_wavlm.py ADDED
@@ -0,0 +1,748 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Fine-tune WavLM-Large Backbone for Phoneme Scoring
3
+ ===================================================
4
+ Same as finetune_backbone.py but uses WavLM-Large instead of wav2vec2-large.
5
+ WavLM has denoising pre-training, making it more robust for non-standard speech (e.g. children).
6
+ """
7
+
8
+ import json
9
+ import re
10
+ import warnings
11
+ from pathlib import Path
12
+
13
+ import numpy as np
14
+ import openpyxl
15
+ import torch
16
+ import torch.nn as nn
17
+ import torch.nn.functional as F
18
+ import torchaudio
19
+ import torchaudio.functional as F_audio
20
+ from g2p_en import G2p
21
+ from huggingface_hub import hf_hub_download
22
+ from sklearn.metrics import precision_recall_fscore_support, roc_auc_score
23
+ from sklearn.model_selection import train_test_split
24
+ from scipy import stats
25
+ from torch.utils.data import DataLoader, Dataset
26
+ from transformers import Wav2Vec2FeatureExtractor, Wav2Vec2Model, Wav2Vec2ForCTC, WavLMModel
27
+
28
+ warnings.filterwarnings("ignore")
29
+
30
+ SAMPLE_RATE = 16000
31
+
32
+ ARPABET_TO_IPA = {
33
+ "aa": ["ɑː", "ɑ", "ɒ", "a"], "ae": ["æ"], "ah": ["ʌ", "ə", "ɐ"],
34
+ "ao": ["ɔː", "ɔ", "ɒ"], "aw": ["aʊ"], "ax": ["ə", "ɐ", "ʌ"], "ay": ["aɪ"],
35
+ "b": ["b"], "ch": ["tʃ"], "d": ["d"], "dh": ["ð"],
36
+ "eh": ["ɛ", "e"], "er": ["ɜː", "ɝ", "ɚ", "ɜ"], "ey": ["eɪ"],
37
+ "f": ["f"], "g": ["ɡ", "g"], "hh": ["h"],
38
+ "ih": ["ɪ", "ᵻ"], "iy": ["iː", "i"],
39
+ "ir": ["ɪɹ"], "jh": ["dʒ"], "k": ["k"], "l": ["l"],
40
+ "m": ["m"], "n": ["n"], "ng": ["ŋ"],
41
+ "ow": ["oʊ", "o", "əʊ"], "oy": ["ɔɪ"],
42
+ "p": ["p"], "r": ["ɹ", "r"], "s": ["s"], "sh": ["ʃ"],
43
+ "t": ["t"], "th": ["θ"], "uh": ["ʊ"], "uw": ["uː", "u"],
44
+ "ur": ["ʊɹ"], "v": ["v"], "w": ["w"], "y": ["j"],
45
+ "z": ["z"], "zh": ["ʒ"],
46
+ "ar": ["ɑːɹ"], "oo": ["ʊ", "uː"], "dr": ["dɹ"], "tr": ["tɹ"],
47
+ "ts": ["ts"], "dz": ["dz"],
48
+ }
49
+
50
+ ALL_PHONES = sorted(ARPABET_TO_IPA.keys())
51
+ PHONE_TO_ID = {ph: i for i, ph in enumerate(ALL_PHONES)}
52
+ N_PHONE_TYPES = len(ALL_PHONES)
53
+
54
+
55
+ # ============================================================
56
+ # Model: Backbone + Scoring Head (end-to-end)
57
+ # ============================================================
58
+ class BackbonePhoneScorer(nn.Module):
59
+ """
60
+ End-to-end model: WavLM-Large backbone → per-phoneme pooling → MLP scorer
61
+ """
62
+ def __init__(self, n_phone_types=N_PHONE_TYPES, phone_emb_dim=32,
63
+ hidden_dim=1024, mlp_dim=512, unfreeze_top_n=6):
64
+ super().__init__()
65
+
66
+ # Load WavLM-Large backbone (disable time masking for fine-tuning)
67
+ self.backbone = WavLMModel.from_pretrained(
68
+ "microsoft/wavlm-large",
69
+ output_hidden_states=False,
70
+ mask_time_prob=0.0,
71
+ )
72
+ # Freeze all layers first
73
+ for param in self.backbone.parameters():
74
+ param.requires_grad = False
75
+
76
+ # Unfreeze top N transformer layers
77
+ n_layers = len(self.backbone.encoder.layers)
78
+ for i in range(n_layers - unfreeze_top_n, n_layers):
79
+ for param in self.backbone.encoder.layers[i].parameters():
80
+ param.requires_grad = True
81
+
82
+ # Also unfreeze layer norm
83
+ if hasattr(self.backbone.encoder, 'layer_norm'):
84
+ for param in self.backbone.encoder.layer_norm.parameters():
85
+ param.requires_grad = True
86
+
87
+ self.fe_backbone = Wav2Vec2FeatureExtractor.from_pretrained("microsoft/wavlm-large")
88
+
89
+ # Phone embedding + MLP scorer (same as PhoneScorerV2)
90
+ self.phone_emb = nn.Embedding(n_phone_types, phone_emb_dim)
91
+ input_dim = hidden_dim + phone_emb_dim + 2 # +2 for GOP and n_frames
92
+
93
+ self.shared = nn.Sequential(
94
+ nn.Linear(input_dim, mlp_dim),
95
+ nn.BatchNorm1d(mlp_dim),
96
+ nn.GELU(),
97
+ nn.Dropout(0.3),
98
+ nn.Linear(mlp_dim, mlp_dim),
99
+ nn.BatchNorm1d(mlp_dim),
100
+ nn.GELU(),
101
+ nn.Dropout(0.3),
102
+ nn.Linear(mlp_dim, 256),
103
+ nn.BatchNorm1d(256),
104
+ nn.GELU(),
105
+ nn.Dropout(0.2),
106
+ )
107
+ self.score_head = nn.Sequential(nn.Linear(256, 64), nn.GELU(), nn.Linear(64, 1))
108
+ self.pherr_head = nn.Sequential(nn.Linear(256, 64), nn.GELU(), nn.Linear(64, 1))
109
+
110
+ def forward(self, h, phone_id, gop, n_frames):
111
+ """Forward pass with pre-extracted hidden states (for batched training)."""
112
+ emb = self.phone_emb(phone_id)
113
+ x = torch.cat([h, emb, gop.unsqueeze(-1), n_frames.unsqueeze(-1)], dim=-1)
114
+ shared = self.shared(x)
115
+ score = self.score_head(shared).squeeze(-1)
116
+ pherr = self.pherr_head(shared).squeeze(-1)
117
+ return score, pherr
118
+
119
+ def extract_hidden(self, waveform):
120
+ """Extract hidden states from raw waveform. waveform: (1, T) tensor."""
121
+ # mask_time_prob=0.0 in config disables masking
122
+ out = self.backbone(waveform)
123
+ return out.last_hidden_state # (1, T_frames, 1024)
124
+
125
+ def backbone_params(self):
126
+ """Return only the trainable backbone parameters."""
127
+ for name, param in self.backbone.named_parameters():
128
+ if param.requires_grad:
129
+ yield param
130
+
131
+ def head_params(self):
132
+ """Return head parameters (phone_emb + MLP)."""
133
+ yield from self.phone_emb.parameters()
134
+ yield from self.shared.parameters()
135
+ yield from self.score_head.parameters()
136
+ yield from self.pherr_head.parameters()
137
+
138
+
139
+ # ============================================================
140
+ # Alignment engine (frozen CTC model)
141
+ # ============================================================
142
+ class AlignmentEngine:
143
+ def __init__(self, device):
144
+ self.device = device
145
+ self.g2p = G2p()
146
+
147
+ print("Loading phoneme CTC model for alignment...")
148
+ self.ctc_model = Wav2Vec2ForCTC.from_pretrained(
149
+ "facebook/wav2vec2-xlsr-53-espeak-cv-ft"
150
+ ).to(device)
151
+ self.ctc_model.eval()
152
+ for p in self.ctc_model.parameters():
153
+ p.requires_grad = False
154
+
155
+ self.fe_ctc = Wav2Vec2FeatureExtractor.from_pretrained(
156
+ "facebook/wav2vec2-xlsr-53-espeak-cv-ft"
157
+ )
158
+
159
+ vocab_path = hf_hub_download("facebook/wav2vec2-xlsr-53-espeak-cv-ft", "vocab.json")
160
+ with open(vocab_path) as f:
161
+ self.vocab = json.load(f)
162
+ self.blank_idx = self.vocab.get("<pad>", 0)
163
+
164
+ def text_to_arpabet(self, text):
165
+ phones = self.g2p(text)
166
+ result = []
167
+ for ph in phones:
168
+ if ph == " ":
169
+ continue
170
+ clean = re.sub(r"\d", "", ph).lower()
171
+ if clean:
172
+ result.append(clean)
173
+ return result
174
+
175
+ def arpabet_to_idx(self, ph):
176
+ for ipa in ARPABET_TO_IPA.get(ph, []):
177
+ if ipa in self.vocab:
178
+ return self.vocab[ipa]
179
+ return self.vocab.get(ph, -1)
180
+
181
+ def viterbi_align(self, emissions, phone_indices):
182
+ T, C = emissions.shape
183
+ S = len(phone_indices)
184
+ if S == 0 or T < S:
185
+ return []
186
+ extended = [self.blank_idx]
187
+ for p in phone_indices:
188
+ extended.append(p)
189
+ extended.append(self.blank_idx)
190
+ S_ext = len(extended)
191
+ dp = np.full((T, S_ext), float("-inf"), dtype=np.float64)
192
+ bp = np.zeros((T, S_ext), dtype=np.int32)
193
+ dp[0][0] = emissions[0, extended[0]].item()
194
+ if S_ext > 1:
195
+ dp[0][1] = emissions[0, extended[1]].item()
196
+ for t in range(1, T):
197
+ for s in range(S_ext):
198
+ emit = emissions[t, extended[s]].item()
199
+ best, best_s = dp[t-1][s], s
200
+ if s > 0 and dp[t-1][s-1] > best:
201
+ best, best_s = dp[t-1][s-1], s-1
202
+ if s > 1 and extended[s] != self.blank_idx and extended[s] != extended[s-2]:
203
+ if dp[t-1][s-2] > best:
204
+ best, best_s = dp[t-1][s-2], s-2
205
+ dp[t][s] = best + emit
206
+ bp[t][s] = best_s
207
+ 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)
208
+ path = []
209
+ for t in range(T-1, -1, -1):
210
+ path.append((t, extended[s]))
211
+ s = bp[t][s]
212
+ path.reverse()
213
+ return path
214
+
215
+ @torch.no_grad()
216
+ def align(self, audio_path, text):
217
+ """Return alignment info: list of (arpabet, ctc_frames, phone_id, gop)."""
218
+ waveform, sr = torchaudio.load(audio_path)
219
+ if waveform.shape[0] > 1:
220
+ waveform = waveform.mean(0, keepdim=True)
221
+ if sr != SAMPLE_RATE:
222
+ waveform = F_audio.resample(waveform, sr, SAMPLE_RATE)
223
+
224
+ wav_np = waveform.squeeze(0).numpy()
225
+
226
+ # CTC emissions
227
+ ctc_inputs = self.fe_ctc(wav_np, sampling_rate=SAMPLE_RATE, return_tensors="pt", padding=True)
228
+ ctc_logits = self.ctc_model(ctc_inputs.input_values.to(self.device)).logits
229
+ emissions = torch.log_softmax(ctc_logits, dim=-1).squeeze(0).cpu()
230
+
231
+ # G2P + alignment
232
+ arpabet = self.text_to_arpabet(text)
233
+ indices, valid = [], []
234
+ for ph in arpabet:
235
+ idx = self.arpabet_to_idx(ph)
236
+ if idx >= 0:
237
+ indices.append(idx)
238
+ valid.append(ph)
239
+ if not indices:
240
+ return [], emissions
241
+
242
+ path = self.viterbi_align(emissions, indices)
243
+ if not path:
244
+ return [], emissions
245
+
246
+ # Group frames by phone
247
+ segments = []
248
+ cur_tok, cur_frames = None, []
249
+ for f, tok in path:
250
+ if tok == self.blank_idx:
251
+ if cur_tok is not None:
252
+ segments.append((cur_tok, cur_frames))
253
+ cur_tok, cur_frames = None, []
254
+ continue
255
+ if tok != cur_tok:
256
+ if cur_tok is not None:
257
+ segments.append((cur_tok, cur_frames))
258
+ cur_tok, cur_frames = tok, [f]
259
+ else:
260
+ cur_frames.append(f)
261
+ if cur_tok is not None:
262
+ segments.append((cur_tok, cur_frames))
263
+
264
+ results = []
265
+ for i, ph in enumerate(valid):
266
+ if i >= len(segments):
267
+ break
268
+ _, frames = segments[i]
269
+
270
+ # GOP
271
+ expected_idx = indices[i]
272
+ seg_emissions = emissions[frames]
273
+ target_lp = seg_emissions[:, expected_idx].mean().item()
274
+ mask = torch.ones(emissions.shape[1], dtype=torch.bool)
275
+ mask[self.blank_idx] = False
276
+ mask[expected_idx] = False
277
+ best_other = seg_emissions[:, mask].max(dim=-1).values.mean().item()
278
+ gop = target_lp - best_other
279
+
280
+ results.append({
281
+ "arpabet": ph,
282
+ "phone_id": PHONE_TO_ID.get(ph, 0),
283
+ "frames": frames,
284
+ "gop": gop,
285
+ "n_frames": len(frames),
286
+ })
287
+
288
+ return results, emissions
289
+
290
+
291
+ # ============================================================
292
+ # Dataset: stores pre-aligned info, extracts backbone features on-the-fly
293
+ # ============================================================
294
+ class PhonemeAlignedDataset(Dataset):
295
+ """
296
+ Each item is one phoneme with its alignment info.
297
+ During __getitem__, we return pre-computed data.
298
+ Backbone features are extracted in batch during training loop.
299
+ """
300
+ def __init__(self, samples, augment=False):
301
+ """
302
+ samples: list of dicts with keys:
303
+ audio_path, phone_id, gop, n_frames, ctc_frames,
304
+ y_score, y_pherr, T_ctc (CTC time steps for this audio)
305
+ """
306
+ self.samples = samples
307
+ self.augment = augment
308
+
309
+ def __len__(self):
310
+ return len(self.samples)
311
+
312
+ def __getitem__(self, idx):
313
+ s = self.samples[idx]
314
+ return {
315
+ "audio_path": s["audio_path"],
316
+ "phone_id": s["phone_id"],
317
+ "gop": s["gop"],
318
+ "n_frames": s["n_frames"],
319
+ "ctc_frames": s["ctc_frames"],
320
+ "T_ctc": s["T_ctc"],
321
+ "y_score": s["y_score"],
322
+ "y_pherr": s["y_pherr"],
323
+ }
324
+
325
+
326
+ # ============================================================
327
+ # Pre-align all data and cache alignment info
328
+ # ============================================================
329
+ def load_gt(xlsx_path):
330
+ wb = openpyxl.load_workbook(xlsx_path, read_only=True)
331
+ ws = wb.active
332
+ records = []
333
+ for row in ws.iter_rows(min_row=2, values_only=True):
334
+ fname, content, raw = row
335
+ if not fname or not content or not raw:
336
+ continue
337
+ s = raw
338
+ for _ in range(5):
339
+ s = s.replace("\\\\", "\\")
340
+ s = s.replace('\\"', '"')
341
+ phones = re.findall(
342
+ r'"char":"([^"]+)","ph2alpha":"([^"]*)".*?"pherr":(\d+).*?"score":([\d.]+)', s
343
+ )
344
+ gt = [{"phone": p.lower(), "pherr": int(e), "score": float(sc)} for p, _, e, sc in phones]
345
+ records.append({"file_name": fname, "content": content, "gt_phones": gt})
346
+ wb.close()
347
+ return records
348
+
349
+
350
+ def pre_align_all(records, audio_dir, aligner, cache_path=None):
351
+ """Pre-compute alignment for all samples. Returns list of per-phoneme dicts."""
352
+ if cache_path and Path(cache_path).exists():
353
+ print(f"Loading cached alignments from {cache_path}")
354
+ return torch.load(cache_path, weights_only=False)
355
+
356
+ all_samples = []
357
+ n_skip = 0
358
+
359
+ for i, rec in enumerate(records):
360
+ audio_path = audio_dir / rec["file_name"]
361
+ if not audio_path.exists():
362
+ n_skip += 1
363
+ continue
364
+ try:
365
+ aligned, emissions = aligner.align(str(audio_path), rec["content"])
366
+ gt = rec["gt_phones"]
367
+ n = min(len(aligned), len(gt))
368
+ T_ctc = emissions.shape[0]
369
+
370
+ for j in range(n):
371
+ all_samples.append({
372
+ "audio_path": str(audio_path),
373
+ "phone_id": aligned[j]["phone_id"],
374
+ "gop": aligned[j]["gop"],
375
+ "n_frames": aligned[j]["n_frames"],
376
+ "ctc_frames": aligned[j]["frames"],
377
+ "T_ctc": T_ctc,
378
+ "y_score": gt[j]["score"],
379
+ "y_pherr": float(gt[j]["pherr"]),
380
+ })
381
+ except Exception as e:
382
+ n_skip += 1
383
+ if (i+1) % 200 == 0:
384
+ print(f" [{i+1}/{len(records)}] phonemes={len(all_samples)}, skip={n_skip}")
385
+
386
+ print(f" Total: {len(all_samples)} phonemes from {len(records)-n_skip} files, skip={n_skip}")
387
+
388
+ if cache_path:
389
+ torch.save(all_samples, cache_path)
390
+ print(f" Cached to {cache_path}")
391
+
392
+ return all_samples
393
+
394
+
395
+ # ============================================================
396
+ # Training with backbone fine-tuning
397
+ # ============================================================
398
+ def extract_phoneme_features(model, audio_paths, frame_lists, T_ctcs, device, fe):
399
+ """
400
+ Extract backbone hidden states for a batch of phonemes.
401
+ Groups phonemes by audio file to avoid redundant forward passes.
402
+ Returns: (B, 1024) tensor of per-phoneme features.
403
+ """
404
+ # Group by audio path
405
+ path_to_indices = {}
406
+ for i, path in enumerate(audio_paths):
407
+ if path not in path_to_indices:
408
+ path_to_indices[path] = []
409
+ path_to_indices[path].append(i)
410
+
411
+ features = torch.zeros(len(audio_paths), 1024, device=device)
412
+
413
+ for path, indices in path_to_indices.items():
414
+ # Load audio once
415
+ waveform, sr = torchaudio.load(path)
416
+ if waveform.shape[0] > 1:
417
+ waveform = waveform.mean(0, keepdim=True)
418
+ if sr != SAMPLE_RATE:
419
+ waveform = F_audio.resample(waveform, sr, SAMPLE_RATE)
420
+
421
+ wav_np = waveform.squeeze(0).numpy()
422
+ inputs = fe(wav_np, sampling_rate=SAMPLE_RATE, return_tensors="pt", padding=True)
423
+ hidden = model.extract_hidden(inputs.input_values.to(device)) # (1, T_h, 1024)
424
+ hidden = hidden.squeeze(0) # (T_h, 1024)
425
+ T_h = hidden.shape[0]
426
+
427
+ for idx in indices:
428
+ T_ctc = T_ctcs[idx]
429
+ scale = T_h / max(T_ctc, 1)
430
+ frames = frame_lists[idx]
431
+ h_start = max(0, int(min(frames) * scale))
432
+ h_end = min(T_h, int((max(frames) + 1) * scale))
433
+ if h_end <= h_start:
434
+ h_end = h_start + 1
435
+ features[idx] = hidden[h_start:h_end].mean(dim=0)
436
+
437
+ return features
438
+
439
+
440
+ def collate_fn(batch):
441
+ """Custom collate that handles variable-length frame lists."""
442
+ return {
443
+ "audio_paths": [b["audio_path"] for b in batch],
444
+ "phone_id": torch.tensor([b["phone_id"] for b in batch], dtype=torch.long),
445
+ "gop": torch.tensor([b["gop"] for b in batch], dtype=torch.float32),
446
+ "n_frames": torch.tensor([b["n_frames"] for b in batch], dtype=torch.float32),
447
+ "ctc_frames": [b["ctc_frames"] for b in batch],
448
+ "T_ctc": [b["T_ctc"] for b in batch],
449
+ "y_score": torch.tensor([b["y_score"] for b in batch], dtype=torch.float32),
450
+ "y_pherr": torch.tensor([b["y_pherr"] for b in batch], dtype=torch.float32),
451
+ }
452
+
453
+
454
+ def train(all_samples, device, epochs=30, batch_size=64, lr_backbone=1e-5, lr_head=5e-4,
455
+ unfreeze_top_n=6, grad_accum=4):
456
+ """
457
+ Fine-tune backbone + train head jointly.
458
+ """
459
+ # Split by audio file to avoid data leakage
460
+ audio_files = list(set(s["audio_path"] for s in all_samples))
461
+ files_train, files_test = train_test_split(audio_files, test_size=0.15, random_state=42)
462
+ files_train, files_val = train_test_split(files_train, test_size=0.12, random_state=42)
463
+
464
+ train_set = set(files_train)
465
+ val_set = set(files_val)
466
+ test_set = set(files_test)
467
+
468
+ train_samples = [s for s in all_samples if s["audio_path"] in train_set]
469
+ val_samples = [s for s in all_samples if s["audio_path"] in val_set]
470
+ test_samples = [s for s in all_samples if s["audio_path"] in test_set]
471
+
472
+ print(f"\nSplit by audio file:")
473
+ print(f" Train: {len(train_samples)} phonemes from {len(files_train)} files")
474
+ print(f" Val: {len(val_samples)} phonemes from {len(files_val)} files")
475
+ print(f" Test: {len(test_samples)} phonemes from {len(files_test)} files")
476
+
477
+ train_ds = PhonemeAlignedDataset(train_samples, augment=True)
478
+ val_ds = PhonemeAlignedDataset(val_samples)
479
+ test_ds = PhonemeAlignedDataset(test_samples)
480
+
481
+ train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True,
482
+ collate_fn=collate_fn, num_workers=0)
483
+ val_loader = DataLoader(val_ds, batch_size=batch_size, collate_fn=collate_fn, num_workers=0)
484
+
485
+ # Model
486
+ model = BackbonePhoneScorer(unfreeze_top_n=unfreeze_top_n).to(device)
487
+
488
+ # Count parameters
489
+ total_params = sum(p.numel() for p in model.parameters())
490
+ trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
491
+ print(f"\n Total params: {total_params/1e6:.1f}M")
492
+ print(f" Trainable params: {trainable_params/1e6:.1f}M")
493
+
494
+ # Differential learning rate
495
+ optimizer = torch.optim.AdamW([
496
+ {"params": model.backbone_params(), "lr": lr_backbone},
497
+ {"params": model.head_params(), "lr": lr_head},
498
+ ], weight_decay=1e-3)
499
+
500
+ scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
501
+
502
+ # Pos weight for imbalanced pherr
503
+ pherr_vals = torch.tensor([s["y_pherr"] for s in train_samples])
504
+ pos_rate = pherr_vals.mean().item()
505
+ pos_weight = torch.tensor([(1 - pos_rate) / max(pos_rate, 0.01)]).to(device)
506
+ print(f" Pherr pos rate: {pos_rate:.3f}, pos_weight: {pos_weight.item():.2f}")
507
+
508
+ best_val_loss = float("inf")
509
+ best_state = None
510
+ patience, patience_counter = 8, 0
511
+
512
+ fe = model.fe_backbone
513
+
514
+ n_steps_per_epoch = len(train_loader)
515
+ print(f" Steps per epoch: {n_steps_per_epoch}")
516
+
517
+ for epoch in range(epochs):
518
+ model.train()
519
+ losses = []
520
+
521
+ optimizer.zero_grad()
522
+ for step, batch in enumerate(train_loader):
523
+ # Extract backbone features (with gradient for unfrozen layers)
524
+ h = extract_phoneme_features(
525
+ model, batch["audio_paths"], batch["ctc_frames"],
526
+ batch["T_ctc"], device, fe
527
+ )
528
+
529
+ # Add noise augmentation
530
+ h = h + torch.randn_like(h) * 0.01
531
+
532
+ pid = batch["phone_id"].to(device)
533
+ gop = batch["gop"].to(device)
534
+ nf = batch["n_frames"].to(device)
535
+ ys = batch["y_score"].to(device)
536
+ yp = batch["y_pherr"].to(device)
537
+
538
+ pred_s, pred_p = model(h, pid, gop, nf)
539
+ loss_s = F.mse_loss(pred_s, ys)
540
+ loss_p = F.binary_cross_entropy_with_logits(pred_p, yp, pos_weight=pos_weight)
541
+ loss = (loss_s / 100.0 + loss_p) / grad_accum
542
+
543
+ loss.backward()
544
+
545
+ if (step + 1) % grad_accum == 0 or (step + 1) == len(train_loader):
546
+ torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
547
+ optimizer.step()
548
+ optimizer.zero_grad()
549
+
550
+ losses.append(loss.item() * grad_accum)
551
+
552
+ if (step + 1) % 50 == 0:
553
+ avg_loss = np.mean(losses[-50:])
554
+ print(f" Epoch {epoch+1} Step {step+1}/{n_steps_per_epoch} loss={avg_loss:.4f}")
555
+
556
+ scheduler.step()
557
+
558
+ # Validate
559
+ model.eval()
560
+ val_losses = []
561
+ with torch.no_grad():
562
+ for batch in val_loader:
563
+ h = extract_phoneme_features(
564
+ model, batch["audio_paths"], batch["ctc_frames"],
565
+ batch["T_ctc"], device, fe
566
+ )
567
+ pid = batch["phone_id"].to(device)
568
+ gop = batch["gop"].to(device)
569
+ nf = batch["n_frames"].to(device)
570
+ ys = batch["y_score"].to(device)
571
+ yp = batch["y_pherr"].to(device)
572
+
573
+ pred_s, pred_p = model(h, pid, gop, nf)
574
+ loss = F.mse_loss(pred_s, ys) / 100.0 + \
575
+ F.binary_cross_entropy_with_logits(pred_p, yp, pos_weight=pos_weight)
576
+ val_losses.append(loss.item())
577
+
578
+ vl = np.mean(val_losses)
579
+ tl = np.mean(losses)
580
+
581
+ if vl < best_val_loss:
582
+ best_val_loss = vl
583
+ best_state = {k: v.cpu().clone() for k, v in model.state_dict().items()}
584
+ patience_counter = 0
585
+ marker = " *"
586
+ else:
587
+ patience_counter += 1
588
+ marker = ""
589
+
590
+ print(f" Epoch {epoch+1:3d}/{epochs} train={tl:.4f} val={vl:.4f} "
591
+ f"best={best_val_loss:.4f} patience={patience_counter}/{patience}{marker}")
592
+
593
+ if patience_counter >= patience:
594
+ print(f" Early stopping at epoch {epoch+1}")
595
+ break
596
+
597
+ # Load best model
598
+ model.load_state_dict(best_state)
599
+ model = model.to(device)
600
+
601
+ # ============================================================
602
+ # Test evaluation
603
+ # ============================================================
604
+ print(f"\n{'='*60}")
605
+ print("TEST EVALUATION")
606
+ print(f"{'='*60}")
607
+
608
+ model.eval()
609
+ all_pred_s, all_pred_p, all_gt_s, all_gt_p = [], [], [], []
610
+
611
+ test_loader = DataLoader(test_ds, batch_size=batch_size, collate_fn=collate_fn, num_workers=0)
612
+ with torch.no_grad():
613
+ for batch in test_loader:
614
+ h = extract_phoneme_features(
615
+ model, batch["audio_paths"], batch["ctc_frames"],
616
+ batch["T_ctc"], device, fe
617
+ )
618
+ pid = batch["phone_id"].to(device)
619
+ gop = batch["gop"].to(device)
620
+ nf = batch["n_frames"].to(device)
621
+
622
+ pred_s, pred_p = model(h, pid, gop, nf)
623
+ all_pred_s.append(pred_s.cpu())
624
+ all_pred_p.append(torch.sigmoid(pred_p).cpu())
625
+ all_gt_s.append(batch["y_score"])
626
+ all_gt_p.append(batch["y_pherr"])
627
+
628
+ pred_s = torch.cat(all_pred_s).numpy()
629
+ pred_p = torch.cat(all_pred_p).numpy()
630
+ gt_s = torch.cat(all_gt_s).numpy()
631
+ gt_p = torch.cat(all_gt_p).numpy()
632
+
633
+ pred_s_clip = np.clip(pred_s, 0, 100)
634
+
635
+ # Score metrics
636
+ mae = np.mean(np.abs(gt_s - pred_s_clip))
637
+ rmse = np.sqrt(np.mean((gt_s - pred_s_clip) ** 2))
638
+ corr, _ = stats.pearsonr(gt_s, pred_s_clip)
639
+ sp, _ = stats.spearmanr(gt_s, pred_s_clip)
640
+
641
+ print(f"\nPHONE SCORE PREDICTION")
642
+ print(f" MAE: {mae:.2f}")
643
+ print(f" RMSE: {rmse:.2f}")
644
+ print(f" Pearson: {corr:.3f}")
645
+ print(f" Spearman: {sp:.3f}")
646
+ abs_err = np.abs(gt_s - pred_s_clip)
647
+ for t in [5, 10, 15, 20, 30]:
648
+ print(f" |err|<{t:2d}: {np.mean(abs_err<t)*100:.1f}%")
649
+
650
+ # Pherr metrics
651
+ auc = roc_auc_score(gt_p, pred_p) if len(np.unique(gt_p)) > 1 else 0
652
+ best_f1, best_thresh = 0, 0.5
653
+ for th in np.arange(0.1, 0.9, 0.05):
654
+ pb = (pred_p >= th).astype(int)
655
+ _, _, f1, _ = precision_recall_fscore_support(gt_p, pb, average="binary", zero_division=0)
656
+ if f1 > best_f1:
657
+ best_f1, best_thresh = f1, th
658
+
659
+ pb = (pred_p >= best_thresh).astype(int)
660
+ prec, rec, f1, _ = precision_recall_fscore_support(gt_p, pb, average="binary")
661
+ tp = ((pb==1) & (gt_p==1)).sum()
662
+ fp = ((pb==1) & (gt_p==0)).sum()
663
+ fn = ((pb==0) & (gt_p==1)).sum()
664
+ tn = ((pb==0) & (gt_p==0)).sum()
665
+
666
+ print(f"\nPHONE ERROR DETECTION")
667
+ print(f" AUC-ROC: {auc:.3f}")
668
+ print(f" Threshold: {best_thresh:.2f}")
669
+ print(f" Precision: {prec:.3f}")
670
+ print(f" Recall: {rec:.3f}")
671
+ print(f" F1: {f1:.3f}")
672
+ print(f" TP={tp:4d} FP={fp:4d}")
673
+ print(f" FN={fn:4d} TN={tn:4d}")
674
+
675
+ # Comparison
676
+ print(f"\n{'='*60}")
677
+ print("COMPARISON WITH PREVIOUS METHODS")
678
+ print(f"{'='*60}")
679
+ print(f" {'Method':<35s} {'AUC':>6s} {'F1':>6s} {'Prec':>6s} {'Rec':>6s} {'Pearson':>8s} {'MAE':>6s}")
680
+ methods = [
681
+ ("GOP threshold (v1.0, 1K data)", 0.738, 0.476, 0.379, 0.638, 0.372, 27.44),
682
+ ("E2E MLP frozen (v2, 1K data)", 0.814, 0.565, 0.500, 0.650, 0.528, 22.57),
683
+ ("Phoneme comparison (1K data)", 0.691, 0.492, 0.379, 0.703, None, None),
684
+ (f"WavLM finetune (11K data)", auc, best_f1, prec, rec, corr, mae),
685
+ ]
686
+ for name, a, f, p, r, c, m in methods:
687
+ c_str = f"{c:.3f}" if c is not None else " N/A"
688
+ m_str = f"{m:.2f}" if m is not None else " N/A"
689
+ print(f" {name:<35s} {a:>6.3f} {f:>6.3f} {p:>6.3f} {r:>6.3f} {c_str:>8s} {m_str:>6s}")
690
+
691
+ # Save model
692
+ save_dir = Path("/mnt/weka/home/jianshu.she/Voice-correction")
693
+ save_path = save_dir / "wavlm_finetuned.pt"
694
+ torch.save({
695
+ "model_state": model.state_dict(),
696
+ "unfreeze_top_n": unfreeze_top_n,
697
+ "metrics": {
698
+ "auc": auc, "f1": best_f1, "precision": float(prec),
699
+ "recall": float(rec), "pearson": corr, "mae": mae,
700
+ },
701
+ }, save_path)
702
+ print(f"\nModel saved to {save_path}")
703
+
704
+ return model
705
+
706
+
707
+ def main():
708
+ device = torch.device("cuda:0")
709
+
710
+ # Load both datasets
711
+ print("Loading datasets...")
712
+ records_sent = load_gt("/mnt/weka/home/jianshu.she/Voice-correction/eval_log.xlsx")
713
+ records_word = load_gt("/mnt/weka/home/jianshu.she/Voice-correction/word_eval_log.xlsx")
714
+ print(f" Sentence data: {len(records_sent)} records")
715
+ print(f" Word data: {len(records_word)} records")
716
+
717
+ # Tag audio directories
718
+ for r in records_sent:
719
+ r["audio_dir"] = "audio_files"
720
+ for r in records_word:
721
+ r["audio_dir"] = "words_audio_files"
722
+
723
+ # Pre-align all data
724
+ aligner = AlignmentEngine(device)
725
+
726
+ base_dir = Path("/mnt/weka/home/jianshu.she/Voice-correction")
727
+
728
+ cache_sent = base_dir / "align_cache_sentences.pt"
729
+ cache_word = base_dir / "align_cache_words.pt"
730
+
731
+ print("\nAligning sentence data...")
732
+ samples_sent = pre_align_all(records_sent, base_dir / "audio_files", aligner, cache_sent)
733
+
734
+ print("\nAligning word data...")
735
+ samples_word = pre_align_all(records_word, base_dir / "words_audio_files", aligner, cache_word)
736
+
737
+ all_samples = samples_sent + samples_word
738
+ print(f"\nTotal: {len(all_samples)} phonemes")
739
+
740
+ pherr_rate = np.mean([s["y_pherr"] for s in all_samples])
741
+ print(f"Overall error rate: {pherr_rate:.3f}")
742
+
743
+ # Train
744
+ train(all_samples, device, epochs=30, batch_size=64, unfreeze_top_n=6)
745
+
746
+
747
+ if __name__ == "__main__":
748
+ main()