Skip to content
Merged
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions mlx_audio/stt/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
lasr_ctc,
parakeet,
qwen3_asr,
sensevoice,
voxtral,
voxtral_realtime,
wav2vec,
Expand Down
39 changes: 39 additions & 0 deletions mlx_audio/stt/models/sensevoice/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
# SenseVoice

MLX implementation of [SenseVoice Small](https://github.com/FunAudioLLM/SenseVoice) from Alibaba DAMO Academy. Non-autoregressive CTC-based model with speech recognition, language identification, emotion recognition, and audio event detection.

## Features

- 50+ languages (strong on zh, en, ja, ko, yue)
- emotion detection (happy, sad, angry, neutral, etc.)
- audio event detection (speech, bgm, laughter, applause)
- non-autoregressive inference (very fast, ~70ms for 10s audio)
- ~234M parameters

## Usage

```python
from mlx_audio.stt import load

# downloads automatically from huggingface
model = load("mlx-community/SenseVoiceSmall")

result = model.generate("audio.wav", language="auto")
print(result.text)
print(result.language)

# rich info (emotion, event) is in the first segment
seg = result.segments[0]
print(seg["emotion"]) # e.g. "neutral"
print(seg["event"]) # e.g. "Speech"
```

### Language options

`"auto"`, `"zh"`, `"en"`, `"ja"`, `"ko"`, `"yue"`, `"nospeech"`

### Inverse text normalization

```python
result = model.generate("audio.wav", language="auto", use_itn=True)
```
3 changes: 3 additions & 0 deletions mlx_audio/stt/models/sensevoice/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from .config import ModelConfig
from .sensevoice import SenseVoiceSmall
from .sensevoice import SenseVoiceSmall as Model
95 changes: 95 additions & 0 deletions mlx_audio/stt/models/sensevoice/config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,95 @@
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional


@dataclass
class EncoderConfig:
output_size: int = 512
attention_heads: int = 4
linear_units: int = 2048
num_blocks: int = 50
tp_blocks: int = 20
dropout_rate: float = 0.1
positional_dropout_rate: float = 0.1
attention_dropout_rate: float = 0.1
kernel_size: int = 11
sanm_shift: int = 0
normalize_before: bool = True

@classmethod
def from_dict(cls, params: Dict[str, Any]) -> "EncoderConfig":
import inspect

valid = {
k: v for k, v in params.items() if k in inspect.signature(cls).parameters
}
# Handle the typo in upstream config ("sanm_shfit" -> "sanm_shift")
if "sanm_shfit" in params and "sanm_shift" not in params:
valid["sanm_shift"] = params["sanm_shfit"]
return cls(**valid)


@dataclass
class FrontendConfig:
fs: int = 16000
window: str = "hamming"
n_mels: int = 80
frame_length: int = 25
frame_shift: int = 10
lfr_m: int = 7
lfr_n: int = 6

@classmethod
def from_dict(cls, params: Dict[str, Any]) -> "FrontendConfig":
import inspect

valid = {
k: v for k, v in params.items() if k in inspect.signature(cls).parameters
}
return cls(**valid)


@dataclass
class ModelConfig:
model_type: str = "sensevoice"
vocab_size: int = 25055
input_size: int = 560
encoder_conf: Optional[EncoderConfig] = None
frontend_conf: Optional[FrontendConfig] = None
# CMVN stats (loaded from am.mvn)
cmvn_means: Optional[List[float]] = None
cmvn_istd: Optional[List[float]] = None
# Path for loading ancillary files
model_path: Optional[str] = None

def __post_init__(self):
if self.encoder_conf is None:
self.encoder_conf = EncoderConfig()
elif isinstance(self.encoder_conf, dict):
self.encoder_conf = EncoderConfig.from_dict(self.encoder_conf)
if self.frontend_conf is None:
self.frontend_conf = FrontendConfig()
elif isinstance(self.frontend_conf, dict):
self.frontend_conf = FrontendConfig.from_dict(self.frontend_conf)

@classmethod
def from_dict(cls, params: Dict[str, Any]) -> "ModelConfig":
import inspect

encoder_conf = params.pop("encoder_conf", None)
frontend_conf = params.pop("frontend_conf", None)
valid = {
k: v for k, v in params.items() if k in inspect.signature(cls).parameters
}
config = cls(**valid)
if encoder_conf is not None:
if isinstance(encoder_conf, dict):
config.encoder_conf = EncoderConfig.from_dict(encoder_conf)
else:
config.encoder_conf = encoder_conf
if frontend_conf is not None:
if isinstance(frontend_conf, dict):
config.frontend_conf = FrontendConfig.from_dict(frontend_conf)
else:
config.frontend_conf = frontend_conf
return config
Loading