99import json
1010import logging
1111import sys
12-
1312from itertools import groupby
13+
14+ from rapidfuzz .distance import Levenshtein
1415import numpy as np
1516import six
1617
@@ -155,7 +156,6 @@ def calculate_cer_ctc(self, ys_hat, ys_pad):
155156 :return: average sentence-level CER score
156157 :rtype float
157158 """
158- import editdistance
159159
160160 cers , char_ref_lens = [], []
161161 for i , y in enumerate (ys_hat ):
@@ -175,7 +175,7 @@ def calculate_cer_ctc(self, ys_hat, ys_pad):
175175 hyp_chars = "" .join (seq_hat )
176176 ref_chars = "" .join (seq_true )
177177 if len (ref_chars ) > 0 :
178- cers .append (editdistance . eval (hyp_chars , ref_chars ))
178+ cers .append (Levenshtein . distance (hyp_chars , ref_chars ))
179179 char_ref_lens .append (len (ref_chars ))
180180
181181 cer_ctc = float (sum (cers )) / sum (char_ref_lens ) if cers else None
@@ -214,16 +214,16 @@ def calculate_cer(self, seqs_hat, seqs_true):
214214 :return: average sentence-level CER score
215215 :rtype float
216216 """
217- import editdistance
218217
219218 char_eds , char_ref_lens = [], []
220219 for i , seq_hat_text in enumerate (seqs_hat ):
221220 seq_true_text = seqs_true [i ]
222221 hyp_chars = seq_hat_text .replace (" " , "" )
223222 ref_chars = seq_true_text .replace (" " , "" )
224- char_eds .append (editdistance . eval (hyp_chars , ref_chars ))
223+ char_eds .append (Levenshtein . distance (hyp_chars , ref_chars ))
225224 char_ref_lens .append (len (ref_chars ))
226- return float (sum (char_eds )) / sum (char_ref_lens )
225+ ref_len = sum (char_ref_lens )
226+ return float (sum (char_eds )) / ref_len if ref_len > 0 else None
227227
228228 def calculate_wer (self , seqs_hat , seqs_true ):
229229 """Calculate sentence-level WER score.
@@ -233,13 +233,13 @@ def calculate_wer(self, seqs_hat, seqs_true):
233233 :return: average sentence-level WER score
234234 :rtype float
235235 """
236- import editdistance
237236
238237 word_eds , word_ref_lens = [], []
239238 for i , seq_hat_text in enumerate (seqs_hat ):
240239 seq_true_text = seqs_true [i ]
241240 hyp_words = seq_hat_text .split ()
242241 ref_words = seq_true_text .split ()
243- word_eds .append (editdistance . eval (hyp_words , ref_words ))
242+ word_eds .append (Levenshtein . distance (hyp_words , ref_words ))
244243 word_ref_lens .append (len (ref_words ))
245- return float (sum (word_eds )) / sum (word_ref_lens )
244+ ref_len = sum (word_ref_lens )
245+ return float (sum (word_eds )) / ref_len if ref_len > 0 else None
0 commit comments