Skip to content

feat(stt): add Granite Speech 4.1 Plus support - #978

Open
pszemraj wants to merge 3 commits into
Blaizzy:mainfrom
pszemraj:feat/granite-plus
Open

pszemraj wants to merge 3 commits into
Blaizzy:mainfrom
pszemraj:feat/granite-plus

Conversation

@pszemraj

@pszemraj pszemraj commented Sep 27, 2026 •

Copy link
Copy Markdown

Context

This PR adds support for IBM's granite-speech-4.1-2b-plus checkpoint to the existing granite_speech model package. Plus uses a wider encoder output and a system turn, and adds two prompt-selected tasks: speaker-attributed transcription and word-level timestamps. Their tagged output is parsed into STTOutput.segments. Loading, conversion and quantization go through the existing paths.

Requested in #779 (the Plus variant only; this does not close the broader request). Streaming detokenization is handled separately in #983.

Description

Checkpoint support.

  • Plus checkpoints have model_type: "granite_speech_plus", which MODEL_REMAPPING routes to granite_speech. GraniteSpeechPlusForConditionalGeneration is added to DETECTION_HINTS.
  • EncoderConfig.cat_hidden_layers (1-based block indices, 0 = the input_linear output) selects encoder states that are concatenated onto the final Conformer output before the projector. As in HF, an exported mid layer includes the mid-layer CTC injection. With cat_hidden_layers=None the encoder is unchanged.
  • Model.is_plus is true for that model_type or a non-empty cat_hidden_layers, so converted checkpoints are recognized after mlx_audio.convert rewrites model_type to granite_speech.
  • Plus prompts put a space between the audio placeholder and the instruction, strip leading whitespace from the instruction, and send the model card's system turn by default. system_prompt= overrides the system turn.
  • The Plus encoder runs in the loaded weight dtype, and the attention mask value is created in the activation dtype.

Tasks and output (Model.generate).

  • New keyword arguments: task="asr" | "saa" | "timestamps", word_timestamps=True (an alias for task="timestamps" on Plus), hotwords: List[str] and system_prompt.
  • Rich tasks use the model card's prompts verbatim. They take priority over language, and passing prompt= with a rich task raises ValueError.
  • hotwords are appended as the model's Keywords: clause using the existing merge_hotwords helper, for any task and any checkpoint.
  • Non-streaming task="saa" returns one segment per turn: {"speaker_id", "text"}, with no timing.
  • Non-streaming task="timestamps" returns one segment with start, end and per-word words. The model's modulo-1000 centisecond clock is resolved, and _ markers advance the clock without producing a word.
  • Output that lacks the requested tags, starts with untagged text, introduces speakers out of order, has a word without a timestamp, or has an absolute timestamp that moves backwards raises StructuredTranscriptError, with the model output on raw_text.
  • stream=True yields the tagged text as generated and does not parse segments.
  • New exceptions UnsupportedTranscriptionTask(ValueError) and StructuredTranscriptError(RuntimeError) live in mlx_audio.stt.models.granite_speech.granite_speech.

Changes in the codebase

  • mlx_audio/stt/models/granite_speech/: the model changes above, the cat_hidden_layers config field and the detection hint.

  • mlx_audio/stt/utils.py: one MODEL_REMAPPING entry.

  • mlx_audio/stt/generate.py:

    • _get_cues skips segments without start.
    • save_as_json writes start/end/duration only when a segment has timing.

    Segments from other models always carry timing, so their output is unchanged.

  • Docs:

    • new docs/models/stt/granite-speech.md and its mkdocs.yml nav entry
    • rows for the Plus checkpoint in README.md and docs/models/stt/index.md
    • rich-task usage in the model README, whose existing streaming example now iterates over StreamingResult.text

Behavior changes for Granite 4.0/4.1 checkpoints.

  • Prompts are unchanged, and the encoder still runs with float32 activations.
  • task="saa" or task="timestamps" raises UnsupportedTranscriptionTask before any audio is processed. Previously an unknown task keyword was ignored.
  • word_timestamps=True is ignored, as with other models that have no word timings.
  • hotwords now adds the Keywords: clause. The 4.0 model card documents this form of biasing.
  • sanitize() decides whether to transpose a conv weight from its singleton dimension, not from whether the checkpoint is quantized. Unquantized checkpoints written by mlx_audio.convert therefore reload without their pointwise convs being transposed a second time.
  • The relative-position attention score uses mx.einsum instead of broadcasting to a [B, blocks, heads, C, C, dim_head] temporary. This changes the float32 reduction order, so encoder features are not bit-identical to main. Transcripts in the checks below are identical.

Changes outside the codebase

None. No converted weights are published and no dependencies change.

Additional information

Testing. mlx_audio/stt/tests/test_granite_speech_plus.py is weight-free and uses tiny configs and a stub tokenizer. It covers:

  • encoder concatenation, including layer-0 and mid-layer exports checked against hand-computed references;
  • the projector consuming the wider features;
  • bf16 masking of non-aligned attention;
  • prompt construction for Plus and non-Plus;
  • is_plus after conversion, idempotent sanitize, and encoder dtype handling;
  • the SAA and timestamp parsers, including rollover, silence and malformed output;
  • prompt resolution, and rich-task rejection on base checkpoints for both streaming and non-streaming;
  • the word_timestamps alias and hotwords;
  • parsed segments through generate() with stubbed decoding;
  • the two generate.py writer guards.

