File size: 5,039 Bytes
92a5c40 df92282 92a5c40 df92282 e57bd8f 92a5c40 df92282 92a5c40 df92282 92a5c40 df92282 92a5c40 df92282 92a5c40 e57bd8f 92a5c40 df92282 92a5c40 df92282 92a5c40 df92282 92a5c40 df92282 92a5c40 df92282 92a5c40 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | import librosa
import numpy as np
import os
import soundfile as sf
import tempfile
import torch
# These will be set by main.py at startup
ASR_MODEL = None
ASR_PROCESSOR = None
DEVICE = "cpu"
def set_asr_globals(model, processor, device):
global ASR_MODEL, ASR_PROCESSOR, DEVICE
ASR_MODEL = model
ASR_PROCESSOR = processor
DEVICE = device
def reduce_noise(audio_array: np.ndarray, sr: int = 16000, noise_reduce_strength: float = 0.7) -> np.ndarray:
try:
stft = librosa.stft(audio_array, n_fft=2048, hop_length=512)
magnitude, phase = librosa.magphase(stft)
frame_energies = np.sum(magnitude, axis=0)
noise_threshold = np.percentile(frame_energies, 10)
noise_frames = magnitude[:, frame_energies <= noise_threshold]
if noise_frames.shape[1] > 0:
noise_profile = np.mean(noise_frames, axis=1, keepdims=True)
else:
noise_profile = np.min(magnitude, axis=1, keepdims=True)
magnitude_denoised = magnitude - (noise_reduce_strength * noise_profile)
magnitude_denoised = np.maximum(magnitude_denoised, 0.0)
smoothing_factor = 0.05
magnitude_denoised = (1 - smoothing_factor) * magnitude_denoised + smoothing_factor * magnitude
stft_denoised = magnitude_denoised * phase
audio_denoised = librosa.istft(stft_denoised, hop_length=512, length=len(audio_array))
original_peak = np.abs(audio_array).max()
denoised_peak = np.abs(audio_denoised).max()
if denoised_peak > 0:
audio_denoised = audio_denoised * (original_peak / denoised_peak)
return audio_denoised
except Exception as e:
print(f"[WARNING] Noise reduction failed: {e}, returning original audio")
return audio_array
def should_reduce_noise() -> bool:
return os.environ.get("USE_NOISE_REDUCTION", "false").lower() == "true"
def convert_audio_to_wav(audio_bytes: bytes, target_sr: int = 16000, filename: str = None) -> np.ndarray:
if not audio_bytes or len(audio_bytes) == 0:
raise ValueError("No audio data received")
if filename:
ext = os.path.splitext(filename)[1].lower()
if not ext:
ext = ".wav"
else:
ext = ".wav"
with tempfile.NamedTemporaryFile(suffix=ext, delete=False) as tmp_file:
tmp_file.write(audio_bytes)
tmp_path = tmp_file.name
try:
try:
audio_array, sr = sf.read(tmp_path, dtype="float32")
if len(audio_array.shape) > 1:
audio_array = audio_array.mean(axis=1)
if sr != target_sr:
audio_array = librosa.resample(audio_array, orig_sr=sr, target_sr=target_sr)
except Exception:
audio_array, sr = librosa.load(
tmp_path,
sr=target_sr,
mono=True,
res_type="kaiser_best",
)
if audio_array is None or len(audio_array) == 0:
raise ValueError("Audio file is empty or unreadable")
max_val = np.abs(audio_array).max()
if max_val > 0:
if max_val > 1.0:
audio_array = audio_array / max_val
else:
raise ValueError("Audio contains only silence")
if should_reduce_noise():
audio_array = reduce_noise(audio_array, sr=target_sr, noise_reduce_strength=0.7)
return audio_array
except Exception as e:
raise ValueError(f"Failed to convert audio file '{filename or 'unknown'}': {str(e)}")
finally:
if os.path.exists(tmp_path):
try:
os.unlink(tmp_path)
except Exception:
pass
def validate_audio_duration(audio_array: np.ndarray, sr: int = 16000) -> bool:
duration = len(audio_array) / sr
if duration < 0.5:
raise ValueError(f"Audio too short: {duration:.2f}s (minimum 0.5s)")
if duration > 20:
raise ValueError(f"Audio too long: {duration:.2f}s (maximum 20s)")
return True
def transcribe_audio(audio_array: np.ndarray, sr: int = 16000, return_ctc_data: bool = False):
if ASR_MODEL is None or ASR_PROCESSOR is None:
raise RuntimeError("ASR model not loaded")
try:
inputs = ASR_PROCESSOR(
audio_array,
sampling_rate=sr,
return_tensors="pt",
padding=True,
)
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
with torch.inference_mode():
logits = ASR_MODEL(**inputs).logits
predicted_ids = torch.argmax(logits, dim=-1)
transcription = ASR_PROCESSOR.batch_decode(predicted_ids)[0]
probs = torch.nn.functional.softmax(logits, dim=-1)
confidence_scores = torch.max(probs, dim=-1)[0].cpu().numpy()[0]
if return_ctc_data:
return transcription.strip(), confidence_scores, logits, predicted_ids
return transcription.strip(), confidence_scores
except Exception as e:
raise RuntimeError(f"ASR transcription failed: {str(e)}")
|