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.
- taters/Taters.py +115 -0
- taters/__init__.py +0 -0
- taters/audio/__init__.py +0 -0
- taters/audio/convert_to_wav.py +112 -0
- taters/audio/diarize_with_thirdparty.py +5 -0
- taters/audio/diarizer/__init__.py +0 -0
- taters/audio/diarizer/whisper-diarization/__init__.py +0 -0
- taters/audio/diarizer/whisper-diarization/diarization/__init__.py +3 -0
- taters/audio/diarizer/whisper-diarization/diarization/msdd/msdd.py +100 -0
- taters/audio/diarizer/whisper-diarization/diarize.py +247 -0
- taters/audio/diarizer/whisper-diarization/diarize_custom.py +374 -0
- taters/audio/diarizer/whisper-diarization/diarize_parallel.py +269 -0
- taters/audio/diarizer/whisper-diarization/helpers.py +552 -0
- taters/audio/diarizer/whisper_diar_wrapper.py +293 -0
- taters/audio/extract_wav_from_video.py +177 -0
- taters/audio/extract_whisper_embeddings.py +243 -0
- taters/audio/extract_whisper_embeddings_subproc.py +574 -0
- taters/audio/split_wav_by_speaker.py +182 -0
- taters/helpers/__init__.py +0 -0
- taters/helpers/feature_gather.py +349 -0
- taters/helpers/find_files.py +222 -0
- taters/helpers/text_gather.py +520 -0
- taters/pipelines/__init__.py +0 -0
- taters/pipelines/run_pipeline.py +384 -0
- taters/text/analyze_with_archetypes.py +238 -0
- taters/text/analyze_with_dictionaries.py +246 -0
- taters/text/dictionary_analyzers/__init__.py +0 -0
- taters/text/dictionary_analyzers/multi_archetype_analyzer.py +125 -0
- taters/text/dictionary_analyzers/multi_dict_analyzer.py +254 -0
- taters/text/extract_sentence_embeddings.py +329 -0
- taters/video/__init__.py +0 -0
- taters/video/extract_features.py +0 -0
- taters/video/features_basic.py +0 -0
- taters/video/features_motion.py +0 -0
- taters/video/features_ocr.py +0 -0
- taters/video/features_shots.py +0 -0
- taters/video/read_video.py +0 -0
- taters/video/windowing.py +0 -0
- taters-0.1.0.dist-info/METADATA +161 -0
- taters-0.1.0.dist-info/RECORD +44 -0
- taters-0.1.0.dist-info/WHEEL +5 -0
- taters-0.1.0.dist-info/entry_points.txt +2 -0
- taters-0.1.0.dist-info/licenses/LICENSE +21 -0
- 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
|
taters/audio/__init__.py
ADDED
|
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()
|
|
File without changes
|
|
File without changes
|
|
@@ -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)
|