Conversation
lucasnewman
reviewed
Sep 28, 2026
lucasnewman
reviewed
Sep 28, 2026
lucasnewman
reviewed
Sep 28, 2026
lucasnewman
requested changes
Sep 28, 2026
lucasnewman
left a comment
Collaborator
There was a problem hiding this comment.
@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.
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 |
2 of 3 tasks
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
force-pushed
the
feat/granite-plus
branch
from
September 29, 2026 18:10
b2bd156 to
0315540
Compare
Author
|
@lucasnewman Thanks for the review. I've split this up as requested:
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Context
This PR adds support for IBM's
granite-speech-4.1-2b-pluscheckpoint to the existinggranite_speechmodel 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 intoSTTOutput.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.
model_type: "granite_speech_plus", whichMODEL_REMAPPINGroutes togranite_speech.GraniteSpeechPlusForConditionalGenerationis added toDETECTION_HINTS.EncoderConfig.cat_hidden_layers(1-based block indices,0= theinput_linearoutput) 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. Withcat_hidden_layers=Nonethe encoder is unchanged.Model.is_plusis true for thatmodel_typeor a non-emptycat_hidden_layers, so converted checkpoints are recognized aftermlx_audio.convertrewritesmodel_typetogranite_speech.system_prompt=overrides the system turn.Tasks and output (
Model.generate).task="asr" | "saa" | "timestamps",word_timestamps=True(an alias fortask="timestamps"on Plus),hotwords: List[str]andsystem_prompt.language, and passingprompt=with a rich task raisesValueError.hotwordsare appended as the model'sKeywords:clause using the existingmerge_hotwordshelper, for any task and any checkpoint.task="saa"returns one segment per turn:{"speaker_id", "text"}, with no timing.task="timestamps"returns one segment withstart,endand per-wordwords. The model's modulo-1000 centisecond clock is resolved, and_markers advance the clock without producing a word.StructuredTranscriptError, with the model output onraw_text.stream=Trueyields the tagged text as generated and does not parse segments.UnsupportedTranscriptionTask(ValueError)andStructuredTranscriptError(RuntimeError)live inmlx_audio.stt.models.granite_speech.granite_speech.Changes in the codebase
mlx_audio/stt/models/granite_speech/: the model changes above, thecat_hidden_layersconfig field and the detection hint.mlx_audio/stt/utils.py: oneMODEL_REMAPPINGentry.mlx_audio/stt/generate.py:_get_cuesskips segments withoutstart.save_as_jsonwritesstart/end/durationonly when a segment has timing.Segments from other models always carry timing, so their output is unchanged.
Docs:
docs/models/stt/granite-speech.mdand itsmkdocs.ymlnav entryREADME.mdanddocs/models/stt/index.mdStreamingResult.textBehavior changes for Granite 4.0/4.1 checkpoints.
task="saa"ortask="timestamps"raisesUnsupportedTranscriptionTaskbefore any audio is processed. Previously an unknowntaskkeyword was ignored.word_timestamps=Trueis ignored, as with other models that have no word timings.hotwordsnow adds theKeywords: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 bymlx_audio.converttherefore reload without their pointwise convs being transposed a second time.mx.einsuminstead 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 tomain. 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.pyis weight-free and uses tiny configs and a stub tokenizer. It covers:is_plusafter conversion, idempotentsanitize, and encoder dtype handling;word_timestampsalias andhotwords;generate()with stubbed decoding;generate.pywriter guards.The full
mlx_audiosuite passes on Apple Silicon exceptmlx_audio/music/tests/test_generate.py::test_main_reads_lyrics_file, which fails the same way onmain.Checked against real checkpoints:
ibm-granite/granite-speech-4.1-2b-plus(revision1454e6e) on the first 8 s ofexamples/voice_prompts/en_man.wavwithtaskset toasr,saaandtimestamps, and withhotwords. For every task, the joinedstream=Truetext equals the non-streaming text.say, forsaaandtimestamps. On the 3-minute recording the timestamp words are monotonic and within the audio duration.mlx_audio.convertload as Plus and produce the same parsed output.ibm-granite/granite-4.0-1b-speechonmainand 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:
Out of scope.
stream=Truedoes not parse segments.max_tokenseither raisesStructuredTranscriptErroror returns the segments parsed from the truncated text.task,hotwordsorsystem_prompt.--languagedefaults toen, which selects the translation prompt fortask="asr". The docs describe passing"language": nullthrough--gen-kwargsfor plain transcription.timestampscontain one utterance cue followed by per-word cues.saacontain no cues, because speaker turns have no timing.granite_speech_plus. Hub repo IDs andmlx_audio.convertoutput load.Checklist