From 0729d8f9c58a9112c02849f0cb4a5611381d5e99 Mon Sep 17 00:00:00 2001 From: pszemraj <74869040+pszemraj@users.noreply.github.com> Date: Tue, 29 Sep 2026 03:34:55 -0700 Subject: [PATCH 1/3] feat(granite_speech): support granite-speech-4.1-2b-plus checkpoints 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. --- .../stt/models/granite_speech/__init__.py | 5 +- mlx_audio/stt/models/granite_speech/config.py | 6 +- .../models/granite_speech/granite_speech.py | 93 ++++-- .../stt/tests/test_granite_speech_plus.py | 296 ++++++++++++++++++ mlx_audio/stt/utils.py | 1 + 5 files changed, 377 insertions(+), 24 deletions(-) create mode 100644 mlx_audio/stt/tests/test_granite_speech_plus.py diff --git a/mlx_audio/stt/models/granite_speech/__init__.py b/mlx_audio/stt/models/granite_speech/__init__.py index 4682774ac..f8888a71e 100644 --- a/mlx_audio/stt/models/granite_speech/__init__.py +++ b/mlx_audio/stt/models/granite_speech/__init__.py @@ -3,7 +3,10 @@ DETECTION_HINTS = { "config_keys": {"encoder_config", "projector_config", "audio_token_index"}, - "architectures": {"GraniteSpeechForConditionalGeneration"}, + "architectures": { + "GraniteSpeechForConditionalGeneration", + "GraniteSpeechPlusForConditionalGeneration", + }, } __all__ = [ diff --git a/mlx_audio/stt/models/granite_speech/config.py b/mlx_audio/stt/models/granite_speech/config.py index 61573d743..8522d5e7b 100644 --- a/mlx_audio/stt/models/granite_speech/config.py +++ b/mlx_audio/stt/models/granite_speech/config.py @@ -1,6 +1,6 @@ import inspect from dataclasses import dataclass, field -from typing import Dict, Optional +from typing import Dict, List, Optional @dataclass @@ -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 diff --git a/mlx_audio/stt/models/granite_speech/granite_speech.py b/mlx_audio/stt/models/granite_speech/granite_speech.py index aa61bf5f0..14ce57907 100644 --- a/mlx_audio/stt/models/granite_speech/granite_speech.py +++ b/mlx_audio/stt/models/granite_speech/granite_speech.py @@ -25,6 +25,14 @@ "ja": "Japanese", } +# Verbatim from the granite-speech-4.1-2b-plus model card. IBM's reference code +# sends this system turn; without one the plus chat template substitutes a +# generic assistant message the model was not trained against. +PLUS_SYSTEM_PROMPT = ( + "Knowledge Cutoff Date: April 2024.\nToday's Date: December 19, 2024.\n" + "You are Granite, developed by IBM. You are a helpful AI assistant" +) + @dataclass class StreamingResult: @@ -113,19 +121,17 @@ def __call__(self, x: mx.array, attention_dists: mx.array) -> mx.array: rel_pos_emb = self.rel_pos_emb(attention_dists) C = self.context_size - pos_attn = ( - mx.sum( - q[:, :, :, :, None, :] * rel_pos_emb[None, None, None, :, :, :], - axis=-1, - ) - * self.scale - ) + # Contract the head dimension directly. Expanding q and rel_pos_emb + # first creates a [B, blocks, heads, C, C, dim_head] temporary; for the + # supported nine-minute input that single allocation can exceed Metal's + # buffer-size limit by itself. + pos_attn = mx.einsum("bnhcd,crd->bnhcr", q, rel_pos_emb) * self.scale if remainder > 0: row_valid = mx.arange(C)[:, None] < remainder col_valid = mx.arange(C)[None, :] < remainder mask = ~(row_valid & col_valid) - mask_value = mx.array(mx.finfo(pos_attn.dtype).min) + mask_value = mx.array(mx.finfo(pos_attn.dtype).min, dtype=pos_attn.dtype) pos_attn_last = mx.where( mask[None, None, None], mask_value, pos_attn[:, -1:, :, :, :] ) @@ -224,11 +230,20 @@ def __init__(self, config: EncoderConfig): def __call__(self, x: mx.array) -> mx.array: x = self.input_linear(x) + cat_layers = set(self.config.cat_hidden_layers or ()) + exported = [x] if 0 in cat_layers else [] for idx, layer in enumerate(self.layers, start=1): x = layer(x, attention_dists=self._attention_dists) if idx == self.num_layers // 2: x_mid = self.out(x) x = x + self.out_mid(mx.softmax(x_mid, axis=-1)) + # Export after the mid-layer CTC injection: HF adds it in place to + # the tensor it has already exported, so an exported mid layer + # carries the injection. + if idx in cat_layers: + exported.append(x) + if exported: + x = mx.concatenate([*exported, x], axis=-1) return x @@ -440,6 +455,15 @@ def __init__(self, config: ModelConfig): def layers(self): return self.language_model.model.layers + @property + def is_plus(self) -> bool: + # mlx_audio.convert rewrites model_type to the module name + # ("granite_speech"), so detect the plus variant by its architectural + # fingerprint, which survives conversion. + return self.config.model_type == "granite_speech_plus" or bool( + self.config.encoder_config.cat_hidden_layers + ) + def make_cache(self) -> List[KVCache]: return [KVCache() for _ in range(len(self.layers))] @@ -474,6 +498,12 @@ def __call__( return logits / self.language_model.logits_scaling def get_audio_features(self, input_features: mx.array) -> mx.array: + # Plus runs its encoder in the loaded weight dtype, as HF does. 4.0/4.1 + # keep their established float32 activations (float32 features promote + # the weights). + encoder_dtype = self.encoder.input_linear.weight.dtype + if self.is_plus and input_features.dtype != encoder_dtype: + input_features = input_features.astype(encoder_dtype) encoder_output = self.encoder(input_features) projected = self.projector(encoder_output) return projected @@ -483,27 +513,27 @@ def model_quant_predicate(self, p: str, m: nn.Module) -> bool: @staticmethod def sanitize(weights: Dict[str, mx.array]) -> Dict[str, mx.array]: - already_converted = any("scales" in k for k in weights) - sanitized = {} for k, v in weights.items(): if "num_batches_tracked" in k: continue if ( - not already_converted - and any(name in k for name in ["up_conv", "down_conv", "depth_conv"]) - and "weight" in k + any(name in k for name in ["up_conv", "down_conv", "depth_conv"]) + and k.endswith("weight") and len(v.shape) == 3 ): # MLX Conv1d expects weights in (out_channels, kernel_size, in_channels) # layout, while PyTorch uses (out_channels, in_channels, kernel_size). - # Models converted from PyTorch checkpoints need transposing; models - # already saved in MLX-native layout (e.g. bf16 safetensors) do not. - # depth_conv (kernel > 1) needs the shape heuristic to distinguish - # PyTorch (out, 1, kernel) from MLX (out, kernel, 1). up_conv and - # down_conv always use kernel_size=1, so they are always transposed. - if "depth_conv" not in k or v.shape[-1] > v.shape[-2]: + # Use each convolution's singleton dimension to distinguish those + # layouts, making sanitization safe both during conversion and when + # loading the resulting unquantized checkpoint. + is_pointwise = "up_conv" in k or "down_conv" in k + pytorch_pointwise = is_pointwise and v.shape[-1] == 1 + pytorch_depthwise = ( + "depth_conv" in k and v.shape[1] == 1 and v.shape[-1] != 1 + ) + if pytorch_pointwise or pytorch_depthwise: v = v.transpose(0, 2, 1) sanitized[k] = v @@ -583,15 +613,24 @@ def _build_prompt( self, num_audio_tokens: int, user_prompt: str = None, + *, + system_prompt: Optional[str] = None, ) -> mx.array: if user_prompt is None: user_prompt = "can you transcribe the speech into a written format?" + # The plus checkpoint was trained with a separator after the audio + # placeholder; the 4.0/4.1 checkpoints concatenate the instruction. audio_placeholder = "<|audio|>" * num_audio_tokens - content = f"{audio_placeholder}{user_prompt}" + if self.is_plus: + content = f"{audio_placeholder} {user_prompt.lstrip()}" + else: + content = f"{audio_placeholder}{user_prompt}" if getattr(self._tokenizer, "chat_template", None): chat = [{"role": "user", "content": content}] + if system_prompt: + chat.insert(0, {"role": "system", "content": system_prompt}) prompt_str = self._tokenizer.apply_chat_template( chat, tokenize=False, add_generation_prompt=True ) @@ -635,6 +674,7 @@ def generate( repetition_context_size: int = 100, prompt: str = None, language: str = None, + system_prompt: Optional[str] = None, prefill_step_size: int = 2048, verbose: bool = False, stream: bool = False, @@ -644,6 +684,9 @@ def generate( lang_name = LANGUAGE_CODES.get(language.lower(), language) prompt = f"Translate the speech to {lang_name}." + if system_prompt is None and self.is_plus: + system_prompt = PLUS_SYSTEM_PROMPT + if stream: return self._stream_generate( audio, @@ -655,6 +698,7 @@ def generate( repetition_penalty=repetition_penalty, repetition_context_size=repetition_context_size, prompt=prompt, + system_prompt=system_prompt, prefill_step_size=prefill_step_size, verbose=verbose, ) @@ -672,7 +716,9 @@ def generate( audio_features = self.get_audio_features(input_features) mx.eval(audio_features) - prompt_ids = self._build_prompt(num_audio_tokens, prompt) + prompt_ids = self._build_prompt( + num_audio_tokens, prompt, system_prompt=system_prompt + ) inputs_embeds = self._build_inputs_embeds(prompt_ids, audio_features) mx.eval(inputs_embeds) @@ -734,6 +780,7 @@ def _stream_generate( repetition_penalty: Optional[float] = None, repetition_context_size: int = 100, prompt: str = None, + system_prompt: Optional[str] = None, prefill_step_size: int = 2048, verbose: bool = False, ) -> Generator[StreamingResult, None, None]: @@ -746,7 +793,9 @@ def _stream_generate( audio_features = self.get_audio_features(input_features) mx.eval(audio_features) - prompt_ids = self._build_prompt(num_audio_tokens, prompt) + prompt_ids = self._build_prompt( + num_audio_tokens, prompt, system_prompt=system_prompt + ) inputs_embeds = self._build_inputs_embeds(prompt_ids, audio_features) mx.eval(inputs_embeds) diff --git a/mlx_audio/stt/tests/test_granite_speech_plus.py b/mlx_audio/stt/tests/test_granite_speech_plus.py new file mode 100644 index 000000000..6e49c94bb --- /dev/null +++ b/mlx_audio/stt/tests/test_granite_speech_plus.py @@ -0,0 +1,296 @@ +"""Weight-free tests for granite-speech-4.1-2b-plus support. + +The plus checkpoint concatenates intermediate Conformer layer outputs onto the +final encoder output (``cat_hidden_layers``) and expects a system turn plus a +space between the audio placeholder and the instruction. +""" + +from types import SimpleNamespace + +import mlx.core as mx +import pytest + +from mlx_audio.stt.models.granite_speech.config import ( + EncoderConfig, + ModelConfig, + ProjectorConfig, + TextConfig, +) +from mlx_audio.stt.models.granite_speech.granite_speech import ( + PLUS_SYSTEM_PROMPT, + ConformerAttention, + CTCEncoder, + EncoderProjector, + Model, +) + +HIDDEN_DIM = 8 + + +def _tiny_encoder_config(**overrides): + params = dict( + input_dim=4, + num_layers=2, + hidden_dim=HIDDEN_DIM, + num_heads=2, + dim_head=4, + output_dim=4, + context_size=8, + max_pos_emb=16, + ) + params.update(overrides) + return EncoderConfig(**params) + + +class TestEncoderCatHiddenLayers: + def _run(self, cat_hidden_layers): + encoder = CTCEncoder(_tiny_encoder_config(cat_hidden_layers=cat_hidden_layers)) + out = encoder(mx.zeros((1, 8, 4))) + mx.eval(out) + return out + + def test_none_keeps_hidden_dim(self): + assert self._run(None).shape == (1, 8, HIDDEN_DIM) + + def test_empty_list_keeps_hidden_dim(self): + assert self._run([]).shape == (1, 8, HIDDEN_DIM) + + def test_single_layer_doubles_dim(self): + assert self._run([1]).shape == (1, 8, 2 * HIDDEN_DIM) + + def test_layer_zero_exports_input_linear_output(self): + assert self._run([0, 1]).shape == (1, 8, 3 * HIDDEN_DIM) + + @staticmethod + def _layer(encoder, idx, hidden): + return encoder.layers[idx - 1](hidden, attention_dists=encoder._attention_dists) + + @staticmethod + def _inject_mid_ctc(encoder, hidden): + return hidden + encoder.out_mid(mx.softmax(encoder.out(hidden), axis=-1)) + + def test_exported_mid_layer_includes_ctc_injection(self): + # With two layers, layer 1 is the mid layer. HF adds the CTC term in + # place, so the state it exports for that layer carries the injection. + encoder = CTCEncoder(_tiny_encoder_config(cat_hidden_layers=[1])) + x = mx.random.normal((1, 8, 4), key=mx.random.key(0)) + h1 = self._inject_mid_ctc( + encoder, self._layer(encoder, 1, encoder.input_linear(x)) + ) + h2 = self._layer(encoder, 2, h1) + + expected = mx.concatenate([h1, h2], axis=-1) + assert mx.allclose(encoder(x), expected, atol=1e-5).item() + + def test_exported_layer_before_mid_is_raw_layer_output(self): + encoder = CTCEncoder(_tiny_encoder_config(num_layers=4, cat_hidden_layers=[1])) + x = mx.random.normal((1, 8, 4), key=mx.random.key(0)) + h1 = self._layer(encoder, 1, encoder.input_linear(x)) + h2 = self._inject_mid_ctc(encoder, self._layer(encoder, 2, h1)) + h4 = self._layer(encoder, 4, self._layer(encoder, 3, h2)) + + expected = mx.concatenate([h1, h4], axis=-1) + assert mx.allclose(encoder(x), expected, atol=1e-5).item() + + def test_config_from_dict_keeps_cat_hidden_layers(self): + cfg = EncoderConfig.from_dict({"cat_hidden_layers": [3], "num_layers": 16}) + assert cfg.cat_hidden_layers == [3] + + def test_projector_consumes_concatenated_features(self): + config = ModelConfig( + encoder_config=_tiny_encoder_config(cat_hidden_layers=[1]), + projector_config=ProjectorConfig( + hidden_size=HIDDEN_DIM, + num_hidden_layers=1, + num_attention_heads=2, + intermediate_size=16, + encoder_hidden_size=2 * HIDDEN_DIM, + ), + text_config=TextConfig(hidden_size=12), + ) + encoder = CTCEncoder(config.encoder_config) + projector = EncoderProjector(config) + out = projector(encoder(mx.zeros((1, 8, 4)))) + mx.eval(out) + # 8 frames -> 1 window of 15 -> window_size // downsample_rate queries + num_queries = config.window_size // config.downsample_rate + assert out.shape == (1, num_queries, 12) + + +def test_non_aligned_attention_preserves_bfloat16(): + config = _tiny_encoder_config(context_size=8) + attention = ConformerAttention(config) + attention.set_dtype(mx.bfloat16) + + seq = mx.arange(config.context_size) + attention_dists = ( + mx.clip( + seq[:, None] - seq[None, :], + -config.context_size, + config.context_size, + ) + + config.max_pos_emb + ) + output = attention( + mx.zeros((1, config.context_size + 1, config.hidden_dim), dtype=mx.bfloat16), + attention_dists, + ) + mx.eval(output) + + assert output.dtype == mx.bfloat16 + + +class StubTokenizer: + """Records what _build_prompt renders; mimics the chat-template contract.""" + + def __init__(self, chat_template="{{ messages }}"): + self.chat_template = chat_template + self.last_prompt = None + + def apply_chat_template(self, chat, tokenize=False, add_generation_prompt=True): + parts = [ + f"<|start_of_role|>{m['role']}<|end_of_role|>{m['content']}<|end_of_text|>\n" + for m in chat + ] + if add_generation_prompt: + parts.append("<|start_of_role|>assistant<|end_of_role|>") + return "".join(parts) + + def encode(self, text): + self.last_prompt = text + return [0] + + +def _build_prompt( + tokenizer, num_audio_tokens=2, prompt="do the thing", *, is_plus=True, **kwargs +): + stub_model = SimpleNamespace(_tokenizer=tokenizer, is_plus=is_plus) + Model._build_prompt(stub_model, num_audio_tokens, prompt, **kwargs) + return tokenizer.last_prompt + + +class TestBuildPrompt: + def test_plus_placeholder_has_reference_space(self): + rendered = _build_prompt(StubTokenizer()) + assert "<|audio|><|audio|> do the thing<|end_of_text|>" in rendered + + def test_non_plus_placeholder_has_no_space(self): + rendered = _build_prompt(StubTokenizer(), is_plus=False) + assert "<|audio|><|audio|>do the thing<|end_of_text|>" in rendered + + def test_plus_placeholder_space_absorbs_leading_whitespace(self): + rendered = _build_prompt(StubTokenizer(), prompt=" do the thing") + assert "<|audio|><|audio|> do the thing<|end_of_text|>" in rendered + + def test_non_plus_prompt_is_verbatim(self): + rendered = _build_prompt( + StubTokenizer(), prompt=" do the thing", is_plus=False + ) + assert "<|audio|><|audio|> do the thing<|end_of_text|>" in rendered + + def test_system_turn_inserted_first(self): + rendered = _build_prompt(StubTokenizer(), system_prompt=PLUS_SYSTEM_PROMPT) + assert rendered.startswith( + f"<|start_of_role|>system<|end_of_role|>{PLUS_SYSTEM_PROMPT}" + ) + + def test_no_system_turn_by_default(self): + rendered = _build_prompt(StubTokenizer(), is_plus=False) + assert rendered.startswith("<|start_of_role|>user<|end_of_role|>") + + def test_no_template_fallback(self): + rendered = _build_prompt(StubTokenizer(chat_template=None)) + assert rendered == "USER: <|audio|><|audio|> do the thing\nASSISTANT:" + + +@pytest.mark.parametrize( + ("is_plus", "expected_system_prompt"), + [(True, PLUS_SYSTEM_PROMPT), (False, None)], +) +def test_generate_defaults_system_turn_for_plus_only(is_plus, expected_system_prompt): + forwarded = {} + + def stream_generate(audio, **kwargs): + forwarded.update(kwargs) + return iter(()) + + stub_model = SimpleNamespace( + is_plus=is_plus, + config=SimpleNamespace(model_type="granite_speech"), + _stream_generate=stream_generate, + ) + + assert list(Model.generate(stub_model, mx.zeros((16000,)), stream=True)) == [] + assert forwarded["system_prompt"] == expected_system_prompt + + +class TestIsPlus: + def _is_plus(self, config): + return Model.is_plus.fget(SimpleNamespace(config=config)) + + def test_by_model_type(self): + assert self._is_plus(ModelConfig(model_type="granite_speech_plus")) + + def test_by_cat_hidden_layers_after_conversion(self): + # mlx_audio.convert rewrites model_type to "granite_speech"; the + # architectural fingerprint must still identify the plus variant. + config = ModelConfig( + model_type="granite_speech", + encoder_config={"cat_hidden_layers": [3]}, + ) + assert self._is_plus(config) + + def test_non_plus(self): + assert not self._is_plus(ModelConfig()) + + +class TestSanitizeWeights: + @pytest.mark.parametrize( + ("name", "pytorch_shape", "mlx_shape"), + [ + ("up_conv", (16, 8, 1), (16, 1, 8)), + ("down_conv", (8, 16, 1), (8, 1, 16)), + ("depth_conv", (16, 1, 5), (16, 5, 1)), + ], + ) + def test_convolution_conversion_is_idempotent(self, name, pytorch_shape, mlx_shape): + key = f"encoder.layers.0.conv.{name}.weight" + source = {key: mx.zeros(pytorch_shape)} + + converted = Model.sanitize(source) + reloaded = Model.sanitize(converted) + + assert converted[key].shape == mlx_shape + assert reloaded[key].shape == mlx_shape + + +@pytest.mark.parametrize( + ("is_plus", "encoder_dtype", "expected_dtype"), + [ + (True, mx.float32, mx.float32), + (True, mx.bfloat16, mx.bfloat16), + # 4.0/4.1 keep float32 activations regardless of the weight dtype. + (False, mx.bfloat16, mx.float32), + ], +) +def test_audio_feature_dtype(is_plus, encoder_dtype, expected_dtype): + class RecordingEncoder: + def __init__(self): + self.input_linear = SimpleNamespace( + weight=mx.zeros((1,), dtype=encoder_dtype) + ) + self.input_dtype = None + + def __call__(self, features): + self.input_dtype = features.dtype + return features + + encoder = RecordingEncoder() + stub_model = SimpleNamespace( + encoder=encoder, projector=lambda features: features, is_plus=is_plus + ) + + output = Model.get_audio_features(stub_model, mx.zeros((1, 2, 4), dtype=mx.float32)) + + assert encoder.input_dtype == expected_dtype + assert output.dtype == expected_dtype diff --git a/mlx_audio/stt/utils.py b/mlx_audio/stt/utils.py index 69affd989..dc069ced8 100644 --- a/mlx_audio/stt/utils.py +++ b/mlx_audio/stt/utils.py @@ -98,6 +98,7 @@ def wired_limit(model: nn.Module, streams: Optional[List[mx.Stream]] = None): "mms": "mms", "granite_speech": "granite_speech", "granite_speech5_ctc": "granite_speech5_ctc", + "granite_speech_plus": "granite_speech", "granite_speech_nar": "granite_speech_nar", "qwen2_audio": "qwen2_audio", "mega_asr": "mega_asr", From ba1f13697fa0246fcfbc75b72dfd5ab5f6bce483 Mon Sep 17 00:00:00 2001 From: pszemraj <74869040+pszemraj@users.noreply.github.com> Date: Tue, 29 Sep 2026 03:38:00 -0700 Subject: [PATCH 2/3] feat(granite_speech): add speaker-attributed and word-timestamp tasks 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. --- mlx_audio/stt/generate.py | 16 +- .../models/granite_speech/granite_speech.py | 229 +++++++++++++- .../stt/tests/test_granite_speech_plus.py | 285 +++++++++++++++++- 3 files changed, 517 insertions(+), 13 deletions(-) diff --git a/mlx_audio/stt/generate.py b/mlx_audio/stt/generate.py index 6d6e20291..bff3eaaa8 100644 --- a/mlx_audio/stt/generate.py +++ b/mlx_audio/stt/generate.py @@ -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"]: @@ -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"] diff --git a/mlx_audio/stt/models/granite_speech/granite_speech.py b/mlx_audio/stt/models/granite_speech/granite_speech.py index 14ce57907..7ec475745 100644 --- a/mlx_audio/stt/models/granite_speech/granite_speech.py +++ b/mlx_audio/stt/models/granite_speech/granite_speech.py @@ -1,4 +1,5 @@ import math +import re import time from dataclasses import dataclass from pathlib import Path @@ -33,6 +34,205 @@ "You are Granite, developed by IBM. You are a helpful AI assistant" ) +# Task prompts verbatim from the model card. An unfamiliar or malformed prompt +# makes the model silently fall back to plain transcription, so these must not +# be reworded. +TASK_PROMPTS = { + "asr": "can you transcribe the speech into a written format?", + "saa": ( + "Speaker attribution: Transcribe and denote who is speaking by adding " + "[Speaker 1]: and [Speaker 2]: tags before speaker turns." + ), + "timestamps": ( + "Timestamps: Transcribe the speech. After each word, add a timestamp tag " + "showing the end time in centiseconds, e.g. hello [T:45] world [T:82]" + ), +} + +_SPEAKER_RE = re.compile(r"\[Speaker (\d+)\]:") +_TS_RE = re.compile(r"\[T:(\d+)\]") + + +class UnsupportedTranscriptionTask(ValueError): + """The loaded checkpoint cannot perform the requested transcription task.""" + + +class StructuredTranscriptError(RuntimeError): + """A rich transcription did not contain the requested structured syntax.""" + + def __init__(self, message: str, *, raw_text: str) -> None: + super().__init__(message) + self.raw_text = raw_text + + +def _normalize_task(task: str, *, word_timestamps: bool = False) -> str: + normalized = (task or "asr").lower().strip() + if normalized not in TASK_PROMPTS: + raise ValueError( + f"Unknown task {task!r}; expected one of {sorted(TASK_PROMPTS)}" + ) + if word_timestamps: + if normalized not in ("asr", "timestamps"): + raise ValueError(f"word_timestamps=True conflicts with task={normalized!r}") + normalized = "timestamps" + return normalized + + +def _parse_saa(text: str) -> List[dict]: + """Parse ``[Speaker N]:`` turns without inventing speaker metadata.""" + matches = list(_SPEAKER_RE.finditer(text)) + if not matches: + raise StructuredTranscriptError( + "Granite speaker-attribution mode produced no [Speaker N]: tags. The " + "model may have fallen back to plain ASR.", + raw_text=text, + ) + if text[: matches[0].start()].strip(): + raise StructuredTranscriptError( + "Granite SAA output begins with unattributed text before the first " + "[Speaker N]: tag.", + raw_text=text, + ) + + seen: List[int] = [] + segments: List[dict] = [] + for index, match in enumerate(matches): + speaker_id = int(match.group(1)) + body_end = matches[index + 1].start() if index + 1 < len(matches) else len(text) + body = text[match.end() : body_end].strip() + if not body: + raise StructuredTranscriptError( + f"[Speaker {speaker_id}]: has no associated transcript text.", + raw_text=text, + ) + if speaker_id not in seen: + expected = len(seen) + 1 + if speaker_id != expected: + raise StructuredTranscriptError( + "SAA speakers must be introduced in order: " + f"expected Speaker {expected}, got Speaker {speaker_id}.", + raw_text=text, + ) + seen.append(speaker_id) + segments.append({"speaker_id": speaker_id, "text": body}) + + return segments + + +def _resolve_timestamp_centiseconds(value: str, previous_cs: Optional[int]) -> int: + """Resolve a modulo-1000 timestamp while enforcing monotonic output.""" + current_cs = int(value) + if current_cs >= 1000: + if previous_cs is not None and current_cs < previous_cs: + raise ValueError( + "Absolute Granite timestamp moved backwards: " + f"{current_cs} < {previous_cs} centiseconds." + ) + return current_cs + if previous_cs is None: + return current_cs + + current_cs += (previous_cs // 1000) * 1000 + while current_cs < previous_cs: + current_cs += 1000 + return current_cs + + +def _timestamp_items(text: str) -> List[Tuple[str, str]]: + """Return validated ``(word, timestamp)`` pairs for one transcript string.""" + tags = _TS_RE.findall(text) + if not tags: + raise StructuredTranscriptError( + "Granite timestamp mode produced no [T:N] tags. The model may have " + "fallen back to plain ASR.", + raw_text=text, + ) + + parts = _TS_RE.split(text) + trailing = parts[-1].strip() + if trailing: + raise StructuredTranscriptError( + "Granite timestamp output ends with content lacking a [T:N] tag: " + f"{trailing!r}.", + raw_text=text, + ) + + items: List[Tuple[str, str]] = [] + for raw_token, raw_timestamp in zip(parts[0::2], tags): + token = raw_token.strip() + if not token: + raise StructuredTranscriptError( + f"[T:{raw_timestamp}] has no preceding word or '_' marker.", + raw_text=text, + ) + if token != "_" and len(token.split()) != 1: + raise StructuredTranscriptError( + "Expected exactly one word or '_' before each timestamp tag, " + f"but {token!r} precedes [T:{raw_timestamp}]. One or more " + "timestamp tags are missing.", + raw_text=text, + ) + items.append((token, raw_timestamp)) + return items + + +def _parse_timestamps(text: str) -> List[dict]: + """Parse a complete word-timestamp sequence without fabricating alignment.""" + cursor = 0.0 + previous_cs: Optional[int] = None + words = [] + + for token, raw_timestamp in _timestamp_items(text): + try: + current_cs = _resolve_timestamp_centiseconds(raw_timestamp, previous_cs) + except ValueError as exc: + raise StructuredTranscriptError(str(exc), raw_text=text) from exc + end = current_cs / 100.0 + # "_" marks silence: it advances the clock but is not a word. + if token != "_": + words.append({"word": token, "start": cursor, "end": end}) + cursor = end + previous_cs = current_cs + + if not words: + return [] + return [ + { + "text": " ".join(word["word"] for word in words), + "start": words[0]["start"], + "end": words[-1]["end"], + "words": words, + } + ] + + +def _resolve_prompt(task: str, prompt: Optional[str], language: Optional[str]) -> str: + # Rich tasks are prompt-controlled, so their canonical prompts are part of + # the output-schema contract. Custom prompts remain available for ASR. + task = _normalize_task(task) + if prompt is not None: + if task != "asr": + raise ValueError( + f"prompt cannot override task={task!r}. Use hotwords=[...] for " + "contextual biasing, or use task='asr' for an unconstrained " + "custom instruction." + ) + return prompt + if task != "asr": + return TASK_PROMPTS[task] + if language is not None: + lang_name = LANGUAGE_CODES.get(language.lower(), language) + return f"Translate the speech to {lang_name}." + return TASK_PROMPTS["asr"] + + +def _parse_segments(task: str, text: str) -> List[dict]: + if task == "saa": + return _parse_saa(text) + if task == "timestamps": + return _parse_timestamps(text) + return [] + @dataclass class StreamingResult: @@ -617,7 +817,7 @@ def _build_prompt( system_prompt: Optional[str] = None, ) -> mx.array: if user_prompt is None: - user_prompt = "can you transcribe the speech into a written format?" + user_prompt = TASK_PROMPTS["asr"] # The plus checkpoint was trained with a separator after the audio # placeholder; the 4.0/4.1 checkpoints concatenate the instruction. @@ -672,17 +872,35 @@ def generate( min_p: float = 0.0, repetition_penalty: Optional[float] = None, repetition_context_size: int = 100, + task: str = "asr", prompt: str = None, language: str = None, system_prompt: Optional[str] = None, + hotwords: Optional[List[str]] = None, + word_timestamps: bool = False, prefill_step_size: int = 2048, verbose: bool = False, stream: bool = False, **kwargs, ) -> Union[STTOutput, Generator[StreamingResult, None, None]]: - if prompt is None and language is not None: - lang_name = LANGUAGE_CODES.get(language.lower(), language) - prompt = f"Translate the speech to {lang_name}." + from mlx_audio.stt.utils import merge_hotwords + + # 4.0/4.1 have no timestamp mode. Like other models without word + # timings, they ignore the generic word_timestamps flag (the server + # forwards it to every STT model); an explicit rich task still raises. + task = _normalize_task(task, word_timestamps=word_timestamps and self.is_plus) + if task != "asr" and not self.is_plus: + raise UnsupportedTranscriptionTask( + f"task={task!r} requires a Granite Speech Plus checkpoint, but the " + f"loaded model has model_type={self.config.model_type!r}. Load " + "ibm-granite/granite-speech-4.1-2b-plus or use task='asr'." + ) + prompt = _resolve_prompt(task, prompt, language) + + # Granite biases toward rare vocabulary via an inline "Keywords:" clause. + keywords = merge_hotwords(None, hotwords) + if keywords: + prompt = f"{prompt} Keywords: {keywords}" if system_prompt is None and self.is_plus: system_prompt = PLUS_SYSTEM_PROMPT @@ -747,6 +965,7 @@ def generate( tokens.append(token) text = self._tokenizer.decode(tokens, skip_special_tokens=True) + segments = _parse_segments(task, text) elapsed = time.time() - start_time gen_tokens = len(tokens) @@ -759,7 +978,7 @@ def generate( return STTOutput( text=text, - segments=[], + segments=segments, prompt_tokens=prompt_tokens, generation_tokens=gen_tokens, total_tokens=prompt_tokens + gen_tokens, diff --git a/mlx_audio/stt/tests/test_granite_speech_plus.py b/mlx_audio/stt/tests/test_granite_speech_plus.py index 6e49c94bb..7b1c8089f 100644 --- a/mlx_audio/stt/tests/test_granite_speech_plus.py +++ b/mlx_audio/stt/tests/test_granite_speech_plus.py @@ -1,15 +1,19 @@ """Weight-free tests for granite-speech-4.1-2b-plus support. The plus checkpoint concatenates intermediate Conformer layer outputs onto the -final encoder output (``cat_hidden_layers``) and expects a system turn plus a -space between the audio placeholder and the instruction. +final encoder output (``cat_hidden_layers``), expects a system turn plus a +space between the audio placeholder and the instruction, and emits +``[Speaker N]:`` / ``[T:N]`` tags that are parsed into the repo's ``segments`` +schema. """ +import json from types import SimpleNamespace import mlx.core as mx import pytest +from mlx_audio.stt.generate import _get_cues, save_as_json from mlx_audio.stt.models.granite_speech.config import ( EncoderConfig, ModelConfig, @@ -18,10 +22,17 @@ ) from mlx_audio.stt.models.granite_speech.granite_speech import ( PLUS_SYSTEM_PROMPT, + TASK_PROMPTS, ConformerAttention, CTCEncoder, EncoderProjector, Model, + StructuredTranscriptError, + UnsupportedTranscriptionTask, + _parse_saa, + _parse_segments, + _parse_timestamps, + _resolve_prompt, ) HIDDEN_DIM = 8 @@ -202,6 +213,10 @@ def test_no_template_fallback(self): rendered = _build_prompt(StubTokenizer(chat_template=None)) assert rendered == "USER: <|audio|><|audio|> do the thing\nASSISTANT:" + def test_default_prompt_is_asr_task(self): + rendered = _build_prompt(StubTokenizer(), prompt=None) + assert TASK_PROMPTS["asr"] in rendered + @pytest.mark.parametrize( ("is_plus", "expected_system_prompt"), @@ -294,3 +309,269 @@ def __call__(self, features): assert encoder.input_dtype == expected_dtype assert output.dtype == expected_dtype + + +class TestOutputParsers: + def test_parse_saa_card_example(self): + text = ( + "[Speaker 1]: Hello how are you " + "[Speaker 2]: I'm fine and how are you feeling " + "[Speaker 1]: I feel wonderful" + ) + segments = _parse_saa(text) + assert [s["speaker_id"] for s in segments] == [1, 2, 1] + assert segments[0]["text"] == "Hello how are you" + assert segments[2]["text"] == "I feel wonderful" + assert all("start" not in s for s in segments) + + @pytest.mark.parametrize( + "text", + [ + "plain transcription without tags", + "intro without attribution [Speaker 1]: tagged turn", + "[Speaker 2]: hello", + "[Speaker 1]: hello [Speaker 3]: world", + "[Speaker 1]:", + ], + ) + def test_parse_saa_rejects_invalid_speaker_structure(self, text): + with pytest.raises(StructuredTranscriptError) as exc_info: + _parse_saa(text) + + assert exc_info.value.raw_text == text + + def test_parse_saa_allows_ordered_speaker_reentry(self): + assert _parse_saa( + "[Speaker 1]: hello " "[Speaker 2]: hi " "[Speaker 1]: welcome back" + ) == [ + {"speaker_id": 1, "text": "hello"}, + {"speaker_id": 2, "text": "hi"}, + {"speaker_id": 1, "text": "welcome back"}, + ] + + def test_parse_timestamps_rollover_and_silence(self): + segments = _parse_timestamps( + "hello [T:995] world [T:012] _ [T:100] again [T:150]" + ) + assert len(segments) == 1 + words = segments[0]["words"] + assert [w["word"] for w in words] == ["hello", "world", "again"] + assert [w["end"] for w in words] == pytest.approx([9.95, 10.12, 11.50]) + # the dropped silence still advances the next word's start + assert words[2]["start"] == pytest.approx(11.00) + assert segments[0]["start"] == 0.0 + assert segments[0]["end"] == pytest.approx(11.50) + assert segments[0]["text"] == "hello world again" + + def test_parse_timestamps_accepts_absolute_values_after_rollover(self): + segments = _parse_timestamps( + "first [T:995] second [T:012] third [T:1234] fourth [T:250]" + ) + + assert [word["end"] for word in segments[0]["words"]] == pytest.approx( + [9.95, 10.12, 12.34, 12.50] + ) + assert segments[0]["text"] == "first second third fourth" + + @pytest.mark.parametrize( + "text", + [ + "no tags here", + "hello world [T:100]", + "hello [T:50] [T:100]", + "hello [T:50] _", + ], + ) + def test_parse_timestamps_rejects_incomplete_word_alignment(self, text): + with pytest.raises(StructuredTranscriptError) as exc_info: + _parse_timestamps(text) + + assert exc_info.value.raw_text == text + + def test_parse_timestamps_rejects_backwards_absolute_value(self): + with pytest.raises(StructuredTranscriptError, match="moved backwards"): + _parse_timestamps("first [T:1234] second [T:1200]") + + def test_parse_segments_dispatch(self): + assert _parse_segments("asr", "[Speaker 1]: hi") == [] + assert _parse_segments("saa", "[Speaker 1]: hi") == [ + {"speaker_id": 1, "text": "hi"} + ] + assert _parse_segments("timestamps", "hi [T:50]")[0]["end"] == 0.5 + + +@pytest.mark.parametrize("task", ["saa", "timestamps"]) +def test_structured_output_rejects_plain_asr_fallback(task): + with pytest.raises(StructuredTranscriptError) as exc_info: + _parse_segments(task, "plain transcript") + + assert exc_info.value.raw_text == "plain transcript" + + +def test_timestamp_output_rejects_untimed_trailing_text(): + with pytest.raises(StructuredTranscriptError, match=r"lacking a \[T:N\] tag"): + _parse_segments("timestamps", "hello [T:50] trailing words") + + +class TestResolvePrompt: + def test_asr_permits_explicit_prompt(self): + assert _resolve_prompt("asr", "custom", "fr") == "custom" + + @pytest.mark.parametrize("task", ["saa", "timestamps"]) + def test_rich_task_rejects_explicit_prompt(self, task): + with pytest.raises(ValueError, match="prompt cannot override"): + _resolve_prompt(task, "custom", "fr") + + def test_rich_task_beats_translation_language(self): + assert _resolve_prompt("saa", None, "en") == TASK_PROMPTS["saa"] + assert _resolve_prompt("timestamps", None, "en") == TASK_PROMPTS["timestamps"] + + def test_asr_with_language_translates(self): + assert _resolve_prompt("asr", None, "fr") == "Translate the speech to French." + + def test_asr_default(self): + assert _resolve_prompt("asr", None, None) == TASK_PROMPTS["asr"] + + def test_unknown_task_raises(self): + with pytest.raises(ValueError, match="Unknown task"): + _resolve_prompt("diarize", None, None) + + +def test_generate_rejects_unknown_task(): + with pytest.raises(ValueError, match="Unknown task"): + Model.generate(SimpleNamespace(), None, task="diarize") + + +def _base_model_stub(): + return SimpleNamespace( + is_plus=False, + config=SimpleNamespace(model_type="granite_speech"), + ) + + +@pytest.mark.parametrize("task", ["saa", "timestamps"]) +@pytest.mark.parametrize("stream", [False, True]) +def test_base_checkpoint_rejects_rich_tasks_before_generation(task, stream): + with pytest.raises(UnsupportedTranscriptionTask, match="requires.*Plus"): + Model.generate(_base_model_stub(), mx.zeros((1,)), task=task, stream=stream) + + +def _streamed_prompt(is_plus=True, **generate_kwargs): + forwarded = {} + + def stream_generate(audio, **kwargs): + forwarded.update(kwargs) + return iter(()) + + stub_model = SimpleNamespace( + is_plus=is_plus, + config=SimpleNamespace( + model_type="granite_speech_plus" if is_plus else "granite_speech" + ), + _stream_generate=stream_generate, + ) + result = Model.generate( + stub_model, mx.zeros((16000,)), stream=True, **generate_kwargs + ) + assert list(result) == [] + return forwarded["prompt"] + + +def test_word_timestamps_alias_selects_rich_timestamp_task(): + assert _streamed_prompt(word_timestamps=True) == TASK_PROMPTS["timestamps"] + + +def test_base_checkpoint_ignores_word_timestamps_alias(): + prompt = _streamed_prompt(is_plus=False, word_timestamps=True) + assert prompt == TASK_PROMPTS["asr"] + + +def test_word_timestamps_alias_rejects_saa_conflict(): + with pytest.raises(ValueError, match="conflicts"): + Model.generate( + SimpleNamespace(is_plus=True), + mx.zeros((16000,)), + task="saa", + word_timestamps=True, + ) + + +def test_hotwords_append_keywords_clause_to_task_prompt(): + assert _streamed_prompt(task="saa", hotwords=["Acme", " QFormer "]) == ( + f"{TASK_PROMPTS['saa']} Keywords: Acme, QFormer" + ) + + +@pytest.mark.parametrize( + ("task", "pieces", "expected_segments"), + [ + ("asr", {10: "hello ", 11: "world"}, []), + ( + "saa", + {10: "[Speaker 1]: ", 11: "hello"}, + [{"speaker_id": 1, "text": "hello"}], + ), + ( + "timestamps", + {10: "hello ", 11: "[T:50]"}, + [ + { + "text": "hello", + "start": 0.0, + "end": 0.5, + "words": [{"word": "hello", "start": 0.0, "end": 0.5}], + } + ], + ), + ], +) +def test_generate_returns_parsed_segments(monkeypatch, task, pieces, expected_segments): + import mlx_audio.lm.generate as lm_generate + + class SequenceTokenizer: + eos_token_id = 99 + + def decode(self, token_ids, **kwargs): + return "".join(pieces[token_id] for token_id in token_ids) + + monkeypatch.setattr( + lm_generate, + "generate_step", + lambda **kwargs: iter([(10, None), (11, None), (99, None)]), + ) + stub_model = SimpleNamespace( + is_plus=True, + config=SimpleNamespace(model_type="granite_speech_plus"), + _tokenizer=SequenceTokenizer(), + _load_audio=lambda audio: audio, + _extract_features=lambda audio: (mx.zeros((1, 1, 160)), 1), + get_audio_features=lambda features: mx.zeros((1, 1, 4)), + _build_prompt=lambda *args, **kwargs: mx.array([0]), + _build_inputs_embeds=lambda *args, **kwargs: mx.zeros((1, 1, 4)), + ) + + result = Model.generate(stub_model, mx.zeros((1,)), task=task) + + assert result.text == "".join(pieces.values()) + assert result.segments == expected_segments + + +def test_untimed_speaker_segments_produce_no_cues(): + output = SimpleNamespace(segments=_parse_saa("[Speaker 1]: hi [Speaker 2]: hey")) + assert _get_cues(output) == [] + + +def test_json_writer_keeps_untimed_speaker_segments(tmp_path): + output = SimpleNamespace( + text="[Speaker 1]: hi [Speaker 2]: hey", + segments=_parse_saa("[Speaker 1]: hi [Speaker 2]: hey"), + ) + output_path = tmp_path / "speakers" + + save_as_json(output, str(output_path)) + + saved = json.loads(output_path.with_suffix(".json").read_text()) + assert saved["segments"] == [ + {"text": "hi", "speaker_id": 1}, + {"text": "hey", "speaker_id": 2}, + ] From 03155409f19596478cf4331b450df233b38a7746 Mon Sep 17 00:00:00 2001 From: pszemraj <74869040+pszemraj@users.noreply.github.com> Date: Tue, 29 Sep 2026 03:40:39 -0700 Subject: [PATCH 3/3] docs(granite_speech): document the Plus checkpoint 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. --- README.md | 2 +- docs/models/stt/granite-speech.md | 76 +++++++++++++++++++ docs/models/stt/index.md | 3 +- mkdocs.yml | 1 + mlx_audio/stt/models/granite_speech/README.md | 51 +++++++++++-- 5 files changed, 126 insertions(+), 7 deletions(-) create mode 100644 docs/models/stt/granite-speech.md diff --git a/README.md b/README.md index 3938b4836..ab8ca3846 100644 --- a/README.md +++ b/README.md @@ -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) | diff --git a/docs/models/stt/granite-speech.md b/docs/models/stt/granite-speech.md new file mode 100644 index 000000000..f82978d83 --- /dev/null +++ b/docs/models/stt/granite-speech.md @@ -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. diff --git a/docs/models/stt/index.md b/docs/models/stt/index.md index 198dc2cef..c2e161728 100644 --- a/docs/models/stt/index.md +++ b/docs/models/stt/index.md @@ -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) | diff --git a/mkdocs.yml b/mkdocs.yml index 418f6bd98..d5c50dcae 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -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 diff --git a/mlx_audio/stt/models/granite_speech/README.md b/mlx_audio/stt/models/granite_speech/README.md index 0d561d067..968ce5470 100644 --- a/mlx_audio/stt/models/granite_speech/README.md +++ b/mlx_audio/stt/models/granite_speech/README.md @@ -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 @@ -78,6 +79,44 @@ 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 @@ -85,10 +124,12 @@ 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