taters 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (44) hide show
  1. taters/Taters.py +115 -0
  2. taters/__init__.py +0 -0
  3. taters/audio/__init__.py +0 -0
  4. taters/audio/convert_to_wav.py +112 -0
  5. taters/audio/diarize_with_thirdparty.py +5 -0
  6. taters/audio/diarizer/__init__.py +0 -0
  7. taters/audio/diarizer/whisper-diarization/__init__.py +0 -0
  8. taters/audio/diarizer/whisper-diarization/diarization/__init__.py +3 -0
  9. taters/audio/diarizer/whisper-diarization/diarization/msdd/msdd.py +100 -0
  10. taters/audio/diarizer/whisper-diarization/diarize.py +247 -0
  11. taters/audio/diarizer/whisper-diarization/diarize_custom.py +374 -0
  12. taters/audio/diarizer/whisper-diarization/diarize_parallel.py +269 -0
  13. taters/audio/diarizer/whisper-diarization/helpers.py +552 -0
  14. taters/audio/diarizer/whisper_diar_wrapper.py +293 -0
  15. taters/audio/extract_wav_from_video.py +177 -0
  16. taters/audio/extract_whisper_embeddings.py +243 -0
  17. taters/audio/extract_whisper_embeddings_subproc.py +574 -0
  18. taters/audio/split_wav_by_speaker.py +182 -0
  19. taters/helpers/__init__.py +0 -0
  20. taters/helpers/feature_gather.py +349 -0
  21. taters/helpers/find_files.py +222 -0
  22. taters/helpers/text_gather.py +520 -0
  23. taters/pipelines/__init__.py +0 -0
  24. taters/pipelines/run_pipeline.py +384 -0
  25. taters/text/analyze_with_archetypes.py +238 -0
  26. taters/text/analyze_with_dictionaries.py +246 -0
  27. taters/text/dictionary_analyzers/__init__.py +0 -0
  28. taters/text/dictionary_analyzers/multi_archetype_analyzer.py +125 -0
  29. taters/text/dictionary_analyzers/multi_dict_analyzer.py +254 -0
  30. taters/text/extract_sentence_embeddings.py +329 -0
  31. taters/video/__init__.py +0 -0
  32. taters/video/extract_features.py +0 -0
  33. taters/video/features_basic.py +0 -0
  34. taters/video/features_motion.py +0 -0
  35. taters/video/features_ocr.py +0 -0
  36. taters/video/features_shots.py +0 -0
  37. taters/video/read_video.py +0 -0
  38. taters/video/windowing.py +0 -0
  39. taters-0.1.0.dist-info/METADATA +161 -0
  40. taters-0.1.0.dist-info/RECORD +44 -0
  41. taters-0.1.0.dist-info/WHEEL +5 -0
  42. taters-0.1.0.dist-info/entry_points.txt +2 -0
  43. taters-0.1.0.dist-info/licenses/LICENSE +21 -0
  44. taters-0.1.0.dist-info/top_level.txt +1 -0
