整理 Q1-Q3 实验代码与结果
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user