Jahmori-R's picture
Sync from Jahmori-R/PureVoiceAI-HG via hub-sync
db05c78 verified
Raw History Blame Contribute Delete
11.7 kB
from time import perf_counter
import glob
import os
from dotenv import load_dotenv
import shutil
from pathlib import Path
import gradio as gr
import spaces
import torch
import torchaudio
import transformers
import numpy as np
import noisereduce as nr
from pyannote.audio import Pipeline
from transformers import AutoModelForSpeechSeq2Seq, AutoProcessor, pipeline
from speechbrain.inference.separation import SepformerSeparation
from pydub import AudioSegment
# Disables TF32 for to prevent lower accuracy in pyannote
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
# Passing .env file as environment variable on machine
load_dotenv()
# Hugging Face authentication token grab from .env
token = os.getenv("HF_TOKEN")
# GPU processing availability check
device = "cuda" if torch.cuda.is_available() else "cpu"
# Name of mixture folder
mixtures = Path(os.getcwd())/"mixtures"
# Name of separated file folder
sources = Path(os.getcwd()) / "sources"
# Path checking for removing source directory every run
def reset_dir():
remove_dir = [mixtures,sources]
# Getting current source directory
for dir in remove_dir:
if os.path.exists(dir):
try:
# Remove directorty recursively
shutil.rmtree(sources)
except:
raise gr.Error("Reset Failed")
gr.Info("System Reset",title="Notice")
# To clear all Gradio Blocks
return None, None, None, None, None
# Formatting the transcription log to Gradio chat display
def chatFormat(transcription_log):
chat = []
for i in transcription_log:
name = i.get("speaker")
text = i.get("text")
start = i.get("start")
end = i.get("end")
# Title format for gradio chatbot block
title_format = f"**({name}) [{start}s - {end}s]**"
role = "user" if name == "speaker_2" else "assistant"
chat.append({"role": role, "content":text, "metadata":{"title":title_format}})
return chat
@spaces.GPU
def process_audio(mixture, progress=gr.Progress()):
# Inference Time Processing
start_time = perf_counter()
progress(0, desc="Starting")
# Gradio processing alert
gr.Info("Processing Has Started", title="Notice")
# Creates mixture directory
Path.mkdir(mixtures, exist_ok=True)
# Creates separated directory
Path.mkdir(sources, exist_ok=True)
# ---------------------------------------------------------
# 1. INITIALIZE MODELS
# ---------------------------------------------------------
progress(0.12, "Loading RE-SepFormer Model...")
print("Intializing models...")
# Loading Resepformer Model
try:
separator = SepformerSeparation.from_hparams(
source=f"Jahmori-R/resepformer-librimix2spk",
hparams_file="hyperparams.yaml",
run_opts={"device": device},
)
print("RE-SepFormer model successfully loaded")
except Exception as e:
raise RuntimeError("Failed to load RE-SepFormer model") from e
progress(0.25, "Loading Pyannote Model...")
# Loading Pyannote Model
try:
diarize = Pipeline.from_pretrained(
"pyannote/speaker-diarization-community-1", token=token
)
# Assigning GPU processing to model
diarize.to(torch.device("cuda"))
print("Pyannote model successfully loaded")
except Exception as e:
raise RuntimeError("Failed to load Pyannote model") from e
progress(0.38, "Loading Whisper Model...")
# Loading Whisper Model
transformers.logging.set_verbosity_error() # Prevents WhisperAI logging in cmd line
try:
transcription_model_id = "openai/whisper-large-v3-turbo"
torch_dtype = torch.float16 if torch.cuda.is_available() else torch.float32
model = AutoModelForSpeechSeq2Seq.from_pretrained(
transcription_model_id,
dtype=torch_dtype,
low_cpu_mem_usage=True,
use_safetensors=True,
)
processor = AutoProcessor.from_pretrained(transcription_model_id)
transcribe = pipeline(
"automatic-speech-recognition",
model=model,
tokenizer=processor.tokenizer,
feature_extractor=processor.feature_extractor,
dtype=torch_dtype,
device="cuda",
)
print("Whisper model successfully loaded")
except Exception as e:
raise RuntimeError("Failed to load Whisper model") from e
# ---------------------------------------------------------
# LOAD & SEPARATE AUDIO
# ---------------------------------------------------------
# Calling ReSepformer model onto audio source
assert separator is not None, "Failed to load the Sepformer model."
progress(0.40,"Denoising Mixture...")
#Denoising mixture
sr, mixture_np = mixture
if len(mixture_np.shape) > 1:
# Swaps channels location with sample rate in numpy array
if mixture_np.shape[0] > mixture_np.shape[1]:
mixture_np = mixture_np.T
# If it has more than 2 channels take the first one
if len(mixture_np.shape) > 1 and mixture_np.shape[0] > 2:
mixture_np = mixture_np[0]
# Converting numpy to float32 for torchaudio
if mixture_np.dtype != np.float32:
# Converting integer value to float between -1.0 and 1.0
if np.issubdtype(mixture_np.dtype, np.integer):
mixture_np = mixture_np.astype(np.float32) / np.iinfo(mixture_np.dtype).max
else:
mixture_np = mixture_np.astype(np.float32)
# Apply noise gate onto waveform
denoised_mixture_np = nr.reduce_noise(y=mixture_np,sr=sr, prop_decrease=0.5,n_fft=512)
# Convert numpy data back to torch for saving
denoised_mixture = torch.from_numpy(denoised_mixture_np)
# Saving new denoised mixture
torchaudio.save(f"{mixtures}/mixture.wav",denoised_mixture,sr)
progress(0.50,"Separating Mixtures...")
# Running mixed audio onto separation model
print("Separating Audio Mixture")
est_sources = separator.separate_file(f"{mixtures}/mixture.wav")
print("Mixture separated")
# Saving separated sources to directory
torchaudio.save(
f"{sources}/speaker_1.wav", est_sources[:, :, 0].detach().cpu(), 16000
)
torchaudio.save(
f"{sources}/speaker_2.wav", est_sources[:, :, 1].detach().cpu(), 16000
)
# ---------------------------------------------------------
# VAD + TRANSCRIPTION ON SOURCES
# ---------------------------------------------------------
# Dictionary to store speaker information
unfiltered_log = [[], []]
progress(0.62, "Sorting Sources...")
# Storing both speakers
file_paths = sorted(
# Glob finds all similar files
glob.glob(f"{sources}/speaker_*.wav")
)
spk_1 = file_paths[0]
spk_2 = file_paths[1]
progress(0.75, "VAD + Transcribing...")
# Individual Speaker Processing
for i, path in enumerate(file_paths):
speaker_key = f"speaker_{i + 1}"
# Diarize speaker audio for time segments
print(f"Diarizing {speaker_key}")
diarization = diarize(path, num_speakers=1)
# Minimum timestamp duration in seconds
minimum_duration = 0.8
# Storing speaker voice activity timestamps to dictionary
for turn, _ in diarization.exclusive_speaker_diarization:
duration = turn.end - turn.start
if duration < minimum_duration:
continue
# Adding timestamp to list
unfiltered_log[i].append(
{
"start": float(round(turn.start, 2)),
"end": float(round(turn.end, 2)),
"speaker": speaker_key,
}
)
# Creating audio segment from detected voice activity
audio_seg = AudioSegment.from_file(path)
print("Transcribing Chunks...")
# Transcribing each voice activity detected by Pyannote
for segment in unfiltered_log[i]:
# Converting segment timestamps to milliseconds for pydub segmentation
start_ms = segment["start"] * 1000
end_ms = segment["end"] * 1000
chunk = audio_seg[start_ms:end_ms]
# Converting pydub AudioSegment to numpy array for whisper model
samples = np.array(chunk.get_array_of_samples())
samples = samples.astype(np.float32) / 32768.0
# Whisper Configutation
kwargs = {
"language": "english", # Language transcription
}
# Generating transcription given chunk segment
transcription = transcribe(samples, generate_kwargs=kwargs)
# Appending each segment transcription to dictionary
segment.update(transcription)
progress(0.88, "Sorting Speaker Logs....")
# Combined both speaker logs
combined_log = unfiltered_log[0] + unfiltered_log[1]
# Sorting by start time of timestamps
filtered_log = sorted(combined_log, key=lambda x: x["start"])
chat = chatFormat(filtered_log)
progress(1.00, "Task Finished")
print("Task Finished")
end_time = perf_counter()
# Inference time calculation
inference_time = f"{end_time - start_time:.2f} seconds"
return (spk_1, spk_2, chat, inference_time)
# ---------------------------------------------------------
# GRADIO INTERFACE CUSTOMIZATION
# ---------------------------------------------------------
# Custom css for gradio interface
css = """
#centered-row {
justify-content: center;
display: flex;
}
"""
# Gradio layout configuration
with gr.Blocks() as demo:
gr.Markdown("PureVoiceAI")
with gr.Row(variant="panel"):
audio_input = gr.Audio(
type="numpy", label="Input Audio Mixture | 2 Speakers"
)
with gr.Row(elem_id="centered-row"):
process_button = gr.Button(
"Process",
variant="primary",
scale=0,
)
reset_button = gr.Button("Restart", variant="stop", interactive=False, scale=0)
with gr.Row(variant="panel", visible=False) as source_row:
with gr.Column():
speaker_1_output = gr.Audio(label="Speaker 1", interactive=False)
speaker_2_output = gr.Audio(label="Speaker 2", interactive=False)
with gr.Column():
inference_time_log = gr.Textbox(label="Inference Time", interactive=False)
with gr.Row(variant="panel") as log_row:
logs = gr.Chatbot(label="Conversation Log",layout="bubble")
# Process Button Functionality
process_button.click(
fn=process_audio,
inputs=audio_input,
outputs=[speaker_1_output, speaker_2_output, logs, inference_time_log],
).then(
fn=lambda: (
gr.Row(visible=True),
gr.Button(interactive=False),
gr.Button(interactive=True),
),
inputs=None,
outputs=[source_row, process_button, reset_button],
)
# Reset Button Functionality
reset_button.click(
fn=reset_dir,
outputs=[
audio_input,
speaker_1_output,
speaker_2_output,
inference_time_log,
logs,
],
).then(
fn=lambda: (
gr.Row(visible=False),
gr.Button(interactive=True),
gr.Button(interactive=False),
),
inputs=None,
outputs=[source_row, process_button, reset_button],
)
if __name__ == "__main__":
demo.launch(theme=gr.Theme.from_hub("pseudolab/huggingface-korea-theme"), css=css)