From e87f30173d9671d20a77e197aa0d83f6ba2a2fad Mon Sep 17 00:00:00 2001 From: Robert H Date: Thu, 28 May 2026 23:14:39 -0500 Subject: [PATCH] get good --- main.py | 60 ++++++++++++++++++++++++++---------------------- requirements.txt | 1 - 2 files changed, 32 insertions(+), 29 deletions(-) diff --git a/main.py b/main.py index b1a3a61..4543a12 100755 --- a/main.py +++ b/main.py @@ -6,9 +6,10 @@ 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 torchaudio import whisper @dataclasses.dataclass @@ -27,12 +28,24 @@ def srt_timestamp(milis): return f"{h:02}:{m:02}:{s:02},{milis:03}" def main(): - parser = argparse.ArgumentParser() + parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) parser.add_argument("input") - parser.add_argument("-o", "--output", default="output.srt") + 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() - model = whisper.load_model("turbo") + + 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"]: @@ -45,36 +58,35 @@ def main(): classifier = EncoderClassifier.from_hparams( source="speechbrain/spkrec-ecapa-voxceleb", + run_opts={"device": str(dev)}, ) - signal, sample_rate = torchaudio.load(args.input) - - print(sample_rate) - # resample to 16khz - if sample_rate != 16000: - resampler = torchaudio.transforms.Resample(sample_rate, 16000) - signal = resampler(signal) - sample_rate = 16000 - - # convert to mono if needed - if signal.shape[0] > 1: - signal = torch.mean(signal, dim=0, keepdim=True) + 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] + chunk = signal[start_sample:end_sample] # normalize chunk /= chunk.abs().max() # SpeechBrain expects [batch, time] - emb = classifier.encode_batch(chunk) + 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) @@ -86,9 +98,10 @@ def main(): ) labels = clustering.fit_predict(embeddings) - for seg, label in zip(segments, labels): + 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)) @@ -96,14 +109,5 @@ def main(): srt_file.write(f"{i}\n{start} --> {end}\n{seg.text}\n\n") print(start, end, seg.text) - from sklearn.decomposition import PCA - import matplotlib.pyplot as plt - - pca = PCA(n_components=2) - points = pca.fit_transform(embeddings) - - plt.scatter(points[:,0], points[:,1]) - plt.show() - if __name__ == "__main__": main() diff --git a/requirements.txt b/requirements.txt index 8af18af..efeb894 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,3 @@ openai-whisper==20250625 scikit-learn==1.8.0 speechbrain==1.1.0 torch==2.12.0 -torchaudio==2.11.0