Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 78 additions & 17 deletions mlx_audio/stt/models/vibevoice_asr/vibevoice_asr.py
Original file line number Diff line number Diff line change
Expand Up @@ -445,25 +445,78 @@ def from_pretrained(cls, model_path: str, **kwargs) -> "Model":

return load(model_path)

def _preprocess_audio(self, audio) -> mx.array:
@staticmethod
def _normalize_audio(
audio: np.ndarray,
target_dB_FS: float = -25.0,
eps: float = 1e-6,
) -> np.ndarray:
"""
Normalize audio to target dB FS level and avoid clipping.

This matches the official VibeVoice AudioNormalizer which adjusts
audio loudness to a consistent level before encoding.

Args:
audio: Input audio waveform as numpy array.
target_dB_FS: Target loudness in dB FS (default: -25).
eps: Small value to avoid division by zero.

Returns:
Normalized audio as numpy array.
"""
rms = np.sqrt(np.mean(audio**2))
scalar = 10 ** (target_dB_FS / 20) / (rms + eps)
audio = audio * scalar
max_val = np.max(np.abs(audio))
if max_val > 1.0:
audio = audio / (max_val + eps)
return audio

def _preprocess_audio(
self, audio, *, sampling_rate: Optional[int] = None
) -> mx.array:
"""
Preprocess audio for the model.

Handles loading, resampling to 24 kHz, normalizing loudness to
-25 dB FS, and trimming to the maximum supported duration.

Args:
audio: Audio path (str), waveform (np.ndarray/mx.array)
sampling_rate: Sample rate of the input waveform. Required when
*audio* is an array that is **not** already at 24 kHz so
that the audio can be resampled correctly.

Returns:
Audio tensor ready for encoding [B, T]
"""
from mlx_audio.stt.utils import load_audio
from mlx_audio.stt.utils import load_audio, resample_audio

SAMPLE_RATE = 24000
MAX_DURATION_SECONDS = 59 * 60 # 59 minutes max

if isinstance(audio, str):
audio = load_audio(audio, sr=SAMPLE_RATE)
elif not isinstance(audio, mx.array):
audio = mx.array(audio)
else:
# Convert to numpy for resampling / normalization
if isinstance(audio, mx.array):
audio_np = np.array(audio, copy=False)
else:
audio_np = np.asarray(audio, dtype=np.float32)

# Flatten to 1-D for resampling
if audio_np.ndim > 1:
audio_np = audio_np.squeeze()

# Resample to 24 kHz when the caller specifies a different rate
if sampling_rate is not None and sampling_rate != SAMPLE_RATE:
audio_np = resample_audio(audio_np, sampling_rate, SAMPLE_RATE)

# Normalize loudness to -25 dB FS (matches official pipeline)
audio_np = self._normalize_audio(audio_np)

audio = mx.array(audio_np)

# Ensure 1D or 2D
if audio.ndim == 3:
Expand Down Expand Up @@ -584,13 +637,14 @@ def generate(
audio,
*,
context: Optional[str] = None,
sampling_rate: Optional[int] = None,
max_tokens: int = 8192,
temperature: float = 0.0,
top_p: float = 0.95,
top_k: int = 25,
min_p: float = 0.02,
top_p: float = 1.0,
top_k: int = 0,
min_p: float = 0.0,
min_tokens_to_keep: int = 1,
repetition_penalty: Optional[float] = 1.2,
repetition_penalty: Optional[float] = 1.0,
repetition_context_size: int = 100,
prefill_step_size: int = 2048,
generation_stream: bool = False,
Expand All @@ -603,11 +657,14 @@ def generate(
Args:
audio: Audio path (str) or waveform (mx.array/np.array)
context: Optional context string (hotwords, metadata)
sampling_rate: Sample rate of the input waveform. When *audio*
is an array not at 24 kHz, provide its actual sample rate
so that it is resampled correctly.
max_tokens: Maximum tokens to generate
temperature: Sampling temperature (0 = greedy)
top_p: Top-p sampling
top_k: Top-k sampling
min_p: Min-p sampling
top_p: Top-p sampling (1.0 = no filtering)
top_k: Top-k sampling (0 = disabled)
min_p: Min-p sampling (0.0 = disabled)
min_tokens_to_keep: Min tokens for sampling
repetition_penalty: Penalty for repeated tokens (1.0 = no penalty)
repetition_context_size: Number of recent tokens to check for repetition
Expand All @@ -623,7 +680,7 @@ def generate(
start_time = time.time()

# Preprocess audio
audio_tensor = self._preprocess_audio(audio)
audio_tensor = self._preprocess_audio(audio, sampling_rate=sampling_rate)

# Encode speech
speech_features = self.encode_speech(audio_tensor, verbose=verbose)
Expand Down Expand Up @@ -695,13 +752,14 @@ def stream_transcribe(
audio,
*,
context: Optional[str] = None,
sampling_rate: Optional[int] = None,
max_tokens: int = 8192,
temperature: float = 0.0,
top_p: float = 0.95,
top_p: float = 1.0,
top_k: int = 0,
min_p: float = 0.0,
min_tokens_to_keep: int = 1,
repetition_penalty: Optional[float] = 1.2,
repetition_penalty: Optional[float] = 1.0,
repetition_context_size: int = 100,
prefill_step_size: int = 2048,
verbose: bool = False,
Expand All @@ -712,10 +770,13 @@ def stream_transcribe(
Args:
audio: Audio path (str) or waveform (mx.array/np.array)
context: Optional context string (hotwords, metadata)
sampling_rate: Sample rate of the input waveform. When *audio*
is an array not at 24 kHz, provide its actual sample rate
so that it is resampled correctly.
max_tokens: Maximum tokens to generate
temperature: Sampling temperature (0 = greedy)
top_p: Top-p sampling
top_k: Top-k sampling
top_p: Top-p sampling (1.0 = no filtering)
top_k: Top-k sampling (0 = disabled)
min_p: Min-p sampling
min_tokens_to_keep: Min tokens for sampling
repetition_penalty: Penalty for repeated tokens (1.0 = no penalty)
Expand All @@ -729,7 +790,7 @@ def stream_transcribe(
from mlx_lm.sample_utils import make_logits_processors, make_sampler

# Preprocess audio
audio_tensor = self._preprocess_audio(audio)
audio_tensor = self._preprocess_audio(audio, sampling_rate=sampling_rate)

# Encode speech
speech_features = self.encode_speech(audio_tensor, verbose=verbose)
Expand Down