Skip to content

Commit e1db71c

Browse files
committed
[OpenVINO] Add MedASR (LASR-CTC) model support
Add OpenVINO export, inference, and quantization support for google/medasr (model_type=lasr_ctc): Export: - Add LasrCtcOpenVINOConfig with custom DummyLasrCtcAudioInputGenerator (input_features [batch, time, features] + attention_mask) - Register lasr_ctc in TasksManager custom classes (AutoModelForCTC) Inference: - Update OVModelForCTC.forward() to handle input_features naming and conditionally pass attention_mask Quantization: - Add OVModelForCTC._preprocess_quantization_config() for automatic processor resolution (mirrors Whisper/Seq2Seq pattern) - Add OVModelForCTC branch in build_from_quantization_config() to route CTC models to speech-to-text calibration datasets - Add OVModelForCTC to build_from_dataset() isinstance check - Add _prepare_ctc_calibration_data() method for collecting audio calibration inputs via InferRequestWrapper - Add CTC model detection in _main_quantize() for weight compression Tests & Docs: - Add gated tests (RUN_SLOW_EXPORT_TESTS=1, transformers>=5.0) - Add MedASR entry to supported models documentation Verified on Intel Arc iGPU + CPU: - FP16 and INT8 weight-only: cosine sim >= 0.9999, token match >= 99% - INT8 full quantization (32 LibriSpeech samples): 2.9x CPU speedup
1 parent e6f612a commit e1db71c

8 files changed

Lines changed: 141 additions & 4 deletions

File tree

docs/source/openvino/models.mdx

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,7 @@ Here is the list of the supported architectures :
9898
- LongT5
9999
- M2M-100
100100
- MAIRA-2
101+
- MedASR (LASR-CTC)
101102
- Mamba
102103
- mBART
103104
- MPNet

optimum/exporters/openvino/__main__.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -727,6 +727,15 @@ def _main_quantize(
727727
except (AttributeError, ImportError, KeyError) as e:
728728
raise RuntimeError(f"Wasn't able to locate OpenVINO class for task {original_task} ({task}).") from e
729729

730+
731+
# For ASR task, detect CTC models (single openvino_model.xml) vs Seq2Seq (encoder/decoder)
732+
if task == "automatic-speech-recognition" and model_cls_name == "OVModelForSpeechSeq2Seq":
733+
if (Path(output) / "openvino_model.xml").exists() and not (
734+
Path(output) / "openvino_encoder_model.xml"
735+
).exists():
736+
model_cls_name = "OVModelForCTC"
737+
model_cls = getattr(__import__("optimum.intel", fromlist=[model_cls_name]), model_cls_name)
738+
730739
# Step 2. Load the exported model
731740
model = model_cls.from_pretrained(
732741
output,

optimum/exporters/openvino/model_configs.py

Lines changed: 52 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -277,6 +277,10 @@ def init_model_configs():
277277
"transformers",
278278
"AutoModelForCausalLM",
279279
)
280+
TasksManager._CUSTOM_CLASSES[("pt", "lasr_ctc", "automatic-speech-recognition")] = (
281+
"transformers",
282+
"AutoModelForCTC",
283+
)
280284

281285
# since transformers v4.46, model can be loaded using default AutoModelForImageTextToText
282286
# https://github.com/huggingface/transformers/blob/v4.46.0/src/transformers/models/auto/modeling_auto.py#L776
@@ -5156,10 +5160,57 @@ class Wav2Vec2OpenVINOConfig(HubertOpenVINOConfig):
51565160
"audio-xvector",
51575161
],
51585162
)
5159-
class Wav2Vec2ConformerOpenVINOConfig(HubertOpenVINOConfig):
5163+
class Wav2Vec2ConformerOpenVINOConfig(Wav2Vec2ConformerOnnxConfig):
51605164
pass
51615165

51625166

