Skip to content
Open
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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -158,7 +158,7 @@ for result in model.generate(
| **Canary** | NVIDIA's multilingual ASR with translation | 25 EU + RU, UK | [README](mlx_audio/stt/models/canary/README.md) |
| **Moonshine** | Useful Sensors' lightweight ASR | EN | [README](mlx_audio/stt/models/moonshine/README.md) |
| **MMS** | Meta's massively multilingual ASR with adapters | 1000+ | [README](mlx_audio/stt/models/mms/README.md) |
| **Granite Speech** | IBM's ASR + speech translation | EN, FR, DE, ES, PT, JA | [README](mlx_audio/stt/models/granite_speech/README.md) |
| **Granite Speech** | IBM's ASR + speech translation; plus variant adds speaker attribution & word timestamps | EN, FR, DE, ES, PT, JA | [README](mlx_audio/stt/models/granite_speech/README.md) · [plus](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) |
| **Granite Speech 5.0 TurboCTC** | IBM's fast encoder-only CTC ASR | EN | [README](mlx_audio/stt/models/granite_speech5_ctc/README.md) |
| **Qwen2-Audio** | Alibaba's multimodal audio understanding (ASR, captioning, emotion, translation) | Multiple | [mlx-community/Qwen2-Audio-7B-Instruct-4bit](https://huggingface.co/mlx-community/Qwen2-Audio-7B-Instruct-4bit) |
| **MOSS-Music** | OpenMOSS music understanding and lyrics ASR | EN, ZH | [README](mlx_audio/stt/models/moss_music/README.md) |
Expand Down
76 changes: 76 additions & 0 deletions docs/models/stt/granite-speech.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
---
title: Granite Speech
---

# Granite Speech

IBM's Granite Speech combines an audio encoder with a language-model decoder. MLX Audio supports the original speech checkpoint and the Plus variant through the same model package.

| Checkpoint | Tasks | Languages |
| --- | --- | --- |
| [Granite 4.0 1B Speech](https://huggingface.co/ibm-granite/granite-4.0-1b-speech) | Transcription, speech translation, keyword biasing | EN, FR, DE, ES, PT, JA |
| [Granite Speech 4.1 2B Plus](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) | Transcription, speaker attribution, word timestamps, keyword biasing | EN, FR, DE, ES, PT |

The examples below use IBM's native Plus checkpoint, which loads directly. Local MLX-converted and quantized checkpoints use the same API.

## Python

```python
from mlx_audio.stt import load

model = load("ibm-granite/granite-speech-4.1-2b-plus")

# Plain transcription is the default.
result = model.generate("audio.wav")
print(result.text)

# Speaker attribution returns speaker IDs and text for each detected turn.
result = model.generate("meeting.wav", task="saa")
for segment in result.segments:
print(segment["speaker_id"], segment["text"])

# Word timestamps are a separate task from speaker attribution.
result = model.generate("audio.wav", task="timestamps", max_tokens=8192)
for segment in result.segments:
for word in segment["words"]:
print(word["word"], word["start"], word["end"])
```

| Option | Behavior |
| --- | --- |
| `task="asr"` | Plain transcription; permits a custom `prompt` |
| `task="saa"` | Speaker-attributed text; requires a Plus checkpoint |
| `task="timestamps"` | Word timings in seconds; requires a Plus checkpoint |
| `word_timestamps=True` | Alias for timestamp mode on Plus checkpoints, ignored by 4.0; cannot be combined with `task="saa"` |
| `hotwords=["Acme", "QFormer"]` | Adds keyword hints to the task prompt |
| `system_prompt="..."` | Replaces the system turn that Plus checkpoints send by default |

Rich tasks use canonical prompts to select their output format, so `prompt=` cannot override `saa` or `timestamps`. If the model returns plain or malformed text instead of the requested tags, `generate()` raises `StructuredTranscriptError` (importable from `mlx_audio.stt.models.granite_speech.granite_speech`) with the model output on its `raw_text` attribute. The original 4.0 checkpoint also accepts `language="fr"` (or another supported target language) for translation.

## CLI

```bash
mlx_audio.stt.generate \
--model ibm-granite/granite-speech-4.1-2b-plus \
--audio meeting.wav \
--output-path transcript \
--format json \
--gen-kwargs '{"task": "saa", "hotwords": ["Acme", "QFormer"]}'
```

For subtitles, use `--format srt` or `--format vtt` with `--gen-kwargs '{"task": "timestamps"}'`; the file holds one cue for the whole utterance followed by one cue per word. Speaker-only segments have no timestamps, so use JSON to preserve their speaker labels.

The CLI passes `--language en` by default, which selects the translation prompt for `task="asr"`. To use the checkpoint's plain transcription prompt instead, add `"language": null` to `--gen-kwargs`.

## Streaming

```python
for chunk in model.generate("audio.wav", stream=True):
print(chunk.text, end="", flush=True)
```

Streaming yields decoder text after the supplied recording has been encoded; it does not ingest live audio. With `task="saa"` or `task="timestamps"` the stream carries the raw `[Speaker N]:` or `[T:N]` tags. Call `generate()` without `stream=True` to receive parsed segments.

## Audio and output limits

The [Plus model card](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) specifies up to nine minutes for ASR or speaker attribution and 3.5 minutes for word timestamps. Timestamp tags need more output tokens than plain text, so raise `max_tokens` for long recordings.
3 changes: 2 additions & 1 deletion docs/models/stt/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,8 @@ MLX Audio provides a range of speech-to-text models optimized for Apple Silicon,
| **Canary** | NVIDIA | ~1B | 25 EU + RU, UK | -- | -- | [README](https://github.com/Blaizzy/mlx-audio/blob/main/mlx_audio/stt/models/canary/README.md) |
| **SenseVoice** | Alibaba DAMO | ~234M | 50+ | -- | -- | [mlx-community/SenseVoiceSmall](https://huggingface.co/mlx-community/SenseVoiceSmall) |
| **FireRedASR2** | Xiaohongshu | ~1.18B | ZH, EN | -- | -- | [mlx-community/FireRedASR2-AED-mlx](https://huggingface.co/mlx-community/FireRedASR2-AED-mlx) |
| **Granite Speech** | IBM | ~1B | EN, FR, DE, ES, PT, JA | Yes | -- | [README](https://github.com/Blaizzy/mlx-audio/blob/main/mlx_audio/stt/models/granite_speech/README.md) |
| [**Granite Speech 4.0**](granite-speech.md) | IBM | ~1B | EN, FR, DE, ES, PT, JA | Yes | -- | [ibm-granite/granite-4.0-1b-speech](https://huggingface.co/ibm-granite/granite-4.0-1b-speech) |
| [**Granite Speech 4.1 Plus**](granite-speech.md) | IBM | ~2B | EN, FR, DE, ES, PT | Yes | Word (separate from speaker attribution) | [ibm-granite/granite-speech-4.1-2b-plus](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) |
| **Moonshine** | Useful Sensors | 27M / 61M | EN | -- | -- | [README](https://github.com/Blaizzy/mlx-audio/blob/main/mlx_audio/stt/models/moonshine/README.md) |
| **MMS** | Meta | 1B | 1000+ | -- | -- | [README](https://github.com/Blaizzy/mlx-audio/blob/main/mlx_audio/stt/models/mms/README.md) |

Expand Down
1 change: 1 addition & 0 deletions mkdocs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@ nav:
- VibeVoice ASR: models/stt/vibevoice-asr.md
- Fun-ASR-Nano: models/stt/fun-asr-nano.md
- Qwen2-Audio: models/stt/qwen2-audio.md
- Granite Speech: models/stt/granite-speech.md
- Speech-to-Speech:
- models/sts/index.md
- MiMo-Audio: models/sts/mimo-audio.md
Expand Down
16 changes: 10 additions & 6 deletions mlx_audio/stt/generate.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,10 @@ def _get_cues(segments):
]
cues = []
for s in segments.segments:
# Speaker-only segments (e.g. Granite SAA) carry no timing, so they
# cannot produce cues.
if "start" not in s:
continue
cues.append({"start": s["start"], "end": s["end"], "text": s["text"]})
if "words" in s and s["words"]:
for w in s["words"]:
Expand Down Expand Up @@ -214,12 +218,12 @@ def save_as_json(segments, output_path: str):
"segments": [],
}
for s in segments.segments:
seg = {
"text": s["text"],
"start": s["start"],
"end": s["end"],
"duration": s["end"] - s["start"],
}
seg = {"text": s["text"]}
# Speaker-only segments (e.g. Granite SAA) carry no timing
if "start" in s:
seg["start"] = s["start"]
seg["end"] = s["end"]
seg["duration"] = s["end"] - s["start"]
# Add word-level timestamps if available
if "words" in s and s["words"]:
seg["words"] = s["words"]
Expand Down
51 changes: 46 additions & 5 deletions mlx_audio/stt/models/granite_speech/README.md
Original file line number Diff line number Diff line change
@@ -1,14 +1,15 @@
# Granite Speech

MLX implementation of IBM's Granite Speech, a speech-to-text model that combines a CTC Conformer encoder with a Granite LLM decoder via a BLIP-2 QFormer projector. Supports ASR (transcription) and AST (speech translation).
MLX implementation of IBM's Granite Speech, a speech-to-text model that combines a CTC Conformer encoder with a Granite LLM decoder via a BLIP-2 QFormer projector. Supports ASR (transcription), AST (speech translation), and, with the plus checkpoint, speaker-attributed ASR and word-level timestamps.

## Available Models

| Model | Parameters | Description |
|-------|------------|-------------|
| [ibm-granite/granite-4.0-1b-speech](https://huggingface.co/ibm-granite/granite-4.0-1b-speech) | ~1B | Speech recognition and translation |
| [ibm-granite/granite-4.0-1b-speech](https://huggingface.co/ibm-granite/granite-4.0-1b-speech) | ~1B | Speech recognition and translation, keyword biasing |
| [ibm-granite/granite-speech-4.1-2b-plus](https://huggingface.co/ibm-granite/granite-speech-4.1-2b-plus) | ~2B | Rich transcription: speaker attribution, word timestamps, keyword biasing (EN, FR, DE, ES, PT) |

**Supported Languages:** English, French, German, Spanish, Portuguese, Japanese
**Supported Languages:** English, French, German, Spanish, Portuguese, Japanese (plus checkpoint: no Japanese)

## CLI Usage

Expand Down Expand Up @@ -78,17 +79,57 @@ print(result.text)

> **Note:** If the model receives an unfamiliar prompt, it falls back to transcription as the default mode.

### Rich Transcription (granite-speech-4.1-2b-plus)

The plus checkpoint selects its mode through the prompt; the `task` parameter picks the right one. Audio limits from the model card: up to 9 minutes for `asr`/`saa`, up to 3.5 minutes for `timestamps`. Timestamps mode emits roughly one tag per word, so budget `max_tokens` accordingly. The model card describes unpunctuated, lowercase output, but generated text can include punctuation and casing; the implementation preserves it.

The canonical `saa` and `timestamps` prompts define their output schemas and cannot be replaced with `prompt=`. Use `hotwords=` for contextual biasing, or use `task="asr"` when supplying a custom instruction. If the checkpoint returns plain or malformed text instead of the requested tags, `generate()` raises `StructuredTranscriptError`; the model output is available on its `raw_text` attribute. Requesting `saa` or `timestamps` from a 4.0/4.1 checkpoint raises `UnsupportedTranscriptionTask`. Both exceptions are importable from `mlx_audio.stt.models.granite_speech.granite_speech`.

Plus checkpoints send the system turn from the model card by default; pass `system_prompt=` to replace it.

```python
from mlx_audio.stt import load

model = load("ibm-granite/granite-speech-4.1-2b-plus")

# Speaker-attributed ASR: [Speaker N]: tags, parsed into segments
result = model.generate("meeting.wav", task="saa")
for seg in result.segments:
print(f"Speaker {seg['speaker_id']}: {seg['text']}")

# Word-level timestamps: [T:N] tags, parsed into at most one segment with word timings
result = model.generate("audio.wav", task="timestamps", max_tokens=8192)
for seg in result.segments:
for word in seg["words"]:
print(f"{word['word']}\t{word['start']:.2f}-{word['end']:.2f}s")

# Keyword biasing (names, technical terms) works with any task and checkpoint
result = model.generate("audio.wav", hotwords=["Nativ", "QFormer"])
```

From the CLI, `task` and `hotwords` go through `--gen-kwargs`:

```bash
mlx_audio.stt.generate --model ibm-granite/granite-speech-4.1-2b-plus \
--audio meeting.wav --output-path output --format json \
--gen-kwargs '{"task": "saa", "hotwords": ["Acme Ledger", "Q3 close"]}'
```

Speaker-attributed segments carry no timing, so save them as JSON; SRT and VTT output needs `task="timestamps"`.

### Streaming

```python
from mlx_audio.stt import load

model = load("ibm-granite/granite-4.0-1b-speech")

for text in model.generate("audio.wav", stream=True):
print(text, end="", flush=True)
for result in model.generate("audio.wav", stream=True):
print(result.text, end="", flush=True)
```

For `saa` and `timestamps`, the stream carries the raw tagged text; call `generate()` without `stream=True` to get parsed segments.

### Generation Parameters

```python
Expand Down
5 changes: 4 additions & 1 deletion mlx_audio/stt/models/granite_speech/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,10 @@

DETECTION_HINTS = {
"config_keys": {"encoder_config", "projector_config", "audio_token_index"},
"architectures": {"GraniteSpeechForConditionalGeneration"},
"architectures": {
"GraniteSpeechForConditionalGeneration",
"GraniteSpeechPlusForConditionalGeneration",
},
}

__all__ = [
Expand Down
6 changes: 5 additions & 1 deletion mlx_audio/stt/models/granite_speech/config.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import inspect
from dataclasses import dataclass, field
from typing import Dict, Optional
from typing import Dict, List, Optional


@dataclass
Expand All @@ -17,6 +17,10 @@ class EncoderConfig:
dropout: float = 0.1
conv_kernel_size: int = 15
conv_expansion_factor: int = 2
# 1-based Conformer block indices whose outputs are concatenated onto the
# final layer output (0 = post input_linear). granite-speech-4.1-2b-plus
# uses [3]; None/[] for 4.0 and 4.1.
cat_hidden_layers: Optional[List[int]] = None
model_type: str = "granite_speech_encoder"

@classmethod
Expand Down
Loading
Loading