5858 OVWeightQuantizationConfig ,
5959 _merge_ignored_scopes ,
6060)
61- from .modeling import OVModelForFeatureExtraction , OVModelForMaskedLM , OVModelForZeroShotImageClassification
61+ from .modeling import OVModelForCTC , OVModelForFeatureExtraction , OVModelForMaskedLM , OVModelForZeroShotImageClassification
6262from .modeling_base import OVBaseModel
6363from .modeling_decoder import OVBaseDecoderModel , OVModelForCausalLM
6464from .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 ,
0 commit comments