Skip to content

Commit 54a3b4a

Browse files
committed
Add "type" identifier to error rate serialization
1 parent 8644eab commit 54a3b4a

10 files changed

Lines changed: 332 additions & 123 deletions

File tree

meeteval/der/md_eval.py

Lines changed: 24 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
from pathlib import Path
99

1010
import meeteval.io
11-
from meeteval.wer.wer.error_rate import ErrorRate
11+
from meeteval.wer.wer.error_rate import BaseErrorRate
1212

1313

1414
def _fix_channel(r):
@@ -21,10 +21,12 @@ def _fix_channel(r):
2121

2222

2323
@dataclasses.dataclass(frozen=True)
24-
class DiaErrorRate:
24+
class DiaErrorRate(BaseErrorRate):
2525
"""
2626
2727
"""
28+
identifier = 'diarization-error-rate'
29+
2830
error_rate: 'float | decimal.Decimal'
2931

3032
scored_speaker_time: 'float | decimal.Decimal'
@@ -36,16 +38,29 @@ class DiaErrorRate:
3638
def zero(cls):
3739
return cls(0, 0, 0, 0, 0)
3840

41+
@classmethod
42+
def from_dict(cls, d: dict) -> 'Self':
43+
return cls(
44+
d['error_rate'],
45+
d['scored_speaker_time'],
46+
d['missed_speaker_time'],
47+
d['falarm_speaker_time'],
48+
d['speaker_error_time'],
49+
)
50+
3951
def __post_init__(self):
4052
assert self.scored_speaker_time >= 0
4153
assert self.missed_speaker_time >= 0
4254
assert self.falarm_speaker_time >= 0
4355
assert self.speaker_error_time >= 0
4456
errors = self.speaker_error_time + self.falarm_speaker_time + self.missed_speaker_time
45-
error_rate = errors / self.scored_speaker_time
57+
if self.scored_speaker_time > 0:
58+
error_rate = errors / self.scored_speaker_time
59+
else:
60+
error_rate = None
4661
if self.error_rate is None:
4762
object.__setattr__(self, 'error_rate', error_rate)
48-
else:
63+
elif error_rate is not None:
4964
# Since md-eval uses float internally, and the printed numbers are
5065
# rounded, it is in corner cases not possible to reproduce the
5166
# exact error rate, that is calculated internally by md-eval.
@@ -76,6 +91,11 @@ def __add__(self, other: 'DiaErrorRate'):
7691
speaker_error_time=self.speaker_error_time + other.speaker_error_time,
7792
)
7893

94+
def asdict(self):
95+
d = dataclasses.asdict(self)
96+
d['type'] = self.identifier
97+
return d
98+
7999

80100
class _FilenameEscaper:
81101
"""

meeteval/viz/visualize.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -414,7 +414,7 @@ def compress(m):
414414
# Add utterances to data. Add total number of words to each utterance
415415
data['utterances'] = [{**l, 'total': len(l['words'].split())} for l in u]
416416

417-
data['info']['wer'] = dataclasses.asdict(wer)
417+
data['info']['wer'] = wer.asdict()
418418

419419
def wer_by_speaker(speaker):
420420
# Get all words from this speaker
@@ -434,15 +434,15 @@ def wer_by_speaker(speaker):
434434
deletions = len(ref_words.filter(
435435
lambda s: not [w for w, _ in s['matches'] if w is not None and words[w]['source'] == 'hypothesis']))
436436

437-
return dataclasses.asdict(ErrorRate(
437+
return ErrorRate(
438438
errors=insertions + deletions + substitutions,
439439
length=len(ref_words),
440440
insertions=insertions,
441441
deletions=deletions,
442442
substitutions=substitutions,
443443
reference_self_overlap=None,
444444
hypothesis_self_overlap=None,
445-
))
445+
).asdict()
446446

447447
data['info']['wer_by_speakers'] = {
448448
speaker: wer_by_speaker(speaker)

meeteval/wer/__main__.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
import meeteval.io
1212
from meeteval.io.smart import _open
13-
from meeteval.wer.wer import combine_error_rates, ErrorRate
13+
from meeteval.wer.wer import combine_error_rates, WordErrorRate
1414
import sys
1515
import meeteval.wer
1616

@@ -119,14 +119,14 @@ def to_str(example_id):
119119

120120
# Save details
121121
_dump({
122-
to_str(example_id): dataclasses.asdict(error_rate)
122+
to_str(example_id): error_rate.asdict()
123123
for example_id, error_rate in per_reco.items()
124124
}, per_reco_out.format(parent=parent, stem=stem))
125125

126126
# Compute and save average
127127
average = combine_error_rates(*per_reco.values())
128128
_dump(
129-
dataclasses.asdict(average),
129+
average.asdict(),
130130
average_out.format(parent=parent, stem=stem),
131131
)
132132
if hasattr(average, 'scored_speaker_time'):
@@ -445,20 +445,20 @@ def _merge(
445445
if 'errors' in d: # Average file
446446
assert average is not False, average
447447
average = True # A single average file forces to do an average
448-
ers.append([None, ErrorRate.from_dict(d)])
448+
ers.append([None, WordErrorRate.from_dict(d)])
449449
else:
450450
for k, v in d.items(): # Details file
451451
if regex is not None and not regex.fullmatch(k):
452452
continue
453453
if 'errors' in v:
454-
ers.append([k, ErrorRate.from_dict(v)])
454+
ers.append([k, WordErrorRate.from_dict(v)])
455455

456456
if average:
457457
er = meeteval.wer.combine_error_rates(*[er for _, er in ers])
458-
out_data = dataclasses.asdict(er)
458+
out_data = er.asdict()
459459
else:
460460
out_data = {
461-
k: dataclasses.asdict(er)
461+
k: er.asdict()
462462
for k, er in ers
463463
}
464464
assert len(out_data) == len(ers), (len(out_data), len(ers), 'Duplicate filenames')

meeteval/wer/wer/cp.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,8 @@ class CPErrorRate(ErrorRate):
3636
>>> combine_error_rates(CPErrorRate(0, 10, 0, 0, 0, None, None, 1, 0, 3), CPErrorRate(5, 10, 0, 0, 5, None, None, 0, 1, 3))
3737
CPErrorRate(error_rate=0.25, errors=5, length=20, insertions=0, deletions=0, substitutions=5, missed_speaker=1, falarm_speaker=1, scored_speaker=6)
3838
"""
39+
identifier = 'cp-error-rate'
40+
3941
missed_speaker: int
4042
falarm_speaker: int
4143
scored_speaker: int

meeteval/wer/wer/di_cp.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,8 +19,13 @@
1919

2020
@dataclasses.dataclass(frozen=True)
2121
class DICPErrorRate(ErrorRate):
22+
identifier = 'di-cp-error-rate'
2223
assignment: Tuple[int, ...]
2324

25+
@classmethod
26+
def zero(cls):
27+
return DICPErrorRate(0, 0, 0, 0, 0, None, None, ())
28+
2429
def apply_assignment(self, reference, hypothesis):
2530
return apply_dicp_assignment(self.assignment, reference, hypothesis)
2631

0 commit comments

Comments
 (0)