5167+
class DummyLasrCtcAudioInputGenerator(DummyAudioInputGenerator):
5168+
SUPPORTED_INPUT_NAMES = ("input_features", "attention_mask")
5169+
5170+
def generate(self, input_name: str, framework: str = "pt", int_dtype: str = "int64", float_dtype: str = "fp32"):
5171+
if input_name == "attention_mask":
5172+
return self.random_mask_tensor(
5173+
shape=[self.batch_size, self.nb_max_frames],
5174+
framework=framework,
5175+
dtype=int_dtype,
5176+
)
5177+
# MedASR expects input_features shape [batch, time, features] unlike most audio models
5178+
return self.random_float_tensor(
5179+
shape=[self.batch_size, self.nb_max_frames, self.feature_size],
5180+
min_value=-1,
5181+
max_value=1,
5182+
framework=framework,
5183+
dtype=float_dtype,
5184+
)
5185+
5186+
5187+
@register_in_tasks_manager("lasr_ctc", *["feature-extraction", "automatic-speech-recognition", "audio-classification"])
5188+
class LasrCtcOpenVINOConfig(OnnxConfig):
5189+
NORMALIZED_CONFIG_CLASS = NormalizedConfig.with_args(
5190+
hidden_size="encoder_config.hidden_size",
5191+
num_attention_heads="encoder_config.num_attention_heads",
5192+
num_hidden_layers="encoder_config.num_hidden_layers",
5193+
allow_new=True,
5194+
feature_size="encoder_config.num_mel_bins",
5195+
)
5196+
DUMMY_INPUT_GENERATOR_CLASSES = (DummyLasrCtcAudioInputGenerator,)
5197+
5198+
@property
5199+
def inputs(self):
5200+
return {
5201+
"input_features": {0: "batch_size", 1: "sequence_length", 2: "feature_size"},
5202+
"attention_mask": {0: "batch_size", 1: "sequence_length"},
5203+
}
5204+
5205+
@property
5206+
def outputs(self):
5207+
return {"logits": {0: "batch_size", 1: "sequence_length"}}
5208+
5209+
5210+
@register_in_tasks_manager("hubert", *["feature-extraction", "automatic-speech-recognition", "audio-classification"])
5211+
class HubertOpenVINOConfig(HubertOnnxConfig):
5212+
pass
5213+
51635214
@register_in_tasks_manager("sew", *["feature-extraction", "automatic-speech-recognition", "audio-classification"])
51645215
class SEWOpenVINOConfig(HubertOpenVINOConfig):
51655216
pass

optimum/intel/openvino/modeling.py

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -728,6 +728,16 @@ class OVModelForCTC(OVModel):
728728
auto_model_class = AutoModelForCTC
729729
export_feature = "automatic-speech-recognition"
730730

