diff --git a/pyproject.toml b/pyproject.toml index 9d955f4..4bebe7e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,4 +30,4 @@ email = "robert@jadedpasta.net" Repository = "https://gitea.jadedpasta.net/jadedpasta/transcribble" [project.scripts] -transcribble = "transcribble:main" +transcribble = "transcribble.cli:main" diff --git a/src/transcribble/__init__.py b/src/transcribble/__init__.py index 4543a12..e69de29 100755 --- a/src/transcribble/__init__.py +++ b/src/transcribble/__init__.py @@ -1,113 +0,0 @@ -#!/usr/bin/env python - -import argparse -import dataclasses - -import numpy as np -from sklearn.cluster import AgglomerativeClustering -from sklearn.preprocessing import normalize -from speechbrain.dataio import audio_io -from speechbrain.dataio.preprocess import AudioNormalizer -from speechbrain.inference.speaker import EncoderClassifier -import torch -import whisper - -@dataclasses.dataclass -class Segment: - start: np.float32 - end: np.float32 - text: str - -def srt_timestamp(milis): - h = milis // (1000 * 60 * 60) - milis %= 1000 * 60 * 60 - m = milis // (1000 * 60) - milis %= 1000 * 60 - s = milis // 1000 - milis %= 1000 - return f"{h:02}:{m:02}:{s:02},{milis:03}" - -def main(): - parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) - parser.add_argument("input") - parser.add_argument("-o", "--output", default="output.srt", - help="Output subtitle (SRT) file path") - parser.add_argument("-s", "--speakers", type=int, default=2, - help="Number of speakers in the audio") - parser.add_argument("-m", "--model", default='turbo', choices=whisper.available_models(), - help="Whisper model to use") - parser.add_argument("--gpu", action="store_true", help="run on GPU if available") - args = parser.parse_args() - - - has_gpu = torch.cuda.is_available() - if args.gpu and not has_gpu: - print("GPU not available. Falling back to CPU") - dev = torch.device("cuda" if has_gpu and args.gpu else "cpu") - - model = whisper.load_model("turbo").to(dev) - result = model.transcribe(args.input, word_timestamps=True, language="en") - segments = [] - for seg in result["segments"]: - words = seg["words"] - start = words[0]["start"] - end = words[-1]["end"] - segments.append(Segment(start, end, seg["text"].strip())) - - print("Transcription done. Starting speaker classification") - - classifier = EncoderClassifier.from_hparams( - source="speechbrain/spkrec-ecapa-voxceleb", - run_opts={"device": str(dev)}, - ) - - signal, sample_rate = audio_io.load(args.input, channels_first=False) - audio_normalizer = AudioNormalizer() - signal = audio_normalizer(signal, sample_rate) - sample_rate = audio_normalizer.sample_rate - - embedding_segments = [] - embeddings = [] - for seg in segments: - if seg.end - seg.start < 0.1: - # Segment is too short to bother trying to classify - continue - - start_sample = int(seg.start * sample_rate) - end_sample = int(seg.end * sample_rate) - - chunk = signal[start_sample:end_sample] - - # normalize - chunk /= chunk.abs().max() - - # SpeechBrain expects [batch, time] - emb = classifier.encode_batch(chunk.unsqueeze(0).to(dev)) - emb = emb.squeeze().detach().cpu().numpy() - - embeddings.append(emb) - embedding_segments.append(seg) - embeddings = np.vstack(embeddings) - embeddings = normalize(embeddings) - - n_speakers = 2 - clustering = AgglomerativeClustering( - n_clusters=n_speakers, - metric="cosine", - linkage="average" - ) - labels = clustering.fit_predict(embeddings) - - for seg, label in zip(embedding_segments, labels): - seg.text = f"Speaker {label + 1}: {seg.text}" - - print("Final transcription:") - with open(args.output, "w") as srt_file: - for i, seg in enumerate(segments, 1): - start = srt_timestamp(round(seg.start * 1000)) - end = srt_timestamp(round(seg.end * 1000)) - srt_file.write(f"{i}\n{start} --> {end}\n{seg.text}\n\n") - print(start, end, seg.text) - -if __name__ == "__main__": - main() diff --git a/src/transcribble/cli.py b/src/transcribble/cli.py new file mode 100644 index 0000000..7e17b9a --- /dev/null +++ b/src/transcribble/cli.py @@ -0,0 +1,58 @@ +import argparse +import sys + +import logging + +class IgnoreHFWarning(logging.Filter): + def filter(self, record): + return "unauthenticated requests to the HF Hub" not in record.getMessage() + +logger = logging.getLogger("huggingface_hub.utils._http") +logger.addFilter(IgnoreHFWarning()) + +def main(): + parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) + parser.add_argument("input") + parser.add_argument("-o", "--output", default="output.srt", + help="Output subtitle (SRT) file path") + parser.add_argument("-s", "--speakers", type=int, default=2, + help="Number of speakers in the audio") + parser.add_argument("-m", "--model", default='turbo', + help="Whisper model to use ('?' to see available options)") + parser.add_argument("--gpu", action="store_true", help="run on GPU if available") + args = parser.parse_args() + + import torch + import whisper + + from .transcribe import transcribe + from .diarize import diarize + from .output import write_srt + + models = whisper.available_models() + + if args.model == "?": + for model in models: + print(model) + return + + if args.model not in models: + parser.error(f"'{args.model}' is not a valid whisper model") + + has_gpu = torch.cuda.is_available() + if args.gpu and not has_gpu: + print("GPU not available. Falling back to CPU") + dev = torch.device("cuda" if has_gpu and args.gpu else "cpu") + + print(f"Starting transcription of {args.input}") + segments = transcribe(args.input, dev) + + if args.speakers >= 2: + print("Transcription done. Starting speaker classification") + diarize(args.input, segments, dev, n_speakers=args.speakers) + + print("Final transcription:") + write_srt(args.output, segments, show=True) + +if __name__ == "__main__": + main() diff --git a/src/transcribble/data.py b/src/transcribble/data.py new file mode 100644 index 0000000..9d88624 --- /dev/null +++ b/src/transcribble/data.py @@ -0,0 +1,10 @@ +import dataclasses + +import numpy as np + +@dataclasses.dataclass +class Segment: + start: np.float32 + end: np.float32 + text: str + diff --git a/src/transcribble/diarize.py b/src/transcribble/diarize.py new file mode 100644 index 0000000..8306f64 --- /dev/null +++ b/src/transcribble/diarize.py @@ -0,0 +1,51 @@ +import numpy as np +from sklearn.cluster import AgglomerativeClustering +from sklearn.preprocessing import normalize +from speechbrain.dataio import audio_io +from speechbrain.dataio.preprocess import AudioNormalizer +from speechbrain.inference.speaker import EncoderClassifier + +def diarize(input_file, segments, dev, n_speakers=2): + classifier = EncoderClassifier.from_hparams( + source="speechbrain/spkrec-ecapa-voxceleb", + run_opts={"device": str(dev)}, + ) + + signal, sample_rate = audio_io.load(input_file, channels_first=False) + audio_normalizer = AudioNormalizer() + signal = audio_normalizer(signal, sample_rate) + sample_rate = audio_normalizer.sample_rate + + embedding_segments = [] + embeddings = [] + for seg in segments: + if seg.end - seg.start < 0.1: + # Segment is too short to bother trying to classify + continue + + start_sample = int(seg.start * sample_rate) + end_sample = int(seg.end * sample_rate) + + chunk = signal[start_sample:end_sample] + + # normalize + chunk /= chunk.abs().max() + + # SpeechBrain expects [batch, time] + emb = classifier.encode_batch(chunk.unsqueeze(0).to(dev)) + emb = emb.squeeze().detach().cpu().numpy() + + embeddings.append(emb) + embedding_segments.append(seg) + embeddings = np.vstack(embeddings) + embeddings = normalize(embeddings) + + clustering = AgglomerativeClustering( + n_clusters=n_speakers, + metric="cosine", + linkage="average" + ) + labels = clustering.fit_predict(embeddings) + + for seg, label in zip(embedding_segments, labels): + seg.text = f"Speaker {label + 1}: {seg.text}" diff --git a/src/transcribble/output.py b/src/transcribble/output.py new file mode 100644 index 0000000..7295aa8 --- /dev/null +++ b/src/transcribble/output.py @@ -0,0 +1,17 @@ +def srt_timestamp(milis): + h = milis // (1000 * 60 * 60) + milis %= 1000 * 60 * 60 + m = milis // (1000 * 60) + milis %= 1000 * 60 + s = milis // 1000 + milis %= 1000 + return f"{h:02}:{m:02}:{s:02},{milis:03}" + +def write_srt(file_path, segments, show=False): + with open(file_path, "w") as srt_file: + for i, seg in enumerate(segments, 1): + start = srt_timestamp(round(seg.start * 1000)) + end = srt_timestamp(round(seg.end * 1000)) + srt_file.write(f"{i}\n{start} --> {end}\n{seg.text}\n\n") + if show: + print(start, end, seg.text) diff --git a/src/transcribble/transcribe.py b/src/transcribble/transcribe.py new file mode 100644 index 0000000..38b4644 --- /dev/null +++ b/src/transcribble/transcribe.py @@ -0,0 +1,16 @@ +import torch +import whisper + +from .data import Segment + +def transcribe(input_file, dev): + model = whisper.load_model("turbo").to(dev) + result = model.transcribe(input_file, word_timestamps=True, language="en", + fp16=dev != torch.device("cpu")) + segments = [] + for seg in result["segments"]: + words = seg["words"] + start = words[0]["start"] + end = words[-1]["end"] + segments.append(Segment(start, end, seg["text"].strip())) + return segments