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
@@ -1,4 +1,5 @@
from . import (
fireredasr2,
glmasr,
lasr_ctc,
parakeet,
Expand Down
35 changes: 35 additions & 0 deletions mlx_audio/stt/models/fireredasr2/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
# FireRedASR2-AED

MLX implementation of Xiaohongshu's FireRedASR2-AED, a conformer encoder + transformer decoder model for Chinese and English speech recognition.

## Model

| Model | Parameters | Description |
|-------|------------|-------------|
| FireRedASR2-AED | ~1.18B | Conformer-AED, Mandarin + English + Chinese dialects |

Converted weights are available on Hugging Face at [`mlx-community/FireRedASR2-AED-mlx`](https://huggingface.co/mlx-community/FireRedASR2-AED-mlx).

## Python Usage

```python
from mlx_audio.stt import load

model = load("mlx-community/FireRedASR2-AED-mlx")

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

# with custom beam search parameters
result = model.generate("audio.wav", beam_size=5, softmax_smoothing=1.25, length_penalty=0.6)
```

## Architecture

- 80-dim Kaldi fbank frontend with CMVN normalization
- Conv2d subsampling (4x temporal downsampling)
- 16 layer conformer encoder with relative positional attention (Transformer-XL style)
- 16 layer transformer decoder with cross attention and GELU FFN
- Beam search decoding with GNMT length penalty
- Hybrid tokenizer: Chinese characters + English BPE (SentencePiece, ~8.7k vocab)
- 1280 hidden dim, 20 attention heads, 5120 FFN dim
3 changes: 3 additions & 0 deletions mlx_audio/stt/models/fireredasr2/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from .config import ModelConfig
from .fireredasr2 import FireRedASR2
from .fireredasr2 import FireRedASR2 as Model
55 changes: 55 additions & 0 deletions mlx_audio/stt/models/fireredasr2/config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
from dataclasses import dataclass, field
from typing import Optional


@dataclass
class EncoderConfig:
n_layers: int = 16
n_head: int = 20
d_model: int = 1280
kernel_size: int = 33
pe_maxlen: int = 5000

@classmethod
def from_dict(cls, d):
return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__})


@dataclass
class DecoderConfig:
n_layers: int = 16
n_head: int = 20
d_model: int = 1280
pe_maxlen: int = 5000

@classmethod
def from_dict(cls, d):
return cls(**{k: v for k, v in d.items() if k in cls.__dataclass_fields__})


@dataclass
class ModelConfig:
model_type: str = "fireredasr2"
idim: int = 80
odim: int = 8667
d_model: int = 1280
sos_id: int = 3
eos_id: int = 4
pad_id: int = 2
blank_id: int = 0
encoder: EncoderConfig = field(default_factory=EncoderConfig)
decoder: DecoderConfig = field(default_factory=DecoderConfig)
model_path: Optional[str] = None

@classmethod
def from_dict(cls, d):
enc_d = d.get("encoder", {})
dec_d = d.get("decoder", {})
enc = EncoderConfig.from_dict(enc_d) if enc_d else EncoderConfig()
dec = DecoderConfig.from_dict(dec_d) if dec_d else DecoderConfig()
top_fields = {
k: v
for k, v in d.items()
if k in cls.__dataclass_fields__ and k not in ("encoder", "decoder")
}
return cls(encoder=enc, decoder=dec, **top_fields)
Loading