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)}")