| 985 | |
| 986 | |
| 987 | class DeserializeSpeechTransformer(TransformerMixin): |
| 988 | def __init__(self, non_speech_label: float) -> None: |
| 989 | super(DeserializeSpeechTransformer, self).__init__() |
| 990 | self._non_speech_label: float = non_speech_label |
| 991 | self.deserialized_speech_results_: Optional[np.ndarray] = None |
| 992 | |
| 993 | def fit(self, fname, *_) -> "DeserializeSpeechTransformer": |
| 994 | speech = np.load(fname) |
| 995 | if hasattr(speech, "files"): |
| 996 | if "speech" in speech.files: |
| 997 | speech = speech["speech"] |
| 998 | else: |
| 999 | raise ValueError( |
| 1000 | 'could not find "speech" array in ' |
| 1001 | "serialized file; only contains: %s" % speech.files |
| 1002 | ) |
| 1003 | speech[speech < 1.0] = self._non_speech_label |
| 1004 | self.deserialized_speech_results_ = speech |
| 1005 | return self |
| 1006 | |
| 1007 | def transform(self, *_) -> np.ndarray: |
| 1008 | assert self.deserialized_speech_results_ is not None |
| 1009 | return self.deserialized_speech_results_ |
| 1010 | |
| 1011 | |
| 1012 | def find_pgs_stream( |
no outgoing calls
no test coverage detected