88from pathlib import Path
99
1010import meeteval .io
11- from meeteval .wer .wer .error_rate import ErrorRate
11+ from meeteval .wer .wer .error_rate import BaseErrorRate
1212
1313
1414def _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
80100class _FilenameEscaper :
81101 """
0 commit comments