taters/Taters.py ADDED
@@ -0,0 +1,115 @@
1
+ # taters/Taters.py
2
+ from __future__ import annotations
3
+ from pathlib import Path
4
+ from typing import Any
5
+ import inspect
6
+
7
+ def _abs(p: str | Path) -> str:
8
+ return str(Path(p).resolve())
9
+
10
+ def _forward(func, kwargs: dict[str, Any]):
11
+ sig = inspect.signature(func)
12
+ try:
13
+ # validate names & requireds; don't execute defaults here
14
+ sig.bind_partial(**kwargs)
15
+ except TypeError as e:
16
+ allowed = ", ".join([str(p) for p in sig.parameters.values()])
17
+ raise TypeError(f"{func.__module__}.{func.__name__}: {e}\nAllowed params: {allowed}")
18
+ return func(**kwargs)
19
+
20
+
21
+ class Taters:
22
+ def __init__(self):
23
+ self.audio = _AudioAPI(self)
24
+ self.text = _TextAPI(self)
25
+ self.helpers = _HelpersAPI(self)
26
+
27
+ # back-compat pass-throughs
28
+
29
+ # audio
30
+ def convert_to_wav(self, **kwargs): return self.audio.convert_to_wav(**kwargs)
31
+ def extract_wavs_from_video(self, **kwargs): return self.audio.extract_wavs_from_video(**kwargs)
32
+ def split_wav_by_speaker(self, **kwargs): return self.audio.split_wav_by_speaker(**kwargs)
33
+ def extract_whisper_embeddings(self, **kwargs): return self.audio.extract_whisper_embeddings(**kwargs)
34
+ def diarize_with_thirdparty(self, **kwargs): return self.audio.diarize_with_thirdparty(**kwargs)
35
+
36
+ #text
37
+ def analyze_with_dictionaries(self, **kwargs): return self.text.analyze_with_dictionaries(**kwargs)
38
+ def analyze_with_archetypes(self, **kwargs): return self.text.analyze_with_archetypes(**kwargs)
39
+ def extract_sentence_embeddings(self, **kwargs): return self.text.extract_sentence_embeddings(**kwargs)
40
+
41
+ # helpers
42
+ def txt_folder_to_analysis_ready_csv(self, **kwargs): return self.helpers.txt_folder_to_analysis_ready_csv(**kwargs)
43
+ def csv_to_analysis_ready_csv(self, **kwargs): return self.helpers.csv_to_analysis_ready_csv(**kwargs)
44
+ def find_files(self, **kwargs): return self.helpers.find_files(**kwargs)
45
+ def feature_gather(self, **kwargs): return self.helpers.feature_gather(**kwargs)
46
+
47
+
48
+ def txt_folder_to_analysis_ready_csv(self, **kwargs):
49
+ from .helpers.text_gather import txt_folder_to_analysis_ready_csv
50
+ return _forward(txt_folder_to_analysis_ready_csv, kwargs)
51
+
52
+ def csv_to_analysis_ready_csv(self, **kwargs):
53
+ from .helpers.text_gather import csv_to_analysis_ready_csv
54
+ return _forward(csv_to_analysis_ready_csv, kwargs)
55
+
56
+
57
+ class _AudioAPI:
58
+ def __init__(self, parent: Taters): self._cs = parent
59
+
60
+ def convert_to_wav(self, **kwargs):
61
+ from .audio.convert_to_wav import convert_audio_to_wav
62
+ return _forward(convert_audio_to_wav, kwargs)
63
+
64
+ def extract_wavs_from_video(self, **kwargs):
65
+ from .audio.extract_wav_from_video import split_audio_streams_to_wav
66
+ return _forward(split_audio_streams_to_wav, kwargs)
67
+
68
+ def split_wav_by_speaker(self, **kwargs):
69
+ from .audio.split_wav_by_speaker import make_speaker_wavs_from_csv
70
+ return _forward(make_speaker_wavs_from_csv, kwargs)
71
+
72
+ def extract_whisper_embeddings(self, **kwargs):
73
+ from .audio.extract_whisper_embeddings import extract_whisper_embeddings
74
+ return _forward(extract_whisper_embeddings, kwargs)
75
+
76
+ def diarize_with_thirdparty(self, **kwargs):
77
+ from .audio.diarizer.whisper_diar_wrapper import run_whisper_diarization_repo
78
+ return _forward(run_whisper_diarization_repo, kwargs)
79
+
80
+
81
+ class _TextAPI:
82
+ def __init__(self, parent: Taters): self._cs = parent
83
+
84
+ def analyze_with_dictionaries(self, **kwargs):
85
+ from .text.analyze_with_dictionaries import analyze_with_dictionaries
86
+ return _forward(analyze_with_dictionaries, kwargs)
87
+
88
+ def analyze_with_archetypes(self, **kwargs):
89
+ from .text.analyze_with_archetypes import analyze_with_archetypes
90
+ return _forward(analyze_with_archetypes, kwargs)
91
+
92
+ def extract_sentence_embeddings(self, **kwargs):
93
+ from .text.extract_sentence_embeddings import analyze_with_sentence_embeddings
94
+ return _forward(analyze_with_sentence_embeddings, kwargs)
95
+
96
+
97
+ class _HelpersAPI:
98
+ def __init__(self, parent: Taters): self._cs = parent
99
+
100
+ def txt_folder_to_analysis_ready_csv(self, **kwargs):
101
+ from .helpers.text_gather import txt_folder_to_analysis_ready_csv
102
+ return _forward(txt_folder_to_analysis_ready_csv, kwargs)
103
+
104
+ def csv_to_analysis_ready_csv(self, **kwargs):
105
+ from .helpers.text_gather import csv_to_analysis_ready_csv
106
+ return _forward(csv_to_analysis_ready_csv, kwargs)
107
+
108
+ def find_files(self, **kwargs):
109
+ from .helpers.find_files import find_files
110
+ return _forward(find_files, kwargs)
111
+
112
+ def feature_gather(self, **kwargs):
113
+ from .helpers import feature_gather
114
+ return _forward(feature_gather, kwargs)
115
+
taters/__init__.py ADDED
File without changes
File without changes
@@ -0,0 +1,112 @@
1
+ # taters/audio/convert_to_wav.py
2
+ from __future__ import annotations
3
+ import shutil
4
+ import subprocess
5
+ from pathlib import Path
6
+ from typing import Optional, Union
7
+
8
+ class FFmpegNotFoundError(RuntimeError):
9
+ pass
10
+
11
+ def _check_ffmpeg():
12
+ if shutil.which("ffmpeg") is None or shutil.which("ffprobe") is None:
13
+ raise FFmpegNotFoundError("ffmpeg and/or ffprobe not found on PATH.")
14
+
15
+ def convert_audio_to_wav(
16
+ input_path: Union[str, Path],
17
+ *,
18
+ output_path: Optional[Union[str, Path]] = None,
19
+ output_dir: Optional[Union[str, Path]] = None,
20
+ sample_rate: int = 16000, # common for ASR
21
+ bit_depth: int = 16, # 16/24/32 signed PCM
22
+ channels: int = 1, # 1=mono, 2=stereo
23
+ overwrite_existing: bool = False, # if the file already exists, let's not overwrite by default
24
+ ) -> Path:
25
+ """
26
+ Convert any audio (or A/V container) to a PCM WAV file using ffmpeg.
27
+
28
+ If output_path is None and output_dir is None, writes <input_stem>.wav next to input.
29
+ If output_dir is given (and output_path is None), writes <output_dir>/<input_stem>.wav.
30
+ If output_path is given, it takes precedence.
31
+
32
+ Returns the Path to the created WAV.
33
+ """
34
+ _check_ffmpeg()
35
+
36
+ in_path = Path(input_path).resolve()
37
+ if not in_path.exists():
38
+ raise FileNotFoundError(f"Input file not found: {in_path}")
39
+
40
+ if output_path and output_dir:
41
+ raise ValueError("Provide at most one of output_path or output_dir.")
42
+
43
+ if output_path:
44
+ out_path = Path(output_path).resolve()
45
+ else:
46
+ base = in_path.stem + ".wav"
47
+ out_dir = Path(output_dir).resolve() if output_dir else Path.cwd() / "audio"
48
+ out_dir.mkdir(parents=True, exist_ok=True)
49
+ out_path = out_dir / base
50
+
51
+ if not overwrite_existing and Path(out_path).is_file():
52
+ print("WAV file already exists; returning existing file.")
53
+ return out_path
54
+
55
+ pcm_map = {16: "pcm_s16le", 24: "pcm_s24le", 32: "pcm_s32le"}
56
+ if bit_depth not in pcm_map:
57
+ raise ValueError("bit_depth must be one of {16, 24, 32}.")
58
+ if channels not in (1, 2):
59
+ raise ValueError("channels must be 1 (mono) or 2 (stereo).")
60
+
61
+ cmd = [
62
+ "ffmpeg",
63
+ "-nostdin",
64
+ "-hide_banner", "-loglevel", "error",
65
+ "-y" if overwrite_existing else "-n",
66
+ "-i", str(in_path),
67
+ "-vn", # ignore video
68
+ "-acodec", pcm_map[bit_depth],
69
+ "-ar", str(sample_rate),
70
+ "-ac", str(channels),
71
+ str(out_path),
72
+ ]
73
+
74
+ result = subprocess.run(cmd, capture_output=True, text=True, stdin=subprocess.DEVNULL)
75
+ if result.returncode != 0:
76
+ if not overwrite_existing and out_path.exists():
77
+ raise FileExistsError(f"Target exists (use overwrite=True): {out_path}")
78
+ raise RuntimeError(f"ffmpeg failed: {result.stderr.strip()}")
79
+
80
+ return out_path
81
+
82
+
83
+ # --- CLI --------------------------------------------------------------------
84
+ def _build_arg_parser():
85
+ import argparse
86
+ p = argparse.ArgumentParser(description="Convert any audio (or A/V) file to PCM WAV via ffmpeg.")
87
+ p.add_argument("input", help="Input file (audio or video container)")
88
+ p.add_argument("--out", dest="output_path", default=None,
89
+ help="Exact output .wav path (overrides --out-dir)")
90
+ p.add_argument("--out-dir", dest="output_dir", default=None,
91
+ help="Directory for output (filename will be <input_stem>.wav)")
92
+ p.add_argument("--sr", dest="sample_rate", type=int, default=16000, help="Sample rate (Hz)")
93
+ p.add_argument("--bit-depth", type=int, choices=[16, 24, 32], default=16, help="PCM bit depth")
94
+ p.add_argument("--channels", type=int, choices=[1, 2], default=1, help="1=mono, 2=stereo")
95
+ p.add_argument("--overwrite_existing", type=bool, default=False, help="Overwrite existing output")
96
+ return p
97
+
98
+ def main():
99
+ args = _build_arg_parser().parse_args()
100
+ out = convert_audio_to_wav(
101
+ args.input,
102
+ output_path=args.output_path,
103
+ output_dir=args.output_dir,
104
+ sample_rate=args.sample_rate,
105
+ bit_depth=args.bit_depth,
106
+ channels=args.channels,
107
+ overwrite_existing=args.overwrite_existing,
108
+ )
109
+ print(str(out))
110
+
111
+ if __name__ == "__main__":
112
+ main()
@@ -0,0 +1,5 @@
1
+ # thin alias (args still pass through)
2
+ from .diarizer.whisper_diar_wrapper import main as main
3
+
4
+ if __name__ == "__main__":
5
+ main()
File without changes
File without changes
@@ -0,0 +1,3 @@
1
+ from .msdd.msdd import MSDDDiarizer
2
+
3
+ __all__ = ["MSDDDiarizer"]
@@ -0,0 +1,100 @@
1
+ import json
2
+ import os
3
+ import tempfile
4
+
5
+ from typing import Union
6
+
7
+ import torch
8
+ import torchaudio
9
+
10
+ from nemo.collections.asr.models.msdd_models import NeuralDiarizer
11
+ from nemo.collections.asr.parts.utils.speaker_utils import rttm_to_labels
12
+ from omegaconf import OmegaConf
13
+
14
+
15
+ class MSDDDiarizer:
16
+ def __init__(self, device: Union[str, torch.device]):
17
+ self.model: NeuralDiarizer = NeuralDiarizer(cfg=create_config()).to(device)
18
+
19
+ def diarize(self, audio: torch.Tensor):
20
+ with tempfile.TemporaryDirectory() as temp_path:
21
+ torchaudio.save(
22
+ os.path.join(temp_path, "mono_file.wav"),
23
+ audio,
24
+ 16000,
25
+ channels_first=True,
26
+ )
27
+
28
+ manifest_path = os.path.join(temp_path, "manifest.json")
29
+ meta = {
30
+ "audio_filepath": os.path.join(temp_path, "mono_file.wav"),
31
+ "offset": 0,
32
+ "duration": None,
33
+ "label": "infer",
34
+ "text": "-",
35
+ "rttm_filepath": None,
36
+ "uem_filepath": None,
37
+ }
38
+
39
+ with open(manifest_path, "w") as f:
40
+ json.dump(meta, f)
41
+
42
+ self.model._initialize_configs(
43
+ manifest_path=manifest_path,
44
+ max_speakers=8,
45
+ num_speakers=None,
46
+ tmpdir=temp_path,
47
+ batch_size=24,
48
+ num_workers=0,
49
+ verbose=True,
50
+ )
51
+ self.model.clustering_embedding.clus_diar_model._diarizer_params.out_dir = (
52
+ temp_path
53
+ )
54
+ self.model.clustering_embedding.clus_diar_model._diarizer_params.manifest_filepath = (
55
+ manifest_path
56
+ )
57
+ self.model.msdd_model.cfg.test_ds.manifest_filepath = manifest_path
58
+ self.model.diarize()
59
+
60
+ pred_labels_clus = rttm_to_labels(
61
+ os.path.join(temp_path, "pred_rttms", "mono_file.rttm")
62
+ )
63
+
64
+ labels = []
65
+ for label in pred_labels_clus:
66
+ start, end, speaker = label.split()
67
+ start, end = float(start), float(end)
68
+ start, end = int(start * 1000), int(end * 1000)
69
+ labels.append((start, end, int(speaker.split("_")[1])))
70
+
71
+ labels = sorted(labels, key=lambda x: x[0])
72
+
73
+ return labels
74
+
75
+
76
+ def create_config():
77
+ config = OmegaConf.load(
78
+ os.path.join(os.path.dirname(__file__), "diar_infer_telephonic.yaml")
79
+ )
80
+ pretrained_vad = "vad_multilingual_marblenet"
81
+ pretrained_speaker_model = "titanet_large"
82
+
83
+ config.diarizer.out_dir = None
84
+ config.diarizer.manifest_filepath = None
85
+ config.diarizer.speaker_embeddings.model_path = pretrained_speaker_model
86
+ config.diarizer.oracle_vad = (
87
+ False # compute VAD provided with model_path to vad config
88
+ )
89
+ config.diarizer.clustering.parameters.oracle_num_speakers = False
90
+
91
+ # Here, we use our in-house pretrained NeMo VAD model
92
+ config.diarizer.vad.model_path = pretrained_vad
93
+ config.diarizer.vad.parameters.onset = 0.8
94
+ config.diarizer.vad.parameters.offset = 0.6
95
+ config.diarizer.vad.parameters.pad_offset = -0.05
96
+ config.diarizer.msdd_model.model_path = (
97
+ "diar_msdd_telephonic" # Telephonic speaker diarization model
98
+ )
99
+
100
+ return config
@@ -0,0 +1,247 @@
1
+ import argparse
2
+ import logging
3
+ import os
4
+ import re
5
+
6
+ import faster_whisper
7
+ import torch
8
+
9
+ from ctc_forced_aligner import (
10
+ generate_emissions,
11
+ get_alignments,
12
+ get_spans,
13
+ load_alignment_model,
14
+ postprocess_results,
15
+ preprocess_text,
16
+ )
17
+ from deepmultilingualpunctuation import PunctuationModel
18
+
19
+ from helpers import (
20
+ cleanup,
21
+ find_numeral_symbol_tokens,
22
+ get_realigned_ws_mapping_with_punctuation,
23
+ get_sentences_speaker_mapping,
24
+ get_speaker_aware_transcript,
25
+ get_words_speaker_mapping,
26
+ langs_to_iso,
27
+ process_language_arg,
28
+ punct_model_langs,
29
+ whisper_langs,
30
+ write_srt,
31
+ )
32
+
33
+ mtypes = {"cpu": "int8", "cuda": "float16"}
34
+
35
+ pid = os.getpid()
36
+ temp_outputs_dir = f"temp_outputs_{pid}"
37
+ temp_path = os.path.join(os.getcwd(), "temp_outputs")
38
+ os.makedirs(temp_path, exist_ok=True)
39
+
40
+ # Initialize parser
41
+ parser = argparse.ArgumentParser()
42
+ parser.add_argument(
43
+ "-a", "--audio", help="name of the target audio file", required=True
44
+ )
45
+ parser.add_argument(
46
+ "--no-stem",
47
+ action="store_false",
48
+ dest="stemming",
49
+ default=True,
50
+ help="Disables source separation."
51
+ "This helps with long files that don't contain a lot of music.",
52
+ )
53
+
54
+ parser.add_argument(
55
+ "--suppress_numerals",
56
+ action="store_true",
57
+ dest="suppress_numerals",
58
+ default=False,
59
+ help="Suppresses Numerical Digits."
60
+ "This helps the diarization accuracy but converts all digits into written text.",
61
+ )
62
+
63
+ parser.add_argument(
64
+ "--whisper-model",
65
+ dest="model_name",
66
+ default="medium.en",
67
+ help="name of the Whisper model to use",
68
+ )
69
+
70
+ parser.add_argument(
71
+ "--batch-size",
72
+ type=int,
73
+ dest="batch_size",
74
+ default=8,
75
+ help="Batch size for batched inference, reduce if you run out of memory, "
76
+ "set to 0 for original whisper longform inference",
77
+ )
78
+
79
+ parser.add_argument(
80
+ "--language",
81
+ type=str,
82
+ default=None,
83
+ choices=whisper_langs,
84
+ help="Language spoken in the audio, specify None to perform language detection",
85
+ )
86
+
87
+ parser.add_argument(
88
+ "--device",
89
+ dest="device",
90
+ default="cuda" if torch.cuda.is_available() else "cpu",
91
+ help="if you have a GPU use 'cuda', otherwise 'cpu'",
92
+ )
93
+
94
+ parser.add_argument(
95
+ "--diarizer",
96
+ default="msdd",
97
+ choices=["msdd"],
98
+ help="Choose the diarization model to use",
99
+ )
100
+
101
+ args = parser.parse_args()
102
+ language = process_language_arg(args.language, args.model_name)
103
+
104
+ if args.stemming:
105
+ # Isolate vocals from the rest of the audio
106
+
107
+ return_code = os.system(
108
+ f'python -m demucs.separate -n htdemucs --two-stems=vocals "{args.audio}" -o "{temp_outputs_dir}" --device "{args.device}"'
109
+ )
110
+
111
+ if return_code != 0:
112
+ logging.warning(
113
+ "Source splitting failed, using original audio file. "
114
+ "Use --no-stem argument to disable it."
115
+ )
116
+ vocal_target = args.audio
117
+ else:
118
+ vocal_target = os.path.join(
119
+ temp_outputs_dir,
120
+ "htdemucs",
121
+ os.path.splitext(os.path.basename(args.audio))[0],
122
+ "vocals.wav",
123
+ )
124
+ else:
125
+ vocal_target = args.audio
126
+
127
+
128
+ # Transcribe the audio file
129
+
130
+ whisper_model = faster_whisper.WhisperModel(
131
+ args.model_name, device=args.device, compute_type=mtypes[args.device]
132
+ )
133
+ whisper_pipeline = faster_whisper.BatchedInferencePipeline(whisper_model)
134
+ audio_waveform = faster_whisper.decode_audio(vocal_target)
135
+ suppress_tokens = (
136
+ find_numeral_symbol_tokens(whisper_model.hf_tokenizer)
137
+ if args.suppress_numerals
138
+ else [-1]
139
+ )
140
+
141
+ if args.batch_size > 0:
142
+ transcript_segments, info = whisper_pipeline.transcribe(
143
+ audio_waveform,
144
+ language,
145
+ suppress_tokens=suppress_tokens,
146
+ batch_size=args.batch_size,
147
+ )
148
+ else:
149
+ transcript_segments, info = whisper_model.transcribe(
150
+ audio_waveform,
151
+ language,
152
+ suppress_tokens=suppress_tokens,
153
+ vad_filter=True,
154
+ )
155
+
156
+ full_transcript = "".join(segment.text for segment in transcript_segments)
157
+
158
+ # clear gpu vram
159
+ del whisper_model, whisper_pipeline
160
+ torch.cuda.empty_cache()
161
+
162
+ # Forced Alignment
163
+ alignment_model, alignment_tokenizer = load_alignment_model(
164
+ args.device,
165
+ dtype=torch.float16 if args.device == "cuda" else torch.float32,
166
+ )
167
+
168
+ emissions, stride = generate_emissions(
169
+ alignment_model,
170
+ torch.from_numpy(audio_waveform)
171
+ .to(alignment_model.dtype)
172
+ .to(alignment_model.device),
173
+ batch_size=args.batch_size,
174
+ )
175
+
176
+ del alignment_model
177
+ torch.cuda.empty_cache()
178
+
179
+ tokens_starred, text_starred = preprocess_text(
180
+ full_transcript,
181
+ romanize=True,
182
+ language=langs_to_iso[info.language],
183
+ )
184
+
185
+ segments, scores, blank_token = get_alignments(
186
+ emissions,
187
+ tokens_starred,
188
+ alignment_tokenizer,
189
+ )
190
+
191
+ spans = get_spans(tokens_starred, segments, blank_token)
192
+
193
+ word_timestamps = postprocess_results(text_starred, spans, stride, scores)
194
+
195
+ if args.diarizer == "msdd":
196
+ from diarization import MSDDDiarizer
197
+
198
+ diarizer_model = MSDDDiarizer(device=args.device)
199
+
200
+ speaker_ts = diarizer_model.diarize(torch.from_numpy(audio_waveform).unsqueeze(0))
201
+ del diarizer_model
202
+ torch.cuda.empty_cache()
203
+
204
+ wsm = get_words_speaker_mapping(word_timestamps, speaker_ts, "start")
205
+
206
+ if info.language in punct_model_langs:
207
+ # restoring punctuation in the transcript to help realign the sentences
208
+ punct_model = PunctuationModel(model="kredor/punctuate-all")
209
+
210
+ words_list = list(map(lambda x: x["word"], wsm))
211
+
212
+ labled_words = punct_model.predict(words_list, chunk_size=230)
213
+
214
+ ending_puncts = ".?!"
215
+ model_puncts = ".,;:!?"
216
+
217
+ # We don't want to punctuate U.S.A. with a period. Right?
218
+ is_acronym = lambda x: re.fullmatch(r"\b(?:[a-zA-Z]\.){2,}", x)
219
+
220
+ for word_dict, labeled_tuple in zip(wsm, labled_words):
221
+ word = word_dict["word"]
222
+ if (
223
+ word
224
+ and labeled_tuple[1] in ending_puncts
225
+ and (word[-1] not in model_puncts or is_acronym(word))
226
+ ):
227
+ word += labeled_tuple[1]
228
+ if word.endswith(".."):
229
+ word = word.rstrip(".")
230
+ word_dict["word"] = word
231
+
232
+ else:
233
+ logging.warning(
234
+ f"Punctuation restoration is not available for {info.language} language."
235
+ " Using the original punctuation."
236
+ )
237
+
238
+ wsm = get_realigned_ws_mapping_with_punctuation(wsm)
239
+ ssm = get_sentences_speaker_mapping(wsm, speaker_ts)
240
+
241
+ with open(f"{os.path.splitext(args.audio)[0]}.txt", "w", encoding="utf-8-sig") as f:
242
+ get_speaker_aware_transcript(ssm, f)
243
+
244
+ with open(f"{os.path.splitext(args.audio)[0]}.srt", "w", encoding="utf-8-sig") as srt:
245
+ write_srt(ssm, srt)
246
+
247
+ cleanup(temp_path)