The full mlx_audio suite passes on Apple Silicon except mlx_audio/music/tests/test_generate.py::test_main_reads_lyrics_file, which fails the same way on main.

Checked against real checkpoints:

  • ibm-granite/granite-speech-4.1-2b-plus (revision 1454e6e) on the first 8 s of examples/voice_prompts/en_man.wav with task set to asr, saa and timestamps, and with hotwords. For every task, the joined stream=True text equals the non-streaming text.
  • The same checkpoint on a 3-minute recording and on a two-voice clip synthesized with macOS say, for saa and timestamps. On the 3-minute recording the timestamp words are monotonic and within the audio duration.
  • bf16 and 4-bit conversions made with mlx_audio.convert load as Plus and produce the same parsed output.
  • ibm-granite/granite-4.0-1b-speech on main and on this branch: English, Japanese and a two-voice clip; language="fr", "de", "ja" and "en"; a custom prompt; word_timestamps=True; and streaming. Rendered prompts and transcripts are identical.

To reproduce:

from mlx_audio.stt.utils import load_model

model = load_model("ibm-granite/granite-speech-4.1-2b-plus")
for task in ("asr", "saa", "timestamps"):
    result = model.generate("examples/voice_prompts/en_man.wav", task=task)
    print(task, result.text, result.segments)

Out of scope.

  • stream=True does not parse segments.
  • There are no completion metadata for rich tasks: a rich task that hits max_tokens either raises StructuredTranscriptError or returns the segments parsed from the truncated text.
  • The HTTP server does not expose task, hotwords or system_prompt.
  • CLI and subtitle writing keep existing behavior:
    • --language defaults to en, which selects the translation prompt for task="asr". The docs describe passing "language": null through --gen-kwargs for plain transcription.
    • SRT/VTT for timestamps contain one utterance cue followed by per-word cues.
    • SRT/VTT for saa contain no cues, because speaker turns have no timing.
  • A native Plus checkpoint in a local directory relies on the existing loader's name matching to resolve granite_speech_plus. Hub repo IDs and mlx_audio.convert output load.

Checklist

Comment thread .github/workflows/granite-speech-plus-checkpoint.yml Outdated
Comment thread mlx_audio/lm/generate.py Outdated
Comment thread tests/model_parity/README.md Outdated

@lucasnewman lucasnewman left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@pszemraj Thanks for the contribution. This is far too invasive to consider at the moment. Please remove the parity testing and all related changes, put your tokenizer changes into a separate PR, and make this PR specifically about adding the model support.

@pszemraj

Copy link
Copy Markdown
Author

Thanks for the feedback! Yes will do. I assume I should do the tokenizer PR first? Then will clean this up accordingly and sort it out wrt tests

The plus encoder concatenates the outputs of the 1-based cat_hidden_layers indices (0 = post input_linear) onto the final Conformer layer output, so the QFormer projector cross-attends over the wider features. As in HF, an exported mid layer includes the mid-layer CTC injection. cat_hidden_layers=None keeps the 4.0/4.1 encoder unchanged. Route model_type granite_speech_plus to this module and add the plus architecture to DETECTION_HINTS.

Plus checkpoints get the model card's system turn by default (overridable via system_prompt) and a single space between the audio placeholder and the instruction; 4.0/4.1 prompts are unchanged.

Also make weight sanitization idempotent for converted unquantized checkpoints, run the plus encoder in the loaded weight dtype with the attention mask kept in that dtype (4.0/4.1 keep float32 activations), and contract the relative-position attention with einsum to avoid a [B, blocks, heads, C, C, dim_head] temporary on long inputs.
The plus checkpoint selects speaker attribution (task="saa") and word timestamps (task="timestamps", or word_timestamps=True) through its model-card prompts. Those canonical prompts take priority over language and cannot be replaced with prompt=. Rich tasks on 4.0/4.1 checkpoints raise UnsupportedTranscriptionTask before any audio is processed; those checkpoints ignore word_timestamps, as other models without word timings do. hotwords are appended as the model's "Keywords:" clause for any task and checkpoint.

Non-streaming generate() parses [Speaker N]: tags into speaker_id segments and [T:N] tags into one segment with word timings, resolving the modulo-1000 centisecond clock. Output that lacks the requested tags or is malformed raises StructuredTranscriptError carrying the raw text rather than returning fabricated structure. Streaming yields the tagged text unchanged.

The CLI writers skip untimed segments when building SRT/VTT cues and omit start/end/duration from JSON segments that have no timing, so speaker-only output can be saved without a KeyError.
Add a Granite Speech page to the STT docs covering the Plus tasks, hotwords, CLI usage, streaming, and the model card's audio limits, list the Plus checkpoint in the README and STT model tables, and extend the model README with rich transcription usage.
@pszemraj

Copy link
Copy Markdown
Author

@lucasnewman Thanks for the review. I've split this up as requested:

  • The parity tests, their CI workflow, and the related dependency and lockfile changes are removed.
  • The tokenizer and streaming detokenization changes are now in fix(stt): tokenizer-safe streaming detokenization #983.
  • This PR now contains only Granite Speech 4.1 Plus model support: 11 files, three commits.

The description is updated to match. Server/API plumbing and the unrelated fixes that were bundled here are left out, and I can send them separately later if they're wanted.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants