summaryrefslogtreecommitdiffstats
path: root/Evaluate
diff options
context:
space:
mode:
authoryum <yum.food.vr@gmail.com>2023-03-03 20:10:32 -0800
committeryum <yum.food.vr@gmail.com>2023-03-03 20:42:10 -0800
commita74ee78dbb79c551851dc182090e7d4292b1e80c (patch)
treea18542faddbdb1a8cc6285e39d0bcb3aad47a19d /Evaluate
parentf7d5741e5c069d759f8412bd40b279e1d7abac4c (diff)
Use logprobs, fix beam candidate selection
Incorrect sort condition resulted in worst 5 beams being picked instead of best 5. Use log probabilities for joint probability calculation instead of linear probabilities. Long beams would have probabilities converge exponentially towards zero; now they converge linearly towards -INFINITY. Using both transcripts in Evaluation/setup.ps1, I see a small edit distance regression (~5%) using beam search vs. greedy.
Diffstat (limited to 'Evaluate')
-rw-r--r--Evaluate/.swpbin12288 -> 0 bytes
-rw-r--r--Evaluate/evaluate.py22
2 files changed, 8 insertions, 14 deletions
diff --git a/Evaluate/.swp b/Evaluate/.swp
deleted file mode 100644
index c1bc460..0000000
--- a/Evaluate/.swp
+++ /dev/null
Binary files differ
diff --git a/Evaluate/evaluate.py b/Evaluate/evaluate.py
index 5e8c85d..81b3edf 100644
--- a/Evaluate/evaluate.py
+++ b/Evaluate/evaluate.py
@@ -1,11 +1,12 @@
import argparse
import editdistance
-import jiwer
import re
import subprocess
import sys
import time
+from whisper.normalizers import EnglishTextNormalizer
+
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("reference_path", type=str, help="Path to reference transcript")
@@ -33,22 +34,15 @@ if __name__ == "__main__":
with open(args.reference_path, "r") as f:
ref_transcript = f.read()
+ # Normalize transcripts before computing edit distance (as described in
+ # whisper paper).
+ normalize = EnglishTextNormalizer()
+ test_transcript = normalize(test_transcript)
+ ref_transcript = normalize(ref_transcript)
+
dist = editdistance.eval(ref_transcript, test_transcript)
- wer_transform = jiwer.Compose([
- jiwer.ToLowerCase(),
- jiwer.RemoveWhiteSpace(replace_by_space=True),
- jiwer.RemoveMultipleSpaces(),
- jiwer.RemovePunctuation(),
- jiwer.ReduceToListOfListOfWords(word_delimiter=" "),
- ])
- wer = jiwer.wer(
- ref_transcript,
- test_transcript,
- truth_transform=wer_transform,
- hypothesis_transform=wer_transform)
print(f"Duration: {t1 - t0}")
print(f"Levenshtein distance: {dist}")
- print(f"Word error rate: {wer}")
print(f"Transcript: {test_transcript}")