整理 Q1-Q3 实验代码与结果

This commit is contained in:
2026-09-24 16:25:15 +08:00
parent 0261ecdfba
commit 8f5c2c3be6
247 changed files with 69828 additions and 19 deletions
+148
View File
@@ -0,0 +1,148 @@
from __future__ import annotations
import re
import subprocess
from dataclasses import dataclass
from pathlib import Path
import numpy as np
import torch
from transformers import AutoModelForCTC, AutoTokenizer
SAMPLE_RATE = 16_000
MODEL_ID = "facebook/wav2vec2-base-960h"
@dataclass
class WordInterval:
word: str
start_s: float
end_s: float
quality: float
valid: bool
def decode_audio(video_path: Path) -> np.ndarray:
result = subprocess.run(
[
"ffmpeg", "-v", "error", "-i", str(video_path), "-map", "0:a:0",
"-ac", "1", "-ar", str(SAMPLE_RATE), "-f", "f32le", "pipe:1",
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
waveform = np.frombuffer(result.stdout, dtype="<f4").copy()
if not len(waveform):
raise RuntimeError(f"no decoded audio in {video_path}")
return waveform
def _ctc_targets(words: list[str], tokenizer) -> tuple[list[int], list[list[int]]]:
vocab = tokenizer.get_vocab()
delimiter = int(tokenizer.convert_tokens_to_ids(tokenizer.word_delimiter_token or "|"))
unknown = int(tokenizer.unk_token_id)
targets: list[int] = []
per_word: list[list[int]] = [[] for _ in words]
for word_index, raw_word in enumerate(words):
if word_index:
targets.append(delimiter)
normalized = re.sub(r"[^a-z']", "", raw_word.lower())
for character in normalized:
per_word[word_index].append(len(targets))
targets.append(int(vocab.get(character, unknown)))
return targets, per_word
def _viterbi(log_probs: np.ndarray, targets: list[int], blank: int) -> np.ndarray | None:
if not targets or log_probs.ndim != 2:
return None
states = np.full(2 * len(targets) + 1, blank, dtype=np.int64)
states[1::2] = np.asarray(targets, dtype=np.int64)
frames, count = log_probs.shape[0], len(states)
if frames == 0 or frames < len(targets):
return None
previous = np.full(count, -np.inf, dtype=np.float64)
previous[0] = float(log_probs[0, blank])
previous[1] = float(log_probs[0, states[1]])
back = np.zeros((frames, count), dtype=np.uint8)
skip = np.zeros(count, dtype=bool)
if count > 2:
skip[2:] = (states[2:] != blank) & (states[2:] != states[:-2])
for frame in range(1, frames):
stay = previous
one = np.full(count, -np.inf, dtype=np.float64)
one[1:] = previous[:-1]
two = np.full(count, -np.inf, dtype=np.float64)
if skip.any():
two[skip] = previous[np.flatnonzero(skip) - 2]
candidates = np.stack((stay, one, two), axis=0)
choice = candidates.argmax(axis=0).astype(np.uint8)
previous = candidates[choice, np.arange(count)] + log_probs[frame, states]
back[frame] = choice
state = count - 1 if previous[-1] >= previous[-2] else count - 2
path = np.empty(frames, dtype=np.int32)
path[-1] = state
for frame in range(frames - 1, 0, -1):
state -= int(back[frame, state])
path[frame - 1] = state
return path
def model_time_constants(model) -> tuple[float, float]:
config = model.config
stride = int(np.prod(config.conv_stride))
receptive = 1
jump = 1
for kernel, local_stride in zip(config.conv_kernel, config.conv_stride):
receptive += (int(kernel) - 1) * jump
jump *= int(local_stride)
return stride / SAMPLE_RATE, receptive / (2 * SAMPLE_RATE)
def align_words(
waveform: np.ndarray,
words: list[str],
tokenizer,
model,
device: torch.device,
) -> list[WordInterval]:
targets, word_targets = _ctc_targets(words, tokenizer)
blank = int(tokenizer.pad_token_id)
if not targets or not len(waveform):
return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words]
with torch.inference_mode():
values = torch.as_tensor(waveform, dtype=torch.float32, device=device).unsqueeze(0)
logits = model(input_values=values).logits[0].float()
log_probs = torch.log_softmax(logits, dim=-1).cpu().numpy()
path = _viterbi(log_probs, targets, blank)
frame_step, center_s = model_time_constants(model)
duration = len(waveform) / SAMPLE_RATE
intervals: list[WordInterval] = []
if path is None:
return [WordInterval(w, float("nan"), float("nan"), 0.0, False) for w in words]
for word, target_indices in zip(words, word_targets):
states = np.asarray([2 * index + 1 for index in target_indices], dtype=np.int32)
frame_indices = np.flatnonzero(np.isin(path, states)) if len(states) else np.empty(0, dtype=np.int64)
if not len(frame_indices):
intervals.append(WordInterval(word, float("nan"), float("nan"), 0.0, False))
continue
first, last = int(frame_indices[0]), int(frame_indices[-1])
start = max(0.0, first * frame_step + center_s - frame_step / 2)
end = min(duration, (last + 1) * frame_step + center_s - frame_step / 2)
char_scores = []
for target_index in target_indices:
selected = np.flatnonzero(path == 2 * target_index + 1)
if len(selected):
char_scores.extend(log_probs[selected, targets[target_index]].tolist())
quality = float(np.exp(np.mean(char_scores))) if char_scores else 0.0
valid = end > start
intervals.append(WordInterval(word, start, end, quality, valid))
return intervals
def load_ctc(device: torch.device):
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForCTC.from_pretrained(MODEL_ID).to(device).eval()
return tokenizer, model