731+
def _preprocess_quantization_config(
732+
self,
733+
quantization_config: OVQuantizationConfigBase,
734+
model_name_or_path: str,
735+
) -> OVQuantizationConfigBase:
736+
if model_name_or_path is not None and quantization_config.processor is None:
737+
quantization_config = quantization_config.clone()
738+
quantization_config.processor = model_name_or_path
739+
return quantization_config
740+
731741
@add_start_docstrings_to_model_forward(
732742
AUDIO_INPUTS_DOCSTRING.format("batch_size, sequence_length")
733743
+ CTC_EXAMPLE.format(
@@ -742,13 +752,20 @@ def forward(
742752
attention_mask: Optional[Union[torch.Tensor, np.ndarray]] = None,
743753
**kwargs,
744754
):
755+
# Support models using input_features (e.g. MedASR/lasr_ctc) instead of input_values
756+
input_features = kwargs.get("input_features", None)
757+
if input_values is None and input_features is not None:
758+
input_values = input_features
759+
745760
np_inputs = isinstance(input_values, np.ndarray)
746761

747762
input_values = ensure_numpy(input_values)
748763
attention_mask = ensure_numpy(attention_mask)
749764

765+
# Use the actual input name expected by the model
766+
input_name = "input_features" if "input_features" in self.input_names else "input_values"
750767
inputs = {
751-
"input_values": input_values,
768+
input_name: input_values,
752769
}
753770

754771
# Add the attention_mask when needed

optimum/intel/openvino/quantization.py

Lines changed: 47 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@
5858
OVWeightQuantizationConfig,
5959
_merge_ignored_scopes,
6060
)
61-
from .modeling import OVModelForFeatureExtraction, OVModelForMaskedLM, OVModelForZeroShotImageClassification
61+
from .modeling import OVModelForCTC, OVModelForFeatureExtraction, OVModelForMaskedLM, OVModelForZeroShotImageClassification
6262
from .modeling_base import OVBaseModel
6363
from .modeling_decoder import OVBaseDecoderModel, OVModelForCausalLM
6464
from .modeling_sam import OVSamModel
@@ -281,6 +281,22 @@ def build_from_quantization_config(self, config: OVQuantizationConfigBase) -> OV
281281

282282
if isinstance(self.model, OVModelForCausalLM):
283283
return self._prepare_causal_lm_calibration_data(config)
284+
elif isinstance(self.model, OVModelForCTC):
285+
if config.processor is None:
286+
raise ValueError(
287+
"`processor` must be specified in order to run data-aware quantization. Please provide it as a"
288+
"model id, or a path to a directory containing all the required configuration files."
289+
)
290+
dataset_metadata = PREDEFINED_SPEECH_TO_TEXT_DATASETS[config.dataset]
291+
return self.build_from_dataset_name(
292+
config,
293+
dataset_metadata["id"],
294+
num_samples=config.num_samples,
295+
dataset_split=dataset_metadata["split"],
296+
streaming=dataset_metadata["streaming"],
297+
data_dir=dataset_metadata.get("data_dir", None),
298+
revision=dataset_metadata.get("revision", None),
299+
)
284300
elif isinstance(
285301
self.model,
286302
(OVModelForVisualCausalLM, _OVModelForWhisper, OVModelForZeroShotImageClassification, OVSamModel),
@@ -486,6 +502,7 @@ def build_from_dataset(
486502
isinstance(
487503
self.model,
488504
(
505+
OVModelForCTC,
489506
OVModelForVisualCausalLM,
490507
_OVModelForWhisper,
491508
OVModelForFeatureExtraction,
@@ -506,7 +523,9 @@ def build_from_dataset(
506523
"`batch_size`, `data_collator` and `remove_unused_columns` are not supported for this type of model."
507524
)
508525

509-
if isinstance(self.model, OVModelForVisualCausalLM):
526+
if isinstance(self.model, OVModelForCTC):
527+
return self._prepare_ctc_calibration_data(quantization_config, dataset)
528+
elif isinstance(self.model, OVModelForVisualCausalLM):
510529
return self._prepare_visual_causal_lm_calibration_data(quantization_config, dataset)
511530
elif isinstance(self.model, _OVModelForWhisper):
512531
return self._prepare_speech_to_text_calibration_data(quantization_config, dataset)
@@ -955,6 +974,32 @@ def _prepare_speech_to_text_calibration_data(
955974

956975
return OVCalibrationDataset(collected_inputs)
957976

977+
def _prepare_ctc_calibration_data(
978+
self, config: OVQuantizationConfigBase, dataset: "Dataset"
979+
) -> OVCalibrationDataset:
980+
"""
981+
Prepares calibration data for CTC (Connectionist Temporal Classification) models by processing audio samples.
982+
"""
983+
collected_inputs = []
984+
self.model.compile()
985+
self.model.request = InferRequestWrapper(self.model.request, collected_inputs, apply_caching=True)
986+
987+
try:
988+
processor = AutoProcessor.from_pretrained(config.processor, trust_remote_code=self.trust_remote_code)
989+
990+
num_samples = config.num_samples or 32
991+
dataset = list(tqdm(dataset.take(num_samples), desc="Downloading audio inputs", total=num_samples))
992+
993+
for item in tqdm(dataset, desc="Collecting calibration data"):
994+
audio = item["audio"]["array"]
995+
sampling_rate = item["audio"]["sampling_rate"]
996+
inputs = processor(audio, sampling_rate=sampling_rate, return_tensors="pt")
997+
self.model(**inputs)
998+
finally:
999+
self.model.request = self.model.request.request
1000+
1001+
return OVCalibrationDataset({"model": nncf.Dataset(collected_inputs)})
1002+
9581003
def _prepare_text_to_text_calibration_data(
9591004
self,
9601005
config: OVQuantizationConfigBase,

tests/openvino/test_export.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
# limitations under the License.
1414

1515

16+
import os
1617
import unittest
1718
from pathlib import Path
1819

@@ -36,6 +37,7 @@
3637
OVLTXPipeline,
3738
OVModelForAudioClassification,
3839
OVModelForCausalLM,
40+
OVModelForCTC,
3941
OVModelForCustomTasks,
4042
OVModelForFeatureExtraction,
4143
OVModelForImageClassification,
@@ -134,6 +136,9 @@ class ExportModelTest(unittest.TestCase):
134136
if is_transformers_version(">=", "5.0"):
135137
SUPPORTED_ARCHITECTURES.update({"lfm2_moe": OVModelForCausalLM})
136138

139+
if os.environ.get("RUN_SLOW_EXPORT_TESTS") == "1" and is_transformers_version(">=", "5.0"):
140+
SUPPORTED_ARCHITECTURES.update({"lasr_ctc": OVModelForCTC})
141+
137142
EXPECTED_DIFFUSERS_SCALE_FACTORS = {
138143
"stable-diffusion-xl": {"vae_encoder": "128.0", "vae_decoder": "128.0"},
139144
"stable-diffusion-3": {"text_encoder_3": "8.0"},

tests/openvino/test_quantization.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
# limitations under the License.
1414
import dataclasses
1515
import inspect
16+
import os
1617

1718
# ruff: noqa
1819

@@ -45,6 +46,7 @@
4546
OVLatentConsistencyModelPipeline,
4647
OVModelForAudioClassification,
4748
OVModelForCausalLM,
49+
OVModelForCTC,
4850
OVModelForFeatureExtraction,
4951
OVModelForImageClassification,
5052
OVModelForMaskedLM,
@@ -1103,6 +1105,9 @@ class OVWeightCompressionTest(unittest.TestCase):
11031105
SUPPORTED_ARCHITECTURES_WITH_AUTO_COMPRESSION.append((OVModelForVisualCausalLM, "gemma4", False))
11041106
SUPPORTED_ARCHITECTURES_WITH_AUTO_COMPRESSION.append((OVModelForVisualCausalLM, "gemma4_moe", False))
11051107

1108+
if os.environ.get("RUN_SLOW_EXPORT_TESTS") == "1" and is_transformers_version(">=", "5.0"):
1109+
SUPPORTED_ARCHITECTURES_WITH_AUTO_COMPRESSION.append((OVModelForCTC, "lasr_ctc", True))
1110+
11061111
SUPPORTED_ARCHITECTURES_WITH_HYBRID_QUANTIZATION = [
11071112
(OVStableDiffusionPipeline, "stable-diffusion", 72, 195),
11081113
(OVStableDiffusionXLPipeline, "stable-diffusion-xl", 84, 331),

tests/openvino/utils_tests.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -356,6 +356,9 @@ def _create_tiny_kokoro_model():
356356
"videochat_flash_qwen": "optimum-intel-internal-testing/tiny-videochat-flash-qwen",
357357
}
358358

359+
if os.environ.get("RUN_SLOW_EXPORT_TESTS") == "1" and is_transformers_version(">=", "5.0"):
360+
MODEL_NAMES["lasr_ctc"] = "google/medasr"
361+
359362
EAGLE3_MODELS = {"qwen3_eagle3": ("AngelSlim/Qwen3-1.7B_eagle3", "Qwen/Qwen3-1.7B")}
360363

361364
# VLM-based Eagle3 draft models (AngelSlim Eagle3LlamaForCausalLM architecture).
@@ -576,6 +579,7 @@ def _create_tiny_kokoro_model():
576579
"qwen3_eagle3",
577580
"qwen3_vl_eagle3",
578581
"qwen3_asr",
582+
"lasr_ctc",
579583
"videochat_flash_qwen",
580584
)
581585

0 commit comments

Comments
